From 63cc8c28fa19ab12918f6fcd663595a0fa41b391 Mon Sep 17 00:00:00 2001 From: Qinren Zhou Date: Tue, 15 Sep 2026 19:30:46 +0800 Subject: [PATCH 01/11] feat: relax name validation and harden recovery --- python/tests/detail/params_helper.py | 25 +- python/tests/detail/test_collection_ddl.py | 23 +- python/tests/detail/test_collection_dml.py | 12 +- python/tests/test_name_validation.py | 245 +++++++ python/zvec/model/_validation.py | 29 + python/zvec/model/collection.py | 36 +- python/zvec/model/convert.py | 30 +- python/zvec/model/schema/collection_schema.py | 29 +- python/zvec/model/schema/field_schema.py | 54 +- src/binding/c/c_api.cc | 202 +++--- src/binding/python/model/python_doc.cc | 6 +- src/db/collection.cc | 97 +-- src/db/common/constants.h | 14 +- src/db/index/common/doc.cc | 571 +++++++--------- .../index/common/manifest/manifest_codec.cc | 24 +- src/db/index/common/manifest_codec.h | 3 +- src/db/index/common/name_validation.cc | 171 +++++ src/db/index/common/name_validation.h | 43 ++ src/db/index/common/schema.cc | 121 ++-- src/db/index/segment/segment.cc | 315 ++++++--- src/db/index/storage/memory_forward_store.cc | 2 +- src/db/index/storage/wal/local_wal_file.cc | 240 ++++--- src/db/index/storage/wal/local_wal_file.h | 15 +- src/db/index/storage/wal/wal_file.h | 9 +- src/db/reranker/reranker.cc | 2 +- src/db/sqlengine/common/util.h | 3 +- src/include/zvec/db/doc.h | 5 - src/include/zvec/db/schema.h | 9 + tests/c/c_api_test.c | 309 +++++++++ .../relaxed_validation_recovery_test.cc | 646 ++++++++++++++++++ tests/db/index/common/doc_test.cc | 174 ++++- .../common/manifest_codec_golden_test.cc | 94 +++ tests/db/index/common/name_validation_test.cc | 289 ++++++++ tests/db/index/common/schema_test.cc | 90 ++- tests/db/index/storage/wal_file_test.cc | 231 +++++-- tests/db/relaxed_validation_test.cc | 412 +++++++++++ 36 files changed, 3660 insertions(+), 920 deletions(-) create mode 100644 python/tests/test_name_validation.py create mode 100644 python/zvec/model/_validation.py create mode 100644 src/db/index/common/name_validation.cc create mode 100644 src/db/index/common/name_validation.h create mode 100644 tests/db/crash_recovery/relaxed_validation_recovery_test.cc create mode 100644 tests/db/index/common/name_validation_test.cc create mode 100644 tests/db/relaxed_validation_test.cc diff --git a/python/tests/detail/params_helper.py b/python/tests/detail/params_helper.py index 24e7d891a..3a9f339c6 100644 --- a/python/tests/detail/params_helper.py +++ b/python/tests/detail/params_helper.py @@ -139,7 +139,7 @@ for param in params ] -COLLECTION_NAME_MAX_LENGTH = 64 +COLLECTION_NAME_MAX_LENGTH = 256 COLLECTION_NAME_VALID_LIST = [ "col", @@ -148,17 +148,20 @@ "collection_2", "123collection-", "a" * COLLECTION_NAME_MAX_LENGTH, -] - -COLLECTION_NAME_INVALID_LIST = [ "l", "1C", - "", " ", - None, - "abcdefghijklmnopqrstuvwxzy123456abcdefghijklmnopqrstuvwxzy1234561", "test/", "!@#$%^&*()test", + "集合名称", +] + +COLLECTION_NAME_INVALID_LIST = [ + "", + None, + "a" * (COLLECTION_NAME_MAX_LENGTH + 1), + "collection\0name", + "collection\nname", ] FIELD_NAME_VALID_LIST = [ @@ -178,7 +181,7 @@ "", " ", None, - "abcdefghijklmnopqrstuvwxzy1234561", + "a" * 65, "test/", "!@#$%^&*()test", "name@with#special$chars", @@ -208,12 +211,12 @@ INCOMPATIBLE_CONSTRUCTOR_ERROR_MSG = "incompatible constructor arguments" -SCHEMA_VALIDATE_ERROR_MSG = "schema validate failed" +SCHEMA_VALIDATE_ERROR_MSG = "Invalid schema" CREATE_READ_ONLY_ERROR_MSG = "Unable to create collection with read-only mode" INCOMPATIBLE_FUNCTION_ERROR_MSG = "incompatible function arguments" INVALID_PATH_ERROR_MSG = "path validate failed" INDEX_NON_EXISTENT_COLUMN_ERROR_MSG = "not found in schema" ACCESS_DESTROYED_COLLECTION_ERROR_MSG = "is already destroyed" COLLECTION_PATH_NOT_EXIST_ERROR_MSG = "not exist" -NOT_SUPPORT_ADD_COLUMN_ERROR_MSG = "Only support basic numeric data type" -NOT_EXIST_COLUMN_TO_DROP_ERROR_MSG = "Column not exists" +NOT_SUPPORT_ADD_COLUMN_ERROR_MSG = "this operation requires a numeric field" +NOT_EXIST_COLUMN_TO_DROP_ERROR_MSG = "field.*not found" diff --git a/python/tests/detail/test_collection_ddl.py b/python/tests/detail/test_collection_ddl.py index a2c682f46..b3e07ebd8 100644 --- a/python/tests/detail/test_collection_ddl.py +++ b/python/tests/detail/test_collection_ddl.py @@ -457,7 +457,7 @@ def check_error_message(exc_info, invalid_name): [ ("", ""), # Empty string (" ", " "), # Space only - ("v" * 33, "v" * 33), # Too long (33 characters, exceeds 32) + ("v" * 33, "v" * 33), # Field does not exist. ("vector name", "vector_name"), # Contains space ("vector@name", "vector@name"), # Contains special character ("vector/name", "vector/name"), # Contains slash @@ -1202,7 +1202,7 @@ def test_add_column_with_index_param(self, basic_collection: Collection): "field_name", [ "a", # Minimum length - "a" * 32, # Maximum length (32 characters) + "a" * 64, # Maximum field name length. "valid_field_name_123", # Alphanumeric with underscore "Valid-Field-Name", # With hyphens "_underscore_start", # Starting with underscore @@ -1241,7 +1241,7 @@ def test_add_column_with_valid_field_names( [ "", # Empty string " ", # Space only - "a" * 33, # Too long (33 characters, exceeds 32) + "a" * 65, # Exceeds the field name byte limit. "field name", # Contains space "field.name", # Contains dot "field@name", # Contains special character @@ -1263,7 +1263,7 @@ def test_add_column_with_invalid_field_names( ) if invalid_field_name is None: - assert "validate failed" in str(exc_info.value), ( + assert "Invalid schema:" in str(exc_info.value), ( "Error message is unreasonable: e=" + str(exc_info.value) ) else: @@ -1303,7 +1303,7 @@ def test_alter_column_non_exist(self, basic_collection: Collection): new_name="new_name", field_schema=FieldSchema("new_name", DataType.STRING), ) - assert "column non_existing not found" in str(exc_info.value), ( + assert "Invalid schema: field[non_existing] not found" in str(exc_info.value), ( "Error message is unreasonable: e=" + str(exc_info.value) ) @@ -1378,9 +1378,9 @@ def test_alter_column_with_various_concurrency_options( [ ("a", "new_a"), # Minimum length ( - "abcdefghijklmnopqrstuvwxyz123456", - "aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa", - ), # Maximum length (32 characters) + "a" * 64, + "b" * 64, + ), # Maximum field name length. ("valid_field_name_123", "new_valid_field"), # Alphanumeric with underscore ("Valid-Field-Name", "New-Field-Name"), # With hyphens ("_underscore_start", "new_underscore"), # Starting with underscore @@ -1427,7 +1427,7 @@ def test_alter_column_field_name_valid( "valid_old_name,invalid_new_name", [ ("temp_field", ""), # Empty new name - ("temp_field", "aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa"), # Too long new name + ("temp_field", "a" * 65), # Exceeds the field name byte limit. ("temp_field", "field name"), # New name with space ("temp_field", "field.name"), # New name with dot ("temp_field", "field@name"), # New name with special character @@ -1499,15 +1499,14 @@ def test_drop_column_exist(self, basic_collection: Collection): assert SCHEMA_VALIDATE_ERROR_MSG in str(exc_info.value) def test_drop_column_non_exist(self, basic_collection: Collection): - with pytest.raises(Exception) as exc_info: + with pytest.raises(Exception, match=NOT_EXIST_COLUMN_TO_DROP_ERROR_MSG): basic_collection.drop_column("non_existing_column") - assert NOT_EXIST_COLUMN_TO_DROP_ERROR_MSG in str(exc_info.value) @pytest.mark.parametrize( "field_name", [ "a", # Minimum length - "a" * 32, # Maximum length (32 characters) + "a" * 64, # Maximum field name length. "valid_field_name_123", # Alphanumeric with underscore "Valid-Field-Name", # With hyphens "_underscore_start", # Starting with underscore diff --git a/python/tests/detail/test_collection_dml.py b/python/tests/detail/test_collection_dml.py index e1b0c97e2..9bd0597d2 100644 --- a/python/tests/detail/test_collection_dml.py +++ b/python/tests/detail/test_collection_dml.py @@ -26,14 +26,18 @@ "123abc", "-!@#$%+=.123abc_+", "abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ123456789012", + "()qsd123", + " ", + "/&AS12", + "订单:2026", + "a" * 1024, ] DOCID_INVALID_LIST = [ None, "", - "()qsd123", - " ", - "/&AS12", - "abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ1234567890123456789012345678901234567890123456789012345678901234567890123456789012345678901234567890121", + "a" * 1025, + "doc\0id", + "doc\nid", ] FIELD_VALUE_VALID_LIST = [ diff --git a/python/tests/test_name_validation.py b/python/tests/test_name_validation.py new file mode 100644 index 000000000..d0ad43a80 --- /dev/null +++ b/python/tests/test_name_validation.py @@ -0,0 +1,245 @@ +# Copyright 2025-present the zvec project +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import pytest +import zvec + + +@pytest.fixture +def collection(tmp_path): + schema = zvec.CollectionSchema( + "validation", fields=zvec.FieldSchema("text", zvec.DataType.STRING) + ) + coll = zvec.create_and_open(str(tmp_path / "collection"), schema) + try: + yield coll + finally: + coll.destroy() + + +@pytest.mark.parametrize("name", ["x", "ab", "集合 / 2026 🔎", "集" * 85 + "a"]) +def test_collection_names_round_trip(tmp_path, name): + path = str(tmp_path / "collection") + schema = zvec.CollectionSchema( + name, fields=zvec.FieldSchema("text", zvec.DataType.STRING) + ) + coll = zvec.create_and_open(path, schema) + try: + assert coll.schema.name == name + coll.close() + coll = zvec.open(path) + assert coll.schema.name == name + finally: + coll.close() + + +def test_document_ids_round_trip_without_normalization(collection): + ids = [ + "user:123", + "https://example.com/articles/42?a=b", + "订单-2026-🙂", + "界" * 341 + "a", # 1024 UTF-8 bytes. + "doc", + " doc", + "doc ", + " ", + "é", + "e\u0301", + ] + text = "正文\nsecond line\tvalue" + statuses = collection.insert( + [zvec.Doc(id=doc_id, fields={"text": text}) for doc_id in ids] + ) + assert all(status.ok() for status in statuses) + collection.flush() + fetched = collection.fetch(ids) + assert set(fetched) == set(ids) + for doc_id in ids: + assert fetched[doc_id].id == doc_id + assert fetched[doc_id].field("text") == text + + +@pytest.mark.parametrize("operation", ["insert", "update", "upsert"]) +@pytest.mark.parametrize( + "doc_id,reason", + [ + ("", "must not be empty"), + ("doc\0id", "null character"), + ("doc\nid", "newline"), + ("界" * 341 + "ab", "exceeds 1024 bytes (got 1025)"), + (b"\xff", "not valid UTF-8"), + ("\ud800", "not valid UTF-8"), + ], +) +def test_invalid_id_rejects_batch_before_writing(collection, operation, doc_id, reason): + if operation == "update": + assert collection.insert(zvec.Doc("valid", fields={"text": "before"})).ok() + + docs = [ + zvec.Doc("valid", fields={"text": "after"}), + zvec.Doc(doc_id, fields={"text": "invalid"}), + ] + with pytest.raises(ValueError) as exc_info: + getattr(collection, operation)(docs) + + message = str(exc_info.value) + assert message.startswith("Invalid doc:") + assert reason in message + assert "document at index 1" in message + assert "offset" not in message + fetched = collection.fetch("valid") + if operation == "update": + assert fetched["valid"].field("text") == "before" + else: + assert fetched == {} + + +@pytest.mark.parametrize("operation", ["insert", "update", "upsert"]) +def test_conversion_type_error_identifies_document_before_writing( + collection, operation +): + if operation == "update": + assert collection.insert(zvec.Doc("valid", fields={"text": "before"})).ok() + + docs = [ + zvec.Doc("valid", fields={"text": "after"}), + zvec.Doc("invalid", fields={"text": 42}), + ] + with pytest.raises(TypeError) as exc_info: + getattr(collection, operation)(docs) + + message = str(exc_info.value) + assert message.endswith(" (document at index 1)") + assert message.count("document at index") == 1 + fetched = collection.fetch("valid") + if operation == "update": + assert fetched["valid"].field("text") == "before" + else: + assert fetched == {} + + +@pytest.mark.parametrize( + "name,reason", + [ + ("", "must not be empty"), + ("collection\0name", "null character"), + ("collection\nname", "newline"), + ("集" * 86, "exceeds 256 bytes (got 258)"), + ], +) +def test_invalid_collection_names_report_the_reason(tmp_path, name, reason): + schema = zvec.CollectionSchema( + name, fields=zvec.FieldSchema("text", zvec.DataType.STRING) + ) + with pytest.raises(ValueError) as exc_info: + zvec.create_and_open(str(tmp_path / "collection"), schema) + message = str(exc_info.value) + assert message.startswith("Invalid schema:") + assert "collection name" in message + assert reason in message + assert "offset" not in message + + +def test_long_field_name_and_rejected_rename_preserve_data(tmp_path): + field_name = "f" * 64 + schema = zvec.CollectionSchema( + "fields", fields=zvec.FieldSchema(field_name, zvec.DataType.INT32) + ) + coll = zvec.create_and_open(str(tmp_path / "collection"), schema) + try: + assert coll.insert(zvec.Doc("doc", fields={field_name: 42})).ok() + with pytest.raises(ValueError) as exc_info: + coll.alter_column(field_name, new_name="f" * 65) + message = str(exc_info.value) + assert message.startswith("Invalid schema:") + assert "exceeds 64 bytes (got 65)" in message + assert "offset" not in message + assert coll.schema.field(field_name) is not None + assert coll.fetch("doc")["doc"].field(field_name) == 42 + finally: + coll.destroy() + + +@pytest.mark.parametrize("operation", ["insert", "update", "upsert"]) +def test_surrogate_id_has_a_readable_encoding_error(collection, operation): + with pytest.raises( + ValueError, + match=r"^Invalid doc: id is not valid UTF-8 \(document at index 0\)$", + ): + getattr(collection, operation)(zvec.Doc("\ud800", fields={"text": "value"})) + assert collection.stats.doc_count == 0 + + +@pytest.mark.parametrize("kind", ["collection", "field", "vector"]) +def test_surrogate_schema_name_has_a_readable_encoding_error(kind): + with pytest.raises(ValueError, match="^Invalid schema: .* is not valid UTF-8$"): + if kind == "collection": + zvec.CollectionSchema("\ud800") + elif kind == "field": + zvec.FieldSchema("\ud800", zvec.DataType.INT32) + else: + zvec.VectorSchema("\ud800", zvec.DataType.VECTOR_FP32, dimension=2) + + +@pytest.mark.parametrize("invalid", [0, False, [], {}, b""]) +@pytest.mark.parametrize("argument", ["new_name", "field_schema"]) +def test_falsey_alter_arguments_are_not_silently_ignored(tmp_path, invalid, argument): + schema = zvec.CollectionSchema( + "fields", fields=zvec.FieldSchema("value", zvec.DataType.INT32) + ) + coll = zvec.create_and_open(str(tmp_path / "collection"), schema) + try: + assert coll.insert(zvec.Doc("doc", fields={"value": 42})).ok() + kwargs = ( + { + "new_name": invalid, + "field_schema": zvec.FieldSchema("renamed", zvec.DataType.INT32), + } + if argument == "new_name" + else {"new_name": "renamed", "field_schema": invalid} + ) + with pytest.raises(TypeError, match="^Invalid schema:"): + coll.alter_column("value", **kwargs) + assert coll.schema.field("value") is not None + assert coll.schema.field("renamed") is None + assert coll.fetch("doc")["doc"].field("value") == 42 + finally: + coll.destroy() + + +@pytest.mark.parametrize("kind", ["field", "vector"]) +@pytest.mark.parametrize("name", ["bad\nname", "x" * 10000]) +def test_duplicate_name_errors_are_escaped_and_bounded(kind, name): + item = ( + zvec.FieldSchema(name, zvec.DataType.INT32) + if kind == "field" + else zvec.VectorSchema(name, zvec.DataType.VECTOR_FP32, dimension=2) + ) + with pytest.raises(ValueError) as exc_info: + zvec.CollectionSchema("duplicates", **{kind + "s": [item, item]}) + message = str(exc_info.value) + assert message.startswith("Invalid schema: duplicate") + assert "\n" not in message + assert len(message) < 256 + if "\n" in name: + assert "\\n" in message + else: + assert "..." in message + + +def test_native_schema_rejects_null_field_pointer(): + from zvec._zvec.schema import _CollectionSchema + + with pytest.raises(ValueError, match="^Invalid schema:"): + _CollectionSchema("fields", [None]) diff --git a/python/zvec/model/_validation.py b/python/zvec/model/_validation.py new file mode 100644 index 000000000..7fb461a58 --- /dev/null +++ b/python/zvec/model/_validation.py @@ -0,0 +1,29 @@ +# Copyright 2025-present the zvec project +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from __future__ import annotations + + +def explain_utf8_conversion_error(value: object, context: str) -> None: + """Explain a failed native string conversion without rescanning valid inputs.""" + if isinstance(value, str): + try: + value.encode("utf-8") + except UnicodeEncodeError: + raise ValueError(f"{context} is not valid UTF-8") from None + + +def format_name_for_error(name: str) -> str: + """Keep a user-supplied name readable, escaped, and bounded in errors.""" + preview = repr(name[:32]) + return preview + "..." if len(name) > 32 else preview diff --git a/python/zvec/model/collection.py b/python/zvec/model/collection.py index 5db24cb19..898d1f9c1 100644 --- a/python/zvec/model/collection.py +++ b/python/zvec/model/collection.py @@ -24,7 +24,8 @@ from ..executor import QueryContext, QueryExecutor from ..extension import ReRanker from ..typing import Status -from .convert import convert_to_cpp_doc, convert_to_py_doc +from ._validation import explain_utf8_conversion_error +from .convert import convert_to_cpp_docs, convert_to_py_doc from .doc import Doc, DocList, GroupResult from .param import ( AddColumnOption, @@ -304,12 +305,21 @@ def alter_column( >>> new_schema = FieldSchema(name="doc_id", dtype=DataType.INT64) >>> collection.alter_column("id", field_schema=new_schema) """ - self._obj.AlterColumn( - old_name, - new_name or "", - field_schema._get_object() if field_schema else None, - option, - ) + if new_name is not None and not isinstance(new_name, str): + raise TypeError("Invalid schema: new column name must be str") + if field_schema is not None and not isinstance(field_schema, FieldSchema): + raise TypeError("Invalid schema: field_schema must be a FieldSchema") + try: + self._obj.AlterColumn( + old_name, + "" if new_name is None else new_name, + field_schema._get_object() if field_schema is not None else None, + option, + ) + except TypeError: + explain_utf8_conversion_error(old_name, "Invalid schema: column name") + explain_utf8_conversion_error(new_name, "Invalid schema: field name") + raise self._schema = CollectionSchema._from_core(self._obj.Schema()) self._querier._schema = self._schema @@ -336,9 +346,7 @@ def insert(self, docs: Union[Doc, list[Doc]]) -> Union[Status, list[Status]]: """ is_single = isinstance(docs, Doc) doc_list = [docs] if is_single else docs - results = self._obj.Insert( - [convert_to_cpp_doc(doc, self.schema) for doc in doc_list] - ) + results = self._obj.Insert(convert_to_cpp_docs(doc_list, self.schema)) return results[0] if is_single else results @overload @@ -361,9 +369,7 @@ def upsert(self, docs: Union[Doc, list[Doc]]) -> Union[Status, list[Status]]: """ is_single = isinstance(docs, Doc) doc_list = [docs] if is_single else docs - results = self._obj.Upsert( - [convert_to_cpp_doc(doc, self.schema) for doc in doc_list] - ) + results = self._obj.Upsert(convert_to_cpp_docs(doc_list, self.schema)) return results[0] if is_single else results @overload @@ -388,9 +394,7 @@ def update(self, docs: Union[Doc, list[Doc]]) -> Union[Status, list[Status]]: """ is_single = isinstance(docs, Doc) doc_list = [docs] if is_single else docs - results = self._obj.Update( - [convert_to_cpp_doc(doc, self.schema) for doc in doc_list] - ) + results = self._obj.Update(convert_to_cpp_docs(doc_list, self.schema)) return results[0] if is_single else results @overload diff --git a/python/zvec/model/convert.py b/python/zvec/model/convert.py index 421bd1741..e37ee06ec 100644 --- a/python/zvec/model/convert.py +++ b/python/zvec/model/convert.py @@ -13,6 +13,7 @@ from zvec._zvec import _Doc +from ._validation import explain_utf8_conversion_error, format_name_for_error from .doc import Doc from .schema import CollectionSchema @@ -24,14 +25,18 @@ def convert_to_cpp_doc(doc: Doc, collection_schema: CollectionSchema) -> _Doc: _doc = _Doc() # set pk - _doc.set_pk(doc.id) + try: + _doc.set_pk(doc.id) + except TypeError: + explain_utf8_conversion_error(doc.id, "Invalid doc: id") + raise # set scalar fields for k, v in doc.fields.items(): field_schema = collection_schema.field(k) if not field_schema: raise ValueError( - f"schema validate failed: {k} not found in collection schema" + f"Invalid schema: {format_name_for_error(k)} not found in collection schema" ) _doc.set_any(k, field_schema._get_object(), v) @@ -40,12 +45,31 @@ def convert_to_cpp_doc(doc: Doc, collection_schema: CollectionSchema) -> _Doc: vector_schema = collection_schema.vector(k) if not vector_schema: raise ValueError( - f"schema validate failed: {k} not found in collection schema" + f"Invalid schema: {format_name_for_error(k)} not found in collection schema" ) _doc.set_any(k, vector_schema._get_object(), v) return _doc +def convert_to_cpp_docs( + docs: list[Doc], collection_schema: CollectionSchema +) -> list[_Doc]: + converted = [] + for index, doc in enumerate(docs): + try: + converted.append(convert_to_cpp_doc(doc, collection_schema)) + except (TypeError, ValueError) as error: + # Preserve the original exception and cause. Unicode error subclasses + # carry structured arguments that must not be replaced with a string. + if type(error) in (TypeError, ValueError): + suffix = f" (document at index {index})" + message = str(error) + if not message.endswith(suffix): + error.args = (message + suffix,) + raise + return converted + + def convert_to_py_doc(doc: _Doc, collection_schema: CollectionSchema) -> Doc: if not doc or not collection_schema: return None diff --git a/python/zvec/model/schema/collection_schema.py b/python/zvec/model/schema/collection_schema.py index 3e8971040..8e2db897e 100644 --- a/python/zvec/model/schema/collection_schema.py +++ b/python/zvec/model/schema/collection_schema.py @@ -18,6 +18,7 @@ from zvec._zvec.schema import _CollectionSchema, _FieldSchema +from .._validation import explain_utf8_conversion_error, format_name_for_error from .field_schema import FieldSchema, VectorSchema __all__ = [ @@ -64,7 +65,7 @@ def __init__( ): if name is None or not isinstance(name, str): raise ValueError( - f"schema validate failed: collection name must be str, got {type(name).__name__}" + f"Invalid schema: collection name must be str, got {type(name).__name__}" ) # handle fields @@ -75,10 +76,14 @@ def __init__( self._check_vectors(vectors, _fields_name, _fields_list) # init - self._cpp_obj = _CollectionSchema( - name=name, - fields=_fields_list, - ) + try: + self._cpp_obj = _CollectionSchema( + name=name, + fields=_fields_list, + ) + except TypeError: + explain_utf8_conversion_error(name, "Invalid schema: collection name") + raise def _check_fields( self, @@ -96,20 +101,20 @@ def _check_fields( field_items = [] else: raise TypeError( - f"schema validate failed: invalid 'fields' type, expected FieldSchema or list[FieldSchema], " + f"Invalid schema: invalid 'fields' type, expected FieldSchema or list[FieldSchema], " f"got {type(fields).__name__}" ) for idx, field in enumerate(field_items): if not isinstance(field, FieldSchema): raise TypeError( - f"schema validate failed: invalid field type in 'fields' list, expected FieldSchema, " + f"Invalid schema: invalid field type in 'fields' list, expected FieldSchema, " f"got {type(field).__name__} at index {idx}" ) if field.name in _fields_name: raise ValueError( - f"schema validate failed: duplicate field name '{field.name}': field names must be unique" + f"Invalid schema: duplicate field name {format_name_for_error(field.name)}: field names must be unique" ) _fields_name.append(field.name) _fields_list.append(field._get_object()) @@ -129,20 +134,20 @@ def _check_vectors( vectors_items = [] else: raise TypeError( - f"schema validate failed: invalid 'vectors' type, expected VectorSchema or list[VectorSchema], " + f"Invalid schema: invalid 'vectors' type, expected VectorSchema or list[VectorSchema], " f"got {type(vectors).__name__}" ) for idx, vector in enumerate(vectors_items): if not isinstance(vector, VectorSchema): raise TypeError( - f"schema validate failed: invalid vector type in 'vectors' list, expected VectorSchema, " + f"Invalid schema: invalid vector type in 'vectors' list, expected VectorSchema, " f"got {type(vector).__name__} at index {idx}" ) if vector.name in _fields_name: raise ValueError( - f"schema validate failed: duplicate vector name '{vector.name}', vector names must be unique " + f"Invalid schema: duplicate vector name {format_name_for_error(vector.name)}, vector names must be unique " f"(conflicts with existing field or vector)" ) _fields_name.append(vector.name) @@ -152,7 +157,7 @@ def _check_vectors( def _from_core(cls, core_collection_schema: _CollectionSchema): inst = cls.__new__(cls) if not core_collection_schema: - raise ValueError("schema validate failed: schema is null") + raise ValueError("Invalid schema: schema is null") inst._cpp_obj = core_collection_schema return inst diff --git a/python/zvec/model/schema/field_schema.py b/python/zvec/model/schema/field_schema.py index 0005459ba..c0b135ca9 100644 --- a/python/zvec/model/schema/field_schema.py +++ b/python/zvec/model/schema/field_schema.py @@ -28,6 +28,8 @@ ) from zvec.typing import DataType +from .._validation import explain_utf8_conversion_error, format_name_for_error + __all__ = [ "FieldSchema", "VectorSchema", @@ -104,28 +106,32 @@ def __init__( ): if name is None or not isinstance(name, str): raise ValueError( - f"schema validate failed: field name must be str, got {type(name).__name__}" + f"Invalid schema: field name must be str, got {type(name).__name__}" ) if data_type not in SUPPORT_SCALAR_DATA_TYPE: raise ValueError( - f"schema validate failed: scalar_field's data_type must be one of " + f"Invalid schema: scalar_field's data_type must be one of " f"{', '.join(str(dt) for dt in SUPPORT_SCALAR_DATA_TYPE)}, " - f"but field[{name}]'s data_type is {data_type}" + f"but field[{format_name_for_error(name)}]'s data_type is {data_type}" ) - self._cpp_obj = _FieldSchema( - name=name, - data_type=data_type, - dimension=0, - nullable=nullable, - index_param=index_param, - ) + try: + self._cpp_obj = _FieldSchema( + name=name, + data_type=data_type, + dimension=0, + nullable=nullable, + index_param=index_param, + ) + except TypeError: + explain_utf8_conversion_error(name, "Invalid schema: field name") + raise @classmethod def _from_core(cls, core_field_schema: _FieldSchema): if core_field_schema is None: - raise ValueError("schema validate failed: field schema is None") + raise ValueError("Invalid schema: field schema is None") inst = cls.__new__(cls) inst._cpp_obj = core_field_schema return inst @@ -229,29 +235,33 @@ def __init__( ): if name is None or not isinstance(name, str): raise ValueError( - f"schema validate failed: field name must be str, got {type(name).__name__}" + f"Invalid schema: field name must be str, got {type(name).__name__}" ) if not isinstance(dimension, int) or dimension < 0: - raise ValueError("schema validate failed: vector's dimension must be >= 0") + raise ValueError("Invalid schema: vector's dimension must be >= 0") if data_type not in SUPPORT_VECTOR_DATA_TYPE: raise ValueError( - f"schema validate failed: vector's data_type must be one of " + f"Invalid schema: vector's data_type must be one of " f"{', '.join(str(dt) for dt in SUPPORT_VECTOR_DATA_TYPE)}, " - f"but field[{name}]'s data_type is {data_type}" + f"but field[{format_name_for_error(name)}]'s data_type is {data_type}" ) if index_param is None: index_param = FlatIndexParam() - self._cpp_obj = _FieldSchema( - name=name, - data_type=data_type, - dimension=dimension, - nullable=False, - index_param=index_param, - ) + try: + self._cpp_obj = _FieldSchema( + name=name, + data_type=data_type, + dimension=dimension, + nullable=False, + index_param=index_param, + ) + except TypeError: + explain_utf8_conversion_error(name, "Invalid schema: field name") + raise @classmethod def _from_core(cls, core_field_schema: _FieldSchema): diff --git a/src/binding/c/c_api.cc b/src/binding/c/c_api.cc index 88c021cd5..b562ac7d9 100644 --- a/src/binding/c/c_api.cc +++ b/src/binding/c/c_api.cc @@ -26,6 +26,7 @@ #include #include #include +#include #include #include #include @@ -100,6 +101,9 @@ SET_LAST_ERROR(ZVEC_ERROR_RESOURCE_EXHAUSTED, \ std::string(msg) + ": " + e.what()); \ return nullptr; \ + } catch (const std::invalid_argument &e) { \ + SET_LAST_ERROR(ZVEC_ERROR_INVALID_ARGUMENT, e.what()); \ + return nullptr; \ } catch (const std::exception &e) { \ SET_LAST_ERROR(ZVEC_ERROR_INTERNAL_ERROR, \ std::string(msg) + ": " + e.what()); \ @@ -118,6 +122,9 @@ SET_LAST_ERROR(ZVEC_ERROR_RESOURCE_EXHAUSTED, \ std::string(msg) + ": " + e.what()); \ return ZVEC_ERROR_RESOURCE_EXHAUSTED; \ + } catch (const std::invalid_argument &e) { \ + SET_LAST_ERROR(ZVEC_ERROR_INVALID_ARGUMENT, e.what()); \ + return ZVEC_ERROR_INVALID_ARGUMENT; \ } catch (const std::exception &e) { \ SET_LAST_ERROR(ZVEC_ERROR_INTERNAL_ERROR, \ std::string(msg) + ": " + e.what()); \ @@ -136,6 +143,9 @@ SET_LAST_ERROR(ZVEC_ERROR_RESOURCE_EXHAUSTED, \ std::string(msg) + ": " + e.what()); \ return (error_val); \ + } catch (const std::invalid_argument &e) { \ + SET_LAST_ERROR(ZVEC_ERROR_INVALID_ARGUMENT, e.what()); \ + return (error_val); \ } catch (const std::exception &e) { \ SET_LAST_ERROR(ZVEC_ERROR_INTERNAL_ERROR, \ std::string(msg) + ": " + e.what()); \ @@ -725,7 +735,7 @@ zvec_error_code_t zvec_initialize(const zvec_config_data_t *config) { // Initialize global configuration auto status = zvec::GlobalConfig::Instance().initialize(cpp_config); if (!status.ok()) { - set_last_error(status.message()); + SET_LAST_ERROR(ZVEC_ERROR_INTERNAL_ERROR, status.message()); return ZVEC_ERROR_INTERNAL_ERROR; } @@ -827,6 +837,15 @@ static zvec_error_code_t status_to_error_code(const zvec::Status &status) { return static_cast(status.code()); } +// Record a returned core error without changing the per-document status mapping. +static zvec_error_code_t handle_status(const zvec::Status &status) { + const auto code = status_to_error_code(status); + if (code != ZVEC_OK) { + set_last_error_details(code, status.message()); + } + return code; +} + // Helper function: handle Expected results template static zvec_error_code_t handle_expected_result( @@ -837,8 +856,7 @@ static zvec_error_code_t handle_expected_result( } return ZVEC_OK; } else { - set_last_error(result.error().message()); - return status_to_error_code(result.error()); + return handle_status(result.error()); } } @@ -910,21 +928,6 @@ static zvec_error_code_t build_write_results( return ZVEC_OK; } -static std::vector collect_doc_pks(const zvec_doc_t **docs, - size_t doc_count) { - std::vector pks; - pks.reserve(doc_count); - for (size_t i = 0; i < doc_count; ++i) { - if (!docs[i]) { - pks.emplace_back(""); - continue; - } - auto *doc_ptr = reinterpret_cast(docs[i]); - pks.emplace_back(doc_ptr->pk_ref()); - } - return pks; -} - // ============================================================================= // Type conversion helpers // ============================================================================= @@ -2325,16 +2328,16 @@ zvec_error_code_t zvec_field_schema_set_index_params( zvec_error_code_t zvec_field_schema_validate(const zvec_field_schema_t *schema, zvec_string_t **error_msg) { + if (error_msg) { + *error_msg = nullptr; + } + if (!schema) { SET_LAST_ERROR(ZVEC_ERROR_INVALID_ARGUMENT, "Field schema pointer cannot be null"); return ZVEC_ERROR_INVALID_ARGUMENT; } - if (error_msg) { - *error_msg = nullptr; - } - ZVEC_TRY_RETURN_ERROR( "Failed to validate field schema", auto *cpp_schema = reinterpret_cast(schema); @@ -2342,7 +2345,7 @@ zvec_error_code_t zvec_field_schema_validate(const zvec_field_schema_t *schema, if (error_msg) { *error_msg = zvec_string_create(status.message().c_str()); } - return status_to_error_code(status); + return handle_status(status); }) return ZVEC_OK; @@ -2430,7 +2433,7 @@ zvec_error_code_t zvec_collection_schema_add_field(zvec_collection_schema_t *sch // Clone the field schema auto cloned_field = std::make_shared(*cpp_field); auto status = cpp_schema->add_field(cloned_field); - return status_to_error_code(status);) + return handle_status(status);) } zvec_error_code_t zvec_collection_schema_alter_field( @@ -2451,7 +2454,7 @@ zvec_error_code_t zvec_collection_schema_alter_field( auto cloned_field = std::make_shared(*cpp_new_field); auto status = cpp_schema->alter_field(std::string(field_name), cloned_field); - return status_to_error_code(status);) + return handle_status(status);) } zvec_error_code_t zvec_collection_schema_drop_field(zvec_collection_schema_t *schema, @@ -2466,7 +2469,7 @@ zvec_error_code_t zvec_collection_schema_drop_field(zvec_collection_schema_t *sc "Failed to drop field", auto *cpp_schema = reinterpret_cast(schema); auto status = cpp_schema->drop_field(std::string(field_name)); - return status_to_error_code(status);) + return handle_status(status);) } bool zvec_collection_schema_has_field(const zvec_collection_schema_t *schema, @@ -2744,16 +2747,16 @@ zvec_error_code_t zvec_collection_schema_set_max_doc_count_per_segment( zvec_error_code_t zvec_collection_schema_validate( const zvec_collection_schema_t *schema, zvec_string_t **error_msg) { + if (error_msg) { + *error_msg = nullptr; + } + if (!schema) { SET_LAST_ERROR(ZVEC_ERROR_INVALID_ARGUMENT, "Collection schema pointer cannot be null"); return ZVEC_ERROR_INVALID_ARGUMENT; } - if (error_msg) { - *error_msg = nullptr; - } - ZVEC_TRY_RETURN_ERROR( "Failed to validate schema", auto *cpp_schema = @@ -2762,7 +2765,7 @@ zvec_error_code_t zvec_collection_schema_validate( if (error_msg) { *error_msg = zvec_string_create(status.message().c_str()); } - return status_to_error_code(status); + return handle_status(status); } return ZVEC_OK;) } @@ -2783,7 +2786,7 @@ zvec_error_code_t zvec_collection_schema_add_index( auto cpp_index_params = convert_c_index_params_to_cpp(index_params); auto status = cpp_schema->add_index(std::string(field_name), cpp_index_params); - return status_to_error_code(status);) + return handle_status(status);) } zvec_error_code_t zvec_collection_schema_drop_index(zvec_collection_schema_t *schema, @@ -3187,8 +3190,15 @@ std::vector extract_binary_array(const void *value, return binary_array; } -static std::vector convert_zvec_docs_to_internal( +static zvec::Result> convert_zvec_docs_to_internal( const zvec_doc_t **zvec_docs, size_t doc_count) { + for (size_t i = 0; i < doc_count; ++i) { + if (!zvec_docs[i]) { + return tl::make_unexpected(zvec::Status::InvalidArgument( + "Invalid doc: document must not be null (document at index ", i, + ")")); + } + } std::vector docs; docs.reserve(doc_count); @@ -4721,8 +4731,10 @@ zvec_error_code_t zvec_doc_serialize(const zvec_doc_t *doc, uint8_t **data, zvec_error_code_t zvec_doc_deserialize(const uint8_t *data, size_t size, zvec_doc_t **doc) { + if (doc) *doc = nullptr; if (!data || !doc || size == 0) { - set_last_error("Invalid arguments"); + SET_LAST_ERROR(ZVEC_ERROR_INVALID_ARGUMENT, + "Invalid doc: data, size and document output must be provided"); return ZVEC_ERROR_INVALID_ARGUMENT; } @@ -4730,8 +4742,9 @@ zvec_error_code_t zvec_doc_deserialize(const uint8_t *data, size_t size, "Failed to deserialize document", auto deserialized_doc = zvec::Doc::deserialize(data, size); if (!deserialized_doc) { - set_last_error("Failed to deserialize document"); - return ZVEC_ERROR_INTERNAL_ERROR; + SET_LAST_ERROR(ZVEC_ERROR_INVALID_ARGUMENT, + "Invalid doc: serialized data is incomplete or invalid"); + return ZVEC_ERROR_INVALID_ARGUMENT; } // Create a new Doc by copying the deserialized content @@ -4792,10 +4805,12 @@ zvec_error_code_t zvec_doc_to_detail_string(const zvec_doc_t *doc, char **detail zvec_error_code_t zvec_collection_create_and_open( const char *path, const zvec_collection_schema_t *schema, const zvec_collection_options_t *options, zvec_collection_t **collection) { + if (collection) *collection = nullptr; ZVEC_TRY_RETURN_ERROR( "Exception in zvec_collection_create_and_open_with_schema", if (!path || !schema || !collection) { - set_last_error("Path, schema, or collection cannot be null"); + SET_LAST_ERROR(ZVEC_ERROR_INVALID_ARGUMENT, + "Path, schema, or collection cannot be null"); return ZVEC_ERROR_INVALID_ARGUMENT; } @@ -4804,7 +4819,7 @@ zvec_error_code_t zvec_collection_create_and_open( auto status = convert_zvec_collection_schema_to_internal(schema, schema_ptr); if (!status.ok()) { - set_last_error(status.message()); + SET_LAST_ERROR(ZVEC_ERROR_INVALID_ARGUMENT, status.message()); return ZVEC_ERROR_INVALID_ARGUMENT; } @@ -4879,9 +4894,7 @@ zvec_error_code_t zvec_collection_destroy(zvec_collection_t *collection) { auto &coll = *reinterpret_cast *>(collection); zvec::Status status = coll->destroy(); - if (!status.ok()) { set_last_error(status.message()); } - - return status_to_error_code(status);) + return handle_status(status);) } zvec_error_code_t zvec_collection_flush(zvec_collection_t *collection) { @@ -4896,9 +4909,7 @@ zvec_error_code_t zvec_collection_flush(zvec_collection_t *collection) { *reinterpret_cast *>(collection); zvec::Status status = coll->flush(); - if (!status.ok()) { set_last_error(status.message()); } - - return status_to_error_code(status);) + return handle_status(status);) } zvec_error_code_t zvec_collection_get_schema(const zvec_collection_t *collection, @@ -6805,7 +6816,7 @@ zvec_error_code_t zvec_collection_create_index( reinterpret_cast(index_params); auto index_params_ptr = cpp_params->clone(); auto status = (*coll_ptr)->create_index(field_name_str, index_params_ptr); - return status_to_error_code(status);) + return handle_status(status);) } /** @@ -6827,9 +6838,7 @@ zvec_error_code_t zvec_collection_drop_index(zvec_collection_t *collection, auto coll_ptr = reinterpret_cast *>(collection); zvec::Status status = (*coll_ptr)->drop_index(column_name); - if (!status.ok()) { set_last_error(status.message()); } - - return status_to_error_code(status);) + return handle_status(status);) } /** @@ -6848,9 +6857,7 @@ zvec_error_code_t zvec_collection_optimize(zvec_collection_t *collection) { auto coll_ptr = reinterpret_cast *>(collection); zvec::Status status = (*coll_ptr)->optimize(); - if (!status.ok()) { set_last_error(status.message()); } - - return status_to_error_code(status);) + return handle_status(status);) } // ============================================================================= @@ -6887,9 +6894,7 @@ zvec_error_code_t zvec_collection_add_column(zvec_collection_t *collection, std::string expr = expression ? expression : ""; zvec::Status status = (*coll_ptr)->add_column(schema, expr); - if (!status.ok()) { set_last_error(status.message()); } - - return status_to_error_code(status);) + return handle_status(status);) } /** @@ -6912,9 +6917,7 @@ zvec_error_code_t zvec_collection_drop_column(zvec_collection_t *collection, reinterpret_cast *>(collection); zvec::Status status = (*coll_ptr)->drop_column(column_name); - if (!status.ok()) { set_last_error(status.message()); } - - return status_to_error_code(status);) + return handle_status(status);) } zvec_error_code_t zvec_collection_alter_column( @@ -6944,9 +6947,7 @@ zvec_error_code_t zvec_collection_alter_column( zvec::Status status = (*coll_ptr)->alter_column(column_name, rename, schema); - if (!status.ok()) { set_last_error(status.message()); } - - return status_to_error_code(status);) + return handle_status(status);) } // ============================================================================= @@ -6957,9 +6958,11 @@ zvec_error_code_t zvec_collection_insert(zvec_collection_t *collection, const zvec_doc_t **docs, size_t doc_count, size_t *success_count, size_t *error_count) { + if (success_count) *success_count = 0; + if (error_count) *error_count = doc_count; if (!collection || !docs || doc_count == 0 || !success_count || !error_count) { - set_last_error( + SET_LAST_ERROR(ZVEC_ERROR_INVALID_ARGUMENT, "Invalid arguments: collection, docs, doc_count, success_count and " "error_count cannot be null/zero"); return ZVEC_ERROR_INVALID_ARGUMENT; @@ -6970,8 +6973,11 @@ zvec_error_code_t zvec_collection_insert(zvec_collection_t *collection, auto coll_ptr = reinterpret_cast *>(collection); - std::vector internal_docs = - convert_zvec_docs_to_internal(docs, doc_count); + auto converted_docs = convert_zvec_docs_to_internal(docs, doc_count); + if (!converted_docs.has_value()) { + return handle_status(converted_docs.error()); + } + auto &internal_docs = converted_docs.value(); auto result = (*coll_ptr)->insert(internal_docs); zvec_error_code_t error_code = handle_expected_result(result); @@ -6999,24 +7005,25 @@ zvec_error_code_t zvec_collection_insert_with_results(zvec_collection_t *collect size_t doc_count, zvec_write_result_t **results, size_t *result_count) { + if (results) *results = nullptr; + if (result_count) *result_count = 0; if (!collection || !docs || doc_count == 0 || !results || !result_count) { - set_last_error( + SET_LAST_ERROR(ZVEC_ERROR_INVALID_ARGUMENT, "Invalid arguments: collection, docs, doc_count, results and " "result_count cannot be null/zero"); return ZVEC_ERROR_INVALID_ARGUMENT; } - *results = nullptr; - *result_count = 0; - ZVEC_TRY_RETURN_ERROR( "Exception in zvec_collection_insert_with_results", auto coll_ptr = reinterpret_cast *>(collection); - std::vector internal_docs = - convert_zvec_docs_to_internal(docs, doc_count); - std::vector pks = collect_doc_pks(docs, doc_count); + auto converted_docs = convert_zvec_docs_to_internal(docs, doc_count); + if (!converted_docs.has_value()) { + return handle_status(converted_docs.error()); + } + auto &internal_docs = converted_docs.value(); auto result = (*coll_ptr)->insert(internal_docs); zvec_error_code_t error_code = handle_expected_result(result); @@ -7030,9 +7037,11 @@ zvec_error_code_t zvec_collection_update(zvec_collection_t *collection, const zvec_doc_t **docs, size_t doc_count, size_t *success_count, size_t *error_count) { + if (success_count) *success_count = 0; + if (error_count) *error_count = doc_count; if (!collection || !docs || doc_count == 0 || !success_count || !error_count) { - set_last_error( + SET_LAST_ERROR(ZVEC_ERROR_INVALID_ARGUMENT, "Invalid arguments: collection, docs, doc_count, success_count and " "error_count cannot be null/zero"); return ZVEC_ERROR_INVALID_ARGUMENT; @@ -7043,8 +7052,11 @@ zvec_error_code_t zvec_collection_update(zvec_collection_t *collection, auto coll_ptr = reinterpret_cast *>(collection); - std::vector internal_docs = - convert_zvec_docs_to_internal(docs, doc_count); + auto converted_docs = convert_zvec_docs_to_internal(docs, doc_count); + if (!converted_docs.has_value()) { + return handle_status(converted_docs.error()); + } + auto &internal_docs = converted_docs.value(); auto result = (*coll_ptr)->update(internal_docs); zvec_error_code_t error_code = handle_expected_result(result); @@ -7069,24 +7081,25 @@ zvec_error_code_t zvec_collection_update_with_results(zvec_collection_t *collect size_t doc_count, zvec_write_result_t **results, size_t *result_count) { + if (results) *results = nullptr; + if (result_count) *result_count = 0; if (!collection || !docs || doc_count == 0 || !results || !result_count) { - set_last_error( + SET_LAST_ERROR(ZVEC_ERROR_INVALID_ARGUMENT, "Invalid arguments: collection, docs, doc_count, results and " "result_count cannot be null/zero"); return ZVEC_ERROR_INVALID_ARGUMENT; } - *results = nullptr; - *result_count = 0; - ZVEC_TRY_RETURN_ERROR( "Exception in zvec_collection_update_with_results", auto coll_ptr = reinterpret_cast *>(collection); - std::vector internal_docs = - convert_zvec_docs_to_internal(docs, doc_count); - std::vector pks = collect_doc_pks(docs, doc_count); + auto converted_docs = convert_zvec_docs_to_internal(docs, doc_count); + if (!converted_docs.has_value()) { + return handle_status(converted_docs.error()); + } + auto &internal_docs = converted_docs.value(); auto result = (*coll_ptr)->update(internal_docs); zvec_error_code_t error_code = handle_expected_result(result); @@ -7100,9 +7113,11 @@ zvec_error_code_t zvec_collection_upsert(zvec_collection_t *collection, const zvec_doc_t **docs, size_t doc_count, size_t *success_count, size_t *error_count) { + if (success_count) *success_count = 0; + if (error_count) *error_count = doc_count; if (!collection || !docs || doc_count == 0 || !success_count || !error_count) { - set_last_error( + SET_LAST_ERROR(ZVEC_ERROR_INVALID_ARGUMENT, "Invalid arguments: collection, docs, doc_count, success_count and " "error_count cannot be null/zero"); return ZVEC_ERROR_INVALID_ARGUMENT; @@ -7113,8 +7128,11 @@ zvec_error_code_t zvec_collection_upsert(zvec_collection_t *collection, auto coll_ptr = reinterpret_cast *>(collection); - std::vector internal_docs = - convert_zvec_docs_to_internal(docs, doc_count); + auto converted_docs = convert_zvec_docs_to_internal(docs, doc_count); + if (!converted_docs.has_value()) { + return handle_status(converted_docs.error()); + } + auto &internal_docs = converted_docs.value(); auto result = (*coll_ptr)->upsert(internal_docs); zvec_error_code_t error_code = handle_expected_result(result); @@ -7139,24 +7157,25 @@ zvec_error_code_t zvec_collection_upsert_with_results(zvec_collection_t *collect size_t doc_count, zvec_write_result_t **results, size_t *result_count) { + if (results) *results = nullptr; + if (result_count) *result_count = 0; if (!collection || !docs || doc_count == 0 || !results || !result_count) { - set_last_error( + SET_LAST_ERROR(ZVEC_ERROR_INVALID_ARGUMENT, "Invalid arguments: collection, docs, doc_count, results and " "result_count cannot be null/zero"); return ZVEC_ERROR_INVALID_ARGUMENT; } - *results = nullptr; - *result_count = 0; - ZVEC_TRY_RETURN_ERROR( "Exception in zvec_collection_upsert_with_results", auto coll_ptr = reinterpret_cast *>(collection); - std::vector internal_docs = - convert_zvec_docs_to_internal(docs, doc_count); - std::vector pks = collect_doc_pks(docs, doc_count); + auto converted_docs = convert_zvec_docs_to_internal(docs, doc_count); + if (!converted_docs.has_value()) { + return handle_status(converted_docs.error()); + } + auto &internal_docs = converted_docs.value(); auto result = (*coll_ptr)->upsert(internal_docs); zvec_error_code_t error_code = handle_expected_result(result); @@ -7259,8 +7278,7 @@ zvec_error_code_t zvec_collection_delete_by_filter(zvec_collection_t *collection reinterpret_cast *>(collection); auto status = (*coll_ptr)->delete_by_filter(filter); if (!status.ok()) { - set_last_error(status.message()); - return status_to_error_code(status); + return handle_status(status); } return ZVEC_OK;) } diff --git a/src/binding/python/model/python_doc.cc b/src/binding/python/model/python_doc.cc index ae03adaaf..48e5faadf 100644 --- a/src/binding/python/model/python_doc.cc +++ b/src/binding/python/model/python_doc.cc @@ -52,7 +52,7 @@ void ZVecPyDoc::bind_doc(py::module_ &m) { doc.def(py::init([]() { return std::make_shared(); })) .def("set_pk", &Doc::set_pk) - .def("pk", &Doc::pk) + .def("pk", &Doc::pk_ref) .def("set_score", &Doc::set_score) .def("score", &Doc::score) .def("has_field", &Doc::has) @@ -325,7 +325,7 @@ py::tuple ZVecPyDoc::doc_to_tuple_with_fields( const FieldSchemaPtrList &vector_fields) { py::tuple result(4); // 1. set doc id and score - result[0] = py::str(self.pk()); + result[0] = py::str(self.pk_ref()); result[1] = py::float_(self.score()); if (self.is_empty()) { @@ -372,4 +372,4 @@ py::tuple ZVecPyDoc::doc_to_tuple_with_fields( } return result; } -} // namespace zvec \ No newline at end of file +} // namespace zvec diff --git a/src/db/collection.cc b/src/db/collection.cc index c6782b11d..fab778299 100644 --- a/src/db/collection.cc +++ b/src/db/collection.cc @@ -46,6 +46,7 @@ #include "db/index/common/delete_store.h" #include "db/index/common/id_map.h" #include "db/index/common/index_filter.h" +#include "db/index/common/name_validation.h" #include "db/index/common/type_helper.h" #include "db/index/common/version_manager.h" #include "db/index/segment/segment.h" @@ -1192,9 +1193,9 @@ Status CollectionImpl::validate(const std::string &column, if (field->data_type() < DataType::INT32 || field->data_type() > DataType::DOUBLE) { return Status::InvalidArgument( - "Only support basic numeric data type [int32, int64, uint32, uint64, " - "float, double]: ", - field->to_string()); + "Invalid schema: this operation requires a numeric field; field[", + FormatNameForError(field->name()), "] has type ", + DataTypeCodeBook::AsString(field->data_type())); } return Status::OK(); }; @@ -1202,15 +1203,14 @@ Status CollectionImpl::validate(const std::string &column, switch (op) { case ColumnOp::ADD: { if (schema == nullptr) { - return Status::InvalidArgument("Column schema is null"); + return Status::InvalidArgument( + "Invalid schema: field schema must not be null"); } - if (schema->name().empty()) { - return Status::InvalidArgument("Column name is empty"); - } if (schema_->has_field(schema->name())) { - return Status::InvalidArgument("column already exists: ", - schema->name()); + return Status::InvalidArgument("Invalid schema: field[", + FormatNameForError(schema->name()), + "] already exists"); } auto s = schema->validate(); @@ -1219,26 +1219,35 @@ Status CollectionImpl::validate(const std::string &column, s = check_data_type(schema.get()); CHECK_RETURN_STATUS(s); - if (expression.empty() && !schema->nullable()) { + if (schema_->forward_fields().size() >= kMaxScalarFieldSize) { return Status::InvalidArgument( - "Add column is not supported for non-nullable column: ", - schema->name()); + "Invalid schema: cannot add field; collection already has ", + kMaxScalarFieldSize, " scalar fields"); + } + + if (expression.empty() && !schema->nullable()) { + return Status::InvalidArgument("Invalid schema: non-nullable field[", + FormatNameForError(schema->name()), + "] requires an expression when added"); } break; } case ColumnOp::ALTER: { if (column.empty()) { - return Status::InvalidArgument("column name is empty"); + return Status::InvalidArgument( + "Invalid schema: field name must not be empty"); } if (!schema_->has_field(column)) { - return Status::InvalidArgument("column ", column, " not found"); + return Status::InvalidArgument("Invalid schema: field[", + FormatNameForError(column), + "] not found"); } if (!rename.empty() && schema) { return Status::InvalidArgument( - "cannot specify both rename and new column schema"); + "Invalid schema: cannot specify both rename and new column schema"); } auto *old_field_schema = schema_->get_field(column); @@ -1247,27 +1256,26 @@ Status CollectionImpl::validate(const std::string &column, if (!rename.empty()) { // rename case + s = ValidateFieldName(rename); + CHECK_RETURN_STATUS(s); if (schema_->has_field(rename)) { - return Status::InvalidArgument("new column name ", rename, - " already exists"); + return Status::InvalidArgument("Invalid schema: field[", + FormatNameForError(rename), + "] already exists"); } } else { // schema change case if (!schema) { - return Status::InvalidArgument("New column schema is null"); + return Status::InvalidArgument( + "Invalid schema: field schema must not be null"); } s = schema->validate(); CHECK_RETURN_STATUS(s); - if (schema->name().empty()) { - return Status::InvalidArgument("new column schema name is empty"); - } - if (!schema->nullable() && old_field_schema->nullable()) { return Status::InvalidArgument( - "new column schema is not nullable, but old column schema is " - "nullable"); + "Invalid schema: cannot make a nullable field non-nullable"); } if (*old_field_schema == *schema) { @@ -1283,7 +1291,14 @@ Status CollectionImpl::validate(const std::string &column, } case ColumnOp::DROP: { if (!schema_->has_field(column)) { - return Status::InvalidArgument("Column not exists: ", column); + return Status::InvalidArgument("Invalid schema: field[", + FormatNameForError(column), + "] not found"); + } + + if (schema_->fields().size() <= 1) { + return Status::InvalidArgument( + "Invalid schema: cannot drop the last field in a collection"); } auto *old_field_schema = schema_->get_field(column); @@ -1311,15 +1326,18 @@ Status CollectionImpl::add_column(const FieldSchema::Ptr &column_schema, CHECK_DESTROY_RETURN_STATUS(destroyed_, false); CHECK_CLOSED_RETURN_STATUS(closed_, false); - // validate - auto s = validate("", column_schema, expression, "", ColumnOp::ADD); + // Keep caller-owned mutable objects out of the published schema. Validate + // and execute the operation using the same independent snapshot. + auto field_snapshot = + column_schema ? std::make_shared(*column_schema) : nullptr; + auto s = validate("", field_snapshot, expression, "", ColumnOp::ADD); CHECK_RETURN_STATUS(s); // forbidden writing until index is ready std::lock_guard write_lock(write_mtx_); auto new_schema = std::make_shared(*schema_); - s = new_schema->add_field(column_schema); + s = new_schema->add_field(field_snapshot); CHECK_RETURN_STATUS(s); if (writing_segment_->has_record()) { @@ -1330,7 +1348,7 @@ Status CollectionImpl::add_column(const FieldSchema::Ptr &column_schema, Version new_version = version_manager_->get_current_version(); // add column on segment manager - s = segment_manager_->add_column(column_schema, expression, + s = segment_manager_->add_column(field_snapshot, expression, options.concurrency_); CHECK_RETURN_STATUS(s); @@ -1465,21 +1483,19 @@ Status CollectionImpl::alter_column(const std::string &column_name, CHECK_DESTROY_RETURN_STATUS(destroyed_, false); CHECK_CLOSED_RETURN_STATUS(closed_, false); - // validate - auto s = - validate(column_name, new_column_schema, "", rename, ColumnOp::ALTER); + auto new_field_schema = + new_column_schema ? std::make_shared(*new_column_schema) + : nullptr; + auto s = validate(column_name, new_field_schema, "", rename, ColumnOp::ALTER); CHECK_RETURN_STATUS(s); // forbidden writing until index is ready std::lock_guard write_lock(write_mtx_); - std::shared_ptr new_field_schema{nullptr}; if (!rename.empty()) { new_field_schema = std::make_shared(*schema_->get_field(column_name)); new_field_schema->set_name(rename); - } else { - new_field_schema = std::make_shared(*new_column_schema); } auto new_schema = std::make_shared(*schema_); @@ -1557,7 +1573,7 @@ Status CollectionImpl::internal_fetch_by_doc(const Doc &doc, // Called from handle_update(), i.e. under write_impl()'s write_mtx_. auto segments = get_all_segments_unsafe(); uint64_t doc_id; - bool has = id_map_->has(doc.pk(), &doc_id); + bool has = id_map_->has(doc.pk_ref(), &doc_id); if (!has) { return Status::NotFound("Document not found"); } @@ -1606,9 +1622,14 @@ Result CollectionImpl::write_impl(std::vector &docs, CHECK_DESTROY_RETURN_STATUS_EXPECTED(destroyed_, false); CHECK_CLOSED_RETURN_STATUS_EXPECTED(closed_, false); - for (auto &&doc : docs) { + for (size_t i = 0; i < docs.size(); ++i) { + auto &doc = docs[i]; auto s = doc.validate_and_sanitize(schema_, mode == WriteMode::UPDATE); - CHECK_RETURN_STATUS_EXPECTED(s); + if (!s.ok()) { + return tl::make_unexpected(Status( + s.code(), + s.message() + " (document at index " + std::to_string(i) + ")")); + } } // TODO: The granularity of the write_lock is too coarse. diff --git a/src/db/common/constants.h b/src/db/common/constants.h index 3aa0512a5..2e16cb7b4 100644 --- a/src/db/common/constants.h +++ b/src/db/common/constants.h @@ -14,7 +14,6 @@ #pragma once #include -#include #include namespace zvec { @@ -32,6 +31,13 @@ const std::string GLOBAL_DOC_ID = "_zvec_g_doc_id_"; const std::string USER_ID = "_zvec_uid_"; +// Query result columns share a namespace with user fields. Keep these names +// available to validation without pulling in the Arrow query utilities. +namespace sqlengine { +inline constexpr const char *kFieldScore = "_zvec_score"; +inline constexpr const char *kFieldGroupId = "_zvec_group_id"; +} // namespace sqlengine + const int kSparseMaxDimSize = 16384; const int64_t kMaxRecordBatchNumRows = 4096; @@ -40,12 +46,6 @@ constexpr uint32_t MAX_ARRAY_FIELD_LEN = 32; const float COMPACT_DELETE_RATIO_THRESHOLD = 0.3f; -const std::regex COLLECTION_NAME_REGEX("^[a-zA-Z0-9_-]{3,64}$"); - -const std::regex FIELD_NAME_REGEX("^[a-zA-Z0-9_-]{1,32}$"); - -const std::regex DOC_PK_REGEX("^[a-zA-Z0-9_!@#$%+=.-]{1,64}$"); - constexpr uint32_t kMaxDenseDimSize = 20000; constexpr uint32_t kMaxScalarFieldSize = 1024; diff --git a/src/db/index/common/doc.cc b/src/db/index/common/doc.cc index f76ad1d86..4935414dd 100644 --- a/src/db/index/common/doc.cc +++ b/src/db/index/common/doc.cc @@ -18,12 +18,12 @@ #include #include #include -#include #include #include #include #include #include "db/common/constants.h" +#include "db/index/common/name_validation.h" #include "db/index/common/type_helper.h" #if defined(__BYTE_ORDER__) && __BYTE_ORDER__ == __ORDER_BIG_ENDIAN__ @@ -123,28 +123,11 @@ namespace { template T byte_swap(T value) { - if constexpr (std::is_same_v) { - uint16_t val; - std::memcpy(&val, static_cast(&value), sizeof(val)); - val = ailego_bswap16(val); - float16_t result; - std::memcpy(static_cast(&result), &val, sizeof(result)); - return result; - } else if constexpr (sizeof(T) == 1) { - return value; - } else if constexpr (sizeof(T) == 2) { - return (value << 8) | ((value >> 8) & 0xFF); - } else if constexpr (sizeof(T) == 4) { - return static_cast(ailego_bswap32(static_cast(value))); - } else if constexpr (sizeof(T) == 8) { - return static_cast(ailego_bswap64(static_cast(value))); - } else { - T result = 0; - for (size_t i = 0; i < sizeof(T); ++i) { - result |= ((value >> (i * 8)) & 0xFF) << ((sizeof(T) - 1 - i) * 8); - } - return result; - } + T result; + const auto *source = reinterpret_cast(&value); + auto *destination = reinterpret_cast(&result); + std::reverse_copy(source, source + sizeof(T), destination); + return result; } template @@ -157,17 +140,179 @@ void write_value_to_buffer(std::vector &buffer, const T &value) { buffer.insert(buffer.end(), bytes, bytes + sizeof(T)); } -template -T read_value_from_buffer(const uint8_t *&data) { - T value; - std::memcpy(&value, data, sizeof(T)); - data += sizeof(T); +// Read persisted values only after checking their complete byte range. Length +// fields are checked before allocation, including each nested array/string. +class DocBufferReader { + public: + DocBufferReader(const uint8_t *data, size_t size) + : data_(data), remaining_(size) {} - if (IS_BIG_ENDIAN) { - value = byte_swap(value); + size_t remaining() const { + return remaining_; + } + + bool ReadBytes(void *destination, size_t size) { + if (size > remaining_) return false; + if (size != 0) { + std::memcpy(destination, data_, size); + data_ += size; + remaining_ -= size; + } + return true; + } + + template + bool ReadNative(T &value) { + return ReadBytes(&value, sizeof(T)); + } + + bool ReadStringBytes(std::string &value, size_t size) { + if (size > remaining_) return false; + value.assign(reinterpret_cast(data_), size); + data_ += size; + remaining_ -= size; + return true; + } + + bool ReadValue(Doc::Value &value) { + uint8_t type; + if (!ReadNative(type)) return false; + switch (type) { + case TYPE_EMPTY: + value = std::monostate{}; + return true; + case TYPE_BOOL: + return ReadAs(value); + case TYPE_INT32: + return ReadAs(value); + case TYPE_UINT32: + return ReadAs(value); + case TYPE_INT64: + return ReadAs(value); + case TYPE_UINT64: + return ReadAs(value); + case TYPE_FLOAT: + return ReadAs(value); + case TYPE_DOUBLE: + return ReadAs(value); + case TYPE_STRING: + return ReadAs(value); + case TYPE_VECTOR_BOOL: + return ReadAs>(value); + case TYPE_VECTOR_INT8: + return ReadAs>(value); + case TYPE_VECTOR_INT16: + return ReadAs>(value); + case TYPE_VECTOR_INT32: + return ReadAs>(value); + case TYPE_VECTOR_INT64: + return ReadAs>(value); + case TYPE_VECTOR_UINT32: + return ReadAs>(value); + case TYPE_VECTOR_UINT64: + return ReadAs>(value); + case TYPE_VECTOR_FLOAT16: + return ReadAs>(value); + case TYPE_VECTOR_FLOAT: + return ReadAs>(value); + case TYPE_VECTOR_DOUBLE: + return ReadAs>(value); + case TYPE_VECTOR_STRING: + return ReadAs>(value); + case TYPE_VECTOR_PAIR_INT_FLOAT: + return ReadAs, std::vector>>( + value); + case TYPE_VECTOR_PAIR_INT_FLOAT16: + return ReadAs, std::vector>>( + value); + default: + return false; + } + } + + private: + template + bool ReadLittle(T &value) { + if (!ReadNative(value)) return false; + if (IS_BIG_ENDIAN) { + auto *bytes = reinterpret_cast(&value); + std::reverse(bytes, bytes + sizeof(T)); + } + return true; + } + + template + bool ReadAs(Doc::Value &out) { + T value; + if (!Read(value)) return false; + out = std::move(value); + return true; + } + + template + bool Read(T &value) { + return ReadLittle(value); } - return value; -} + + bool Read(bool &value) { + static_assert(sizeof(bool) == sizeof(uint8_t)); + uint8_t byte; + if (!ReadNative(byte) || byte > 1) return false; + value = byte != 0; + return true; + } + + bool Read(std::string &value) { + uint32_t size; + return ReadLittle(size) && ReadStringBytes(value, size); + } + + template + bool Read(std::vector &values) { + uint32_t count; + if (!ReadLittle(count)) return false; + if constexpr (std::is_same_v) { + // Each string contains at least its four-byte length prefix. + if (count > remaining_ / sizeof(uint32_t)) return false; + values.reserve(count); + for (uint32_t i = 0; i < count; ++i) { + std::string value; + if (!Read(value)) return false; + values.push_back(std::move(value)); + } + } else if constexpr (std::is_same_v) { + if (count > remaining_ / sizeof(bool)) return false; + values.reserve(count); + for (uint32_t i = 0; i < count; ++i) { + bool value; + if (!Read(value)) return false; + values.push_back(value); + } + } else { + // Division avoids overflow before checking the allocation/copy size. + if (count > remaining_ / sizeof(T)) return false; + values.resize(count); + if (!ReadBytes(values.data(), static_cast(count) * sizeof(T))) { + return false; + } + if (IS_BIG_ENDIAN) { + for (auto &value : values) { + auto *bytes = reinterpret_cast(&value); + std::reverse(bytes, bytes + sizeof(T)); + } + } + } + return true; + } + + template + bool Read(std::pair, std::vector> &value) { + return Read(value.first) && Read(value.second); + } + + const uint8_t *data_; + size_t remaining_; +}; template std::string vec_to_string(const std::vector &v) { @@ -199,11 +344,6 @@ void Doc::write_to_buffer(std::vector &buffer, const void *src, buffer.insert(buffer.end(), bytes, bytes + size); } -void Doc::read_from_buffer(const uint8_t *&data, void *dest, size_t size) { - std::memcpy(dest, data, size); - data += size; -} - void Doc::serialize_value(std::vector &buffer, const Value &value) { std::visit( [&buffer](const auto &v) { @@ -437,238 +577,6 @@ void Doc::serialize_value(std::vector &buffer, const Value &value) { } -Doc::Value Doc::deserialize_value(const uint8_t *&data) { - uint8_t type; - read_from_buffer(data, &type, sizeof(type)); - - switch (type) { - case TYPE_EMPTY: { - return std::monostate{}; - } - case TYPE_BOOL: { - bool v; - read_from_buffer(data, &v, sizeof(v)); - return v; - } - case TYPE_INT32: { - return read_value_from_buffer(data); - } - case TYPE_INT64: { - return read_value_from_buffer(data); - } - case TYPE_UINT32: { - return read_value_from_buffer(data); - } - case TYPE_UINT64: { - return read_value_from_buffer(data); - } - case TYPE_FLOAT: { - return read_value_from_buffer(data); - } - case TYPE_DOUBLE: { - return read_value_from_buffer(data); - } - case TYPE_STRING: { - uint32_t len = read_value_from_buffer(data); - std::string v(reinterpret_cast(data), len); - data += len; - return v; - } - case TYPE_VECTOR_BOOL: { - uint32_t len = read_value_from_buffer(data); - std::vector v; - v.reserve(len); - for (uint32_t i = 0; i < len; ++i) { - bool b; - read_from_buffer(data, &b, sizeof(b)); - v.push_back(b); - } - return v; - } - case TYPE_VECTOR_INT8: { - uint32_t len = read_value_from_buffer(data); - std::vector v(len); - read_from_buffer(data, v.data(), len * sizeof(int8_t)); - return v; - } - case TYPE_VECTOR_INT16: { - uint32_t len = read_value_from_buffer(data); - std::vector v(len); - if (IS_BIG_ENDIAN) { - for (uint32_t i = 0; i < len; ++i) { - v[i] = byte_swap(read_value_from_buffer(data)); - } - } else { - read_from_buffer(data, v.data(), len * sizeof(int16_t)); - } - return v; - } - case TYPE_VECTOR_INT32: { - uint32_t len = read_value_from_buffer(data); - std::vector v(len); - if (IS_BIG_ENDIAN) { - for (uint32_t i = 0; i < len; ++i) { - v[i] = byte_swap(read_value_from_buffer(data)); - } - } else { - read_from_buffer(data, v.data(), len * sizeof(int32_t)); - } - return v; - } - case TYPE_VECTOR_INT64: { - uint32_t len = read_value_from_buffer(data); - std::vector v(len); - if (IS_BIG_ENDIAN) { - for (uint32_t i = 0; i < len; ++i) { - v[i] = byte_swap(read_value_from_buffer(data)); - } - } else { - read_from_buffer(data, v.data(), len * sizeof(int64_t)); - } - return v; - } - case TYPE_VECTOR_UINT32: { - uint32_t len = read_value_from_buffer(data); - std::vector v(len); - if (IS_BIG_ENDIAN) { - for (uint32_t i = 0; i < len; ++i) { - v[i] = byte_swap(read_value_from_buffer(data)); - } - } else { - read_from_buffer(data, v.data(), len * sizeof(uint32_t)); - } - return v; - } - case TYPE_VECTOR_UINT64: { - uint32_t len = read_value_from_buffer(data); - std::vector v(len); - if (IS_BIG_ENDIAN) { - for (uint32_t i = 0; i < len; ++i) { - v[i] = byte_swap(read_value_from_buffer(data)); - } - } else { - read_from_buffer(data, v.data(), len * sizeof(uint64_t)); - } - return v; - } - case TYPE_VECTOR_FLOAT: { - uint32_t len = read_value_from_buffer(data); - std::vector v(len); - if (IS_BIG_ENDIAN) { - for (uint32_t i = 0; i < len; ++i) { - v[i] = byte_swap(read_value_from_buffer(data)); - } - } else { - read_from_buffer(data, v.data(), len * sizeof(float)); - } - return v; - } - case TYPE_VECTOR_DOUBLE: { - uint32_t len = read_value_from_buffer(data); - std::vector v(len); - if (IS_BIG_ENDIAN) { - for (uint32_t i = 0; i < len; ++i) { - v[i] = byte_swap(read_value_from_buffer(data)); - } - } else { - read_from_buffer(data, v.data(), len * sizeof(double)); - } - return v; - } - case TYPE_VECTOR_FLOAT16: { - uint32_t len = read_value_from_buffer(data); - std::vector v(len); - if (IS_BIG_ENDIAN) { - for (uint32_t i = 0; i < len; ++i) { - v[i] = byte_swap(read_value_from_buffer(data)); - } - } else { - read_from_buffer(data, v.data(), len * sizeof(float16_t)); - } - return v; - } - case TYPE_VECTOR_STRING: { - uint32_t len = read_value_from_buffer(data); - std::vector v; - v.reserve(len); - for (uint32_t i = 0; i < len; ++i) { - uint32_t str_len = read_value_from_buffer(data); - std::string s(reinterpret_cast(data), str_len); - data += str_len; - v.push_back(s); - } - return v; - } - case TYPE_VECTOR_PAIR_INT_FLOAT: { - uint32_t len = read_value_from_buffer(data); - std::pair, std::vector> v; - v.first.reserve(len); - if (IS_BIG_ENDIAN) { - for (uint32_t i = 0; i < len; ++i) { - v.first.push_back( - byte_swap(read_value_from_buffer(data))); - } - } else { - for (uint32_t i = 0; i < len; ++i) { - uint32_t first; - read_from_buffer(data, &first, sizeof(first)); - v.first.push_back(first); - } - } - len = read_value_from_buffer(data); - v.second.reserve(len); - if (IS_BIG_ENDIAN) { - for (uint32_t i = 0; i < len; ++i) { - v.second.push_back( - byte_swap(read_value_from_buffer(data))); - } - } else { - for (uint32_t i = 0; i < len; ++i) { - float second; - read_from_buffer(data, &second, sizeof(second)); - v.second.push_back(second); - } - } - return v; - } - case TYPE_VECTOR_PAIR_INT_FLOAT16: { - uint32_t len = read_value_from_buffer(data); - std::pair, std::vector> v; - v.first.reserve(len); - if (IS_BIG_ENDIAN) { - for (uint32_t i = 0; i < len; ++i) { - v.first.push_back( - byte_swap(read_value_from_buffer(data))); - } - } else { - for (uint32_t i = 0; i < len; ++i) { - uint32_t first; - read_from_buffer(data, &first, sizeof(first)); - v.first.push_back(first); - } - } - len = read_value_from_buffer(data); - v.second.reserve(len); - if (IS_BIG_ENDIAN) { - for (uint32_t i = 0; i < len; ++i) { - v.second.push_back( - byte_swap(read_value_from_buffer(data))); - } - } else { - for (uint32_t i = 0; i < len; ++i) { - float16_t second; - read_from_buffer(data, &second, sizeof(second)); - v.second.push_back(second); - } - } - return v; - } - - default: - throw std::runtime_error("Unknown value type: " + std::to_string(type)); - } -} - std::vector Doc::serialize() const { std::vector buffer; uint32_t pk_len = static_cast(pk_.size()); @@ -693,37 +601,40 @@ std::vector Doc::serialize() const { return buffer; } -Doc::Ptr Doc::deserialize(const uint8_t *data, size_t /*size*/) { - const uint8_t *ptr = data; - Doc::Ptr doc = std::make_shared(); - - uint32_t pk_len = read_value_from_buffer(ptr); - std::string pk(reinterpret_cast(ptr), pk_len); - ptr += pk_len; - doc->set_pk(pk); - - float score = read_value_from_buffer(ptr); - doc->set_score(score); - - uint64_t doc_id = read_value_from_buffer(ptr); - doc->set_doc_id(doc_id); - - Operator op; - read_from_buffer(ptr, &op, sizeof(op)); - doc->set_operator(op); - - uint32_t field_count = read_value_from_buffer(ptr); - +Doc::Ptr Doc::deserialize(const uint8_t *data, size_t size) { + if (!data) return nullptr; + DocBufferReader reader(data, size); + auto doc = std::make_shared(); + uint32_t pk_length; + uint32_t operation; + uint32_t field_count; + // The document header and field-name lengths retain their existing native + // representation; value payloads use the existing little-endian encoding. + if (!reader.ReadNative(pk_length) || + !reader.ReadStringBytes(doc->pk_, pk_length) || + !reader.ReadNative(doc->score_) || !reader.ReadNative(doc->doc_id_) || + !reader.ReadNative(operation) || + operation > static_cast(Operator::DELETE) || + !reader.ReadNative(field_count)) { + return nullptr; + } + doc->op_ = static_cast(operation); + // Even an empty name and a null value require a length prefix and type byte. + if (field_count > reader.remaining() / (sizeof(uint32_t) + sizeof(uint8_t))) { + return nullptr; + } for (uint32_t i = 0; i < field_count; ++i) { - uint32_t name_len = read_value_from_buffer(ptr); - std::string field_name(reinterpret_cast(ptr), name_len); - ptr += name_len; - - Doc::Value value = deserialize_value(ptr); - doc->fields_[field_name] = value; + uint32_t name_length; + std::string name; + Value value; + if (!reader.ReadNative(name_length) || + !reader.ReadStringBytes(name, name_length) || + !reader.ReadValue(value) || + !doc->fields_.emplace(std::move(name), std::move(value)).second) { + return nullptr; + } } - - return doc; + return reader.remaining() == 0 ? doc : nullptr; } Status Doc::validate_and_sanitize(const CollectionSchema::Ptr &schema, @@ -732,20 +643,17 @@ Status Doc::validate_and_sanitize(const CollectionSchema::Ptr &schema, return Status::InternalError("schema is null during doc validation"); } - if (pk_.empty()) { - return Status::InvalidArgument("Invalid doc: id (primary key) is not set"); - } - - if (!std::regex_match(pk_, DOC_PK_REGEX)) { - return Status::InvalidArgument("Invalid doc: doc[", pk_, - "] contains invalid characters"); + auto id_status = ValidateDocumentId(pk_); + if (!id_status.ok()) { + return id_status; } // check doc fields match schema for (auto &[name, value] : fields_) { if (!schema->has_field(name)) { return Status::InvalidArgument( - "Invalid doc[", pk_, "]: field[", name, + "Invalid doc: doc[", FormatNameForError(pk_), "]: field[", + FormatNameForError(name), "] does not exist in the collection schema"); } } @@ -758,16 +666,17 @@ Status Doc::validate_and_sanitize(const CollectionSchema::Ptr &schema, if (field_schema->nullable() || is_update) { continue; } - return Status::InvalidArgument("Invalid doc[", pk_, "]: field[", - field_name, - "] is required but not provided"); + return Status::InvalidArgument( + "Invalid doc: doc[", FormatNameForError(pk_), "]: field[", + FormatNameForError(field_name), "] is required but not provided"); } else { if (std::holds_alternative(field_pair->second)) { if (field_schema->nullable()) { continue; } - return Status::InvalidArgument("Invalid doc[", pk_, "]: field[", - field_name, + return Status::InvalidArgument("Invalid doc: doc[", + FormatNameForError(pk_), "]: field[", + FormatNameForError(field_name), "] is required but its value is null"); } } @@ -898,12 +807,14 @@ Status Doc::validate_and_sanitize(const CollectionSchema::Ptr &schema, field_value); if (sparse_values.size() != sparse_indices.size()) { return Status::InvalidArgument( - "Invalid doc[", pk_, "]: sparse vector field[", field_name, + "Invalid doc: doc[", FormatNameForError(pk_), + "]: sparse vector field[", FormatNameForError(field_name), "] has mismatched indices and values sizes"); } if (sparse_indices.size() > kSparseMaxDimSize) { return Status::InvalidArgument( - "Invalid doc[", pk_, "]: sparse vector field[", field_name, + "Invalid doc: doc[", FormatNameForError(pk_), + "]: sparse vector field[", FormatNameForError(field_name), "] exceeds the maximum number of sparse indices (", kSparseMaxDimSize, ")"); } @@ -911,7 +822,8 @@ Status Doc::validate_and_sanitize(const CollectionSchema::Ptr &schema, sparse_indices.size()); if (status == SparseIndicesStatus::kHasDuplicate) { return Status::InvalidArgument( - "Invalid doc[", pk_, "]: sparse vector field[", field_name, + "Invalid doc: doc[", FormatNameForError(pk_), + "]: sparse vector field[", FormatNameForError(field_name), "] contains duplicate indices"); } if (status == SparseIndicesStatus::kNeedSort) { @@ -920,7 +832,8 @@ Status Doc::validate_and_sanitize(const CollectionSchema::Ptr &schema, reinterpret_cast(sparse_values.data()), sparse_indices.size(), sizeof(float16_t))) { return Status::InvalidArgument( - "Invalid doc[", pk_, "]: sparse vector field[", field_name, + "Invalid doc: doc[", FormatNameForError(pk_), + "]: sparse vector field[", FormatNameForError(field_name), "] contains duplicate indices"); } } @@ -936,12 +849,14 @@ Status Doc::validate_and_sanitize(const CollectionSchema::Ptr &schema, field_value); if (sparse_values.size() != sparse_indices.size()) { return Status::InvalidArgument( - "Invalid doc[", pk_, "]: sparse vector field[", field_name, + "Invalid doc: doc[", FormatNameForError(pk_), + "]: sparse vector field[", FormatNameForError(field_name), "] has mismatched indices and values sizes"); } if (sparse_indices.size() > kSparseMaxDimSize) { return Status::InvalidArgument( - "Invalid doc[", pk_, "]: sparse vector field[", field_name, + "Invalid doc: doc[", FormatNameForError(pk_), + "]: sparse vector field[", FormatNameForError(field_name), "] exceeds the maximum number of sparse indices (", kSparseMaxDimSize, ")"); } @@ -949,7 +864,8 @@ Status Doc::validate_and_sanitize(const CollectionSchema::Ptr &schema, sparse_indices.size()); if (status == SparseIndicesStatus::kHasDuplicate) { return Status::InvalidArgument( - "Invalid doc[", pk_, "]: sparse vector field[", field_name, + "Invalid doc: doc[", FormatNameForError(pk_), + "]: sparse vector field[", FormatNameForError(field_name), "] contains duplicate indices"); } if (status == SparseIndicesStatus::kNeedSort) { @@ -958,7 +874,8 @@ Status Doc::validate_and_sanitize(const CollectionSchema::Ptr &schema, reinterpret_cast(sparse_values.data()), sparse_indices.size(), sizeof(float))) { return Status::InvalidArgument( - "Invalid doc[", pk_, "]: sparse vector field[", field_name, + "Invalid doc: doc[", FormatNameForError(pk_), + "]: sparse vector field[", FormatNameForError(field_name), "] contains duplicate indices"); } } @@ -966,25 +883,25 @@ Status Doc::validate_and_sanitize(const CollectionSchema::Ptr &schema, break; } default: - return Status::InvalidArgument("Invalid doc[", pk_, "]: field[", - field_name, - "] has unsupported data type"); + return Status::InvalidArgument( + "Invalid doc: doc[", FormatNameForError(pk_), "]: field[", + FormatNameForError(field_name), "] has unsupported data type"); break; } if (!type_match) { return Status::InvalidArgument( - "Invalid doc[", pk_, "]: field[", field_name, - "] type mismatch, expected ", + "Invalid doc: doc[", FormatNameForError(pk_), "]: field[", + FormatNameForError(field_name), "] type mismatch, expected ", DataTypeCodeBook::AsString(expected_type), " but got ", get_value_type_name(field_value, field_schema->is_vector_field())); } if (field_schema->is_dense_vector()) { if (value_dimension != field_schema->dimension()) { return Status::InvalidArgument( - "Invalid doc[", pk_, "]: field[", field_name, - "] dimension mismatch, expected ", field_schema->dimension(), - " but got ", value_dimension); + "Invalid doc: doc[", FormatNameForError(pk_), "]: field[", + FormatNameForError(field_name), "] dimension mismatch, expected ", + field_schema->dimension(), " but got ", value_dimension); } } } diff --git a/src/db/index/common/manifest/manifest_codec.cc b/src/db/index/common/manifest/manifest_codec.cc index 06af8e43b..1b22ef5ba 100644 --- a/src/db/index/common/manifest/manifest_codec.cc +++ b/src/db/index/common/manifest/manifest_codec.cc @@ -803,7 +803,7 @@ void ManifestCodec::EncodeCollectionSchema(const CollectionSchema &schema, schema.max_doc_count_per_segment()); } -CollectionSchema::Ptr ManifestCodec::DecodeCollectionSchema( +Result ManifestCodec::DecodeCollectionSchema( std::string_view buf) { auto schema = std::make_shared(); // The protobuf-based converter read max_doc_count_per_segment straight from @@ -816,9 +816,14 @@ CollectionSchema::Ptr ManifestCodec::DecodeCollectionSchema( case f_collection::kName: schema->set_name(r.string_value()); break; - case f_collection::kFields: - schema->add_field(DecodeFieldSchema(r.bytes())); + case f_collection::kFields: { + auto status = schema->add_field(DecodeFieldSchema(r.bytes())); + if (!status.ok()) { + return tl::make_unexpected(Status::InternalError( + "Malformed manifest schema: ", status.message())); + } break; + } case f_collection::kMaxDocCountPerSegment: schema->set_max_doc_count_per_segment(r.varint()); break; @@ -826,6 +831,10 @@ CollectionSchema::Ptr ManifestCodec::DecodeCollectionSchema( break; } } + if (!r.ok()) { + return tl::make_unexpected( + Status::InternalError("Malformed manifest schema")); + } return schema; } @@ -952,9 +961,14 @@ Status ManifestCodec::Decode(std::string_view buf, ManifestData *data) { case f_manifest::kVersion: data->version = r.uint32_value(); break; - case f_manifest::kSchema: - data->schema = DecodeCollectionSchema(r.bytes()); + case f_manifest::kSchema: { + auto schema = DecodeCollectionSchema(r.bytes()); + if (!schema.has_value()) { + return schema.error(); + } + data->schema = std::move(schema).value(); break; + } case f_manifest::kEnableMmap: data->enable_mmap = r.bool_value(); break; diff --git a/src/db/index/common/manifest_codec.h b/src/db/index/common/manifest_codec.h index 4dc7c3154..c6a22b86c 100644 --- a/src/db/index/common/manifest_codec.h +++ b/src/db/index/common/manifest_codec.h @@ -67,7 +67,8 @@ struct ManifestCodec { static void EncodeCollectionSchema(const CollectionSchema &schema, std::string *out); - static CollectionSchema::Ptr DecodeCollectionSchema(std::string_view buf); + static Result DecodeCollectionSchema( + std::string_view buf); static void EncodeBlockMeta(const BlockMeta &meta, std::string *out); static BlockMeta::Ptr DecodeBlockMeta(std::string_view buf); diff --git a/src/db/index/common/name_validation.cc b/src/db/index/common/name_validation.cc new file mode 100644 index 000000000..349cc4348 --- /dev/null +++ b/src/db/index/common/name_validation.cc @@ -0,0 +1,171 @@ +// Copyright 2025-present the zvec project +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include "name_validation.h" +#include +#include +#include +#include +#include "db/common/constants.h" + +namespace zvec { +namespace { + +const char *ForbiddenCodepointReason(utf8proc_int32_t codepoint) { + if (codepoint == 0) { + return "contains a null character"; + } + if (codepoint == '\n' || codepoint == '\r') { + return "contains a newline"; + } + if (codepoint == '\t') { + return "contains a tab"; + } + if (codepoint <= 0x1F || (codepoint >= 0x7F && codepoint <= 0x9F)) { + return "contains a control character"; + } + if (codepoint == 0x2028) { + return "contains a line separator"; + } + if (codepoint == 0x2029) { + return "contains a paragraph separator"; + } + return nullptr; +} + +Status ValidateUtf8Name(std::string_view value, size_t max_bytes, + const char *prefix) { + if (value.empty()) { + return Status::InvalidArgument(prefix, " must not be empty"); + } + if (value.size() > max_bytes) { + return Status::InvalidArgument(prefix, " exceeds ", max_bytes, + " bytes (got ", value.size(), ")"); + } + + const auto *data = reinterpret_cast(value.data()); + size_t position = 0; + while (position < value.size()) { + utf8proc_int32_t codepoint; + auto bytes = utf8proc_iterate( + data + position, static_cast(value.size() - position), + &codepoint); + if (bytes <= 0) { + return Status::InvalidArgument(prefix, " is not valid UTF-8"); + } + if (const char *reason = ForbiddenCodepointReason(codepoint)) { + return Status::InvalidArgument(prefix, " ", reason); + } + position += static_cast(bytes); + } + return Status::OK(); +} + +bool IsReservedFieldName(std::string_view name) { + static const std::array reserved_names{ + LOCAL_ROW_ID, GLOBAL_DOC_ID, USER_ID, sqlengine::kFieldScore, + sqlengine::kFieldGroupId}; + return std::find(reserved_names.begin(), reserved_names.end(), name) != + reserved_names.end(); +} + +} // namespace + +std::string FormatNameForError(std::string_view name) { + constexpr size_t kMaxPreviewBytes = 32; + constexpr char kHexDigits[] = "0123456789ABCDEF"; + auto length = std::min(name.size(), kMaxPreviewBytes); + std::string preview; + preview.reserve(length); + for (size_t i = 0; i < length; ++i) { + auto byte = static_cast(name[i]); + switch (byte) { + case '\0': + preview += "\\0"; + break; + case '\n': + preview += "\\n"; + break; + case '\r': + preview += "\\r"; + break; + case '\t': + preview += "\\t"; + break; + case '\\': + case '[': + case ']': + preview += '\\'; + preview += static_cast(byte); + break; + default: + if (byte >= 0x20 && byte <= 0x7E) { + preview += static_cast(byte); + } else { + preview += "\\x"; + preview += kHexDigits[byte >> 4]; + preview += kHexDigits[byte & 0x0F]; + } + break; + } + } + if (length < name.size()) { + preview += "..."; + } + return preview; +} + +Status ValidateDocumentId(std::string_view id) { + return ValidateUtf8Name(id, kMaxDocumentIdBytes, "Invalid doc: id"); +} + +Status ValidateCollectionName(std::string_view name) { + return ValidateUtf8Name(name, kMaxCollectionNameBytes, + "Invalid schema: collection name"); +} + +Status ValidateFieldName(std::string_view name) { + if (name.empty()) { + return Status::InvalidArgument( + "Invalid schema: field name must not be empty"); + } + if (name.size() > kMaxFieldNameBytes) { + return Status::InvalidArgument("Invalid schema: field name exceeds ", + kMaxFieldNameBytes, " bytes (got ", + name.size(), ")"); + } + for (unsigned char byte : name) { + if ((byte >= 'A' && byte <= 'Z') || (byte >= 'a' && byte <= 'z') || + (byte >= '0' && byte <= '9') || byte == '_' || byte == '-') { + continue; + } + const char *reason = byte >= 0x80 ? "contains a non-ASCII character" + : ForbiddenCodepointReason(byte); + if (!reason) { + reason = byte == ' ' ? "contains a space" + : "contains an unsupported character"; + } + return Status::InvalidArgument( + "Invalid schema: field[", FormatNameForError(name), "] ", reason, + "; use letters (A-Z, a-z), digits, underscores (_) or hyphens (-)"); + } + if (IsReservedFieldName(name)) { + return Status::InvalidArgument("Invalid schema: field[", + FormatNameForError(name), + "] is reserved; use a different name"); + } + return Status::OK(); +} + +} // namespace zvec diff --git a/src/db/index/common/name_validation.h b/src/db/index/common/name_validation.h new file mode 100644 index 000000000..49adcf3b2 --- /dev/null +++ b/src/db/index/common/name_validation.h @@ -0,0 +1,43 @@ +// Copyright 2025-present the zvec project +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#pragma once + +#include +#include +#include +#include + +namespace zvec { + +inline constexpr size_t kMaxDocumentIdBytes = 1024; +inline constexpr size_t kMaxCollectionNameBytes = 256; +inline constexpr size_t kMaxFieldNameBytes = 64; + +// Validate new input without changing its bytes. Document IDs and collection +// names are nonempty UTF-8 strings without C0/C1 controls or line/paragraph +// separators. Other spaces, including strings consisting only of spaces, are +// allowed. These validators do not normalize, trim, or change case. +Status ValidateDocumentId(std::string_view id); +Status ValidateCollectionName(std::string_view name); + +// Field names retain the ASCII letters, digits, underscore, and hyphen set, +// excluding exact names used by storage and query execution. +Status ValidateFieldName(std::string_view name); + +// Bounded, escaped preview for errors. Never includes raw control characters +// or malformed UTF-8 bytes, even when the supplied name has not been validated. +std::string FormatNameForError(std::string_view name); + +} // namespace zvec diff --git a/src/db/index/common/schema.cc b/src/db/index/common/schema.cc index cbd8f1c67..7a2b302f9 100644 --- a/src/db/index/common/schema.cc +++ b/src/db/index/common/schema.cc @@ -12,7 +12,6 @@ // See the License for the specific language governing permissions and // limitations under the License. -#include #include #include #include @@ -26,6 +25,7 @@ #include "db/common/utils.h" #include "db/index/column/fts_column/fts_types.h" #include "db/index/column/fts_column/tokenizer/tokenizer_factory.h" +#include "db/index/common/name_validation.h" #include "db/index/common/type_helper.h" namespace zvec { @@ -62,7 +62,7 @@ static Status validate_fts_index_params(const FieldSchema &field) { auto params = std::dynamic_pointer_cast(field.index_params()); if (!params) { return Status::InvalidArgument( - "schema validate failed: FTS index requires FtsIndexParams, but field[", + "Invalid schema: FTS index requires FtsIndexParams, but field[", field.name(), "] has incompatible index params"); } @@ -74,30 +74,24 @@ static Status validate_fts_index_params(const FieldSchema &field) { auto pipeline = fts::TokenizerFactory::create(internal_params); if (!pipeline.has_value()) { return Status::InvalidArgument( - "schema validate failed: invalid FTS index params for field[", - field.name(), "]: ", pipeline.error().message()); + "Invalid schema: invalid FTS index params for field[", field.name(), + "]: ", pipeline.error().message()); } return Status::OK(); } Status FieldSchema::validate() const { + auto name_status = ValidateFieldName(name_); + CHECK_RETURN_STATUS(name_status); + if (data_type_ == DataType::UNDEFINED) { - return Status::InvalidArgument("schema validate failed: field[", name_, + return Status::InvalidArgument("Invalid schema: field[", name_, "]'s data_type is not defined"); } - if (name_.empty()) { - return Status::InvalidArgument("schema validate failed: field[", name_, - "]'s name is empty"); - } - if (!std::regex_match(name_, FIELD_NAME_REGEX)) { - return Status::InvalidArgument( - "schema validate failed: field[", name_, - "]'s name cannot pass the regex verification"); - } if (is_vector_field()) { auto is_sparse = is_sparse_vector(); if (!is_sparse && (dimension_ == 0 || dimension() > kMaxDenseDimSize)) { - return Status::InvalidArgument("schema validate failed: field[", name_, + return Status::InvalidArgument("Invalid schema: field[", name_, "]'s dimension must be in (0,20000]"); } @@ -105,7 +99,7 @@ Status FieldSchema::validate() const { if (support_dense_vector_type.find(data_type_) == support_dense_vector_type.end()) { return Status::InvalidArgument( - "schema validate failed: dense_vector's data type only " + "Invalid schema: dense_vector's data type only " "support FP32, " "but field[", name_, "]'s data type is ", DataTypeCodeBook::AsString(data_type_)); @@ -114,7 +108,7 @@ Status FieldSchema::validate() const { if (support_sparse_vector_type.find(data_type_) == support_sparse_vector_type.end()) { return Status::InvalidArgument( - "schema validate failed: sparse_vector's data type only " + "Invalid schema: sparse_vector's data type only " "support FP32, " "but field[", name_, "]'s data type is ", DataTypeCodeBook::AsString(data_type_)); @@ -129,7 +123,7 @@ Status FieldSchema::validate() const { if (support_sparse_vector_index.find(index_params_->type()) == support_sparse_vector_index.end()) { return Status::InvalidArgument( - "schema validate failed: sparse_vector's index_params only " + "Invalid schema: sparse_vector's index_params only " "support FLAT|HNSW index, " "but field[", name_, "]'s index_type is ", @@ -137,7 +131,7 @@ Status FieldSchema::validate() const { } if (vector_index_params->metric_type() != MetricType::IP) { return Status::InvalidArgument( - "schema validate failed: sparse_vector's index_params only " + "Invalid schema: sparse_vector's index_params only " "support IP metric, but " "field[", name_, "]'s metric is ", @@ -148,7 +142,7 @@ Status FieldSchema::validate() const { if (support_dense_vector_index.find(index_params_->type()) == support_dense_vector_index.end()) { return Status::InvalidArgument( - "schema validate failed: dense_vector's index_params only " + "Invalid schema: dense_vector's index_params only " "support FLAT|HNSW|HNSW_RABITQ|IVF|IVF_RABITQ|DISKANN|VAMANA " "index, but " "field[", @@ -161,20 +155,20 @@ Status FieldSchema::validate() const { index_params_->type() == IndexType::IVF_RABITQ) { if (dimension_ < kMinRabitqDimSize || dimension_ > kMaxRabitqDimSize) { return Status::InvalidArgument( - "schema validate failed: RabitQ index only support " + "Invalid schema: RabitQ index only support " "dimension in [", kMinRabitqDimSize, ", ", kMaxRabitqDimSize, "]"); } if (data_type_ != DataType::VECTOR_FP32) { return Status::InvalidArgument( - "schema validate failed: RabitQ index only support FP32 " + "Invalid schema: RabitQ index only support FP32 " "data types"); } auto metric_type = vector_index_params->metric_type(); if (metric_type != MetricType::L2 && metric_type != MetricType::IP && metric_type != MetricType::COSINE) { return Status::InvalidArgument( - "schema validate failed: RabitQ index only support " + "Invalid schema: RabitQ index only support " "L2/IP/COSINE metric"); } #if !RABITQ_SUPPORTED @@ -196,17 +190,17 @@ Status FieldSchema::validate() const { std::dynamic_pointer_cast(index_params_); if (!ivf_rabitq_params) { return Status::InvalidArgument( - "schema validate failed: IVF_RABITQ index requires " + "Invalid schema: IVF_RABITQ index requires " "IvfRabitqIndexParams"); } if (ivf_rabitq_params->nlist() <= 0) { return Status::InvalidArgument( - "schema validate failed: IVF_RABITQ nlist must be greater than " + "Invalid schema: IVF_RABITQ nlist must be greater than " "0"); } if (ivf_rabitq_params->sample_count() < 0) { return Status::InvalidArgument( - "schema validate failed: IVF_RABITQ sample_count must be " + "Invalid schema: IVF_RABITQ sample_count must be " "greater than or equal to 0"); } } @@ -214,7 +208,7 @@ Status FieldSchema::validate() const { if (index_params_->type() == IndexType::IVF && vector_index_params->quantize_type() == QuantizeType::RABITQ) { return Status::InvalidArgument( - "schema validate failed: IVF index does not support RABITQ " + "Invalid schema: IVF index does not support RABITQ " "quantization; use the dedicated IVF_RABITQ index instead"); } @@ -246,20 +240,20 @@ Status FieldSchema::validate() const { flat_data_type != DataType::VECTOR_FP16 && flat_data_type != DataType::VECTOR_UINT8) { return Status::InvalidArgument( - "schema validate failed: field[", name_, + "Invalid schema: field[", name_, "]'s flat_data_type must be VECTOR_FP32, VECTOR_FP16, " "or VECTOR_UINT8, but got ", DataTypeCodeBook::AsString(flat_data_type)); } if (is_sparse && flat_data_type != DataType::VECTOR_FP32) { return Status::InvalidArgument( - "schema validate failed: non-FP32 flat_data_type is only " + "Invalid schema: non-FP32 flat_data_type is only " "supported for dense vector fields"); } if (!is_sparse && flat_data_type == DataType::VECTOR_UINT8 && vector_index_params->metric_type() != MetricType::L2) { return Status::InvalidArgument( - "schema validate failed: field[", name_, + "Invalid schema: field[", name_, "] can only use VECTOR_UINT8 Flat reference storage with L2 " "metric"); } @@ -296,8 +290,7 @@ Status FieldSchema::validate() const { auto iter = quantize_type_map.find(data_type_); if (iter == quantize_type_map.end()) { return Status::InvalidArgument( - "schema validate failed: ", - is_sparse ? "sparse_vector" : "dense_vector", + "Invalid schema: ", is_sparse ? "sparse_vector" : "dense_vector", "'s index_params of ", DataTypeCodeBook::AsString(data_type_), " do not support quantize, but field[", name_, "]'s quantize_type is ", @@ -307,7 +300,7 @@ Status FieldSchema::validate() const { if (iter->second.find(vector_index_params->quantize_type()) == iter->second.end()) { return Status::InvalidArgument( - "schema validate failed: ", + "Invalid schema: ", is_sparse ? "sparse_vector" : "dense_vector", "'s index_params of ", DataTypeCodeBook::AsString(data_type_), " support ", QuantizeTypeCodeBook::AsString(iter->second), @@ -322,7 +315,7 @@ Status FieldSchema::validate() const { if (data_type_ != DataType::VECTOR_FP16 && data_type_ != DataType::VECTOR_FP32) { return Status::InvalidArgument( - "schema validate failed: IVF index only support FP32/FP16 data " + "Invalid schema: IVF index only support FP32/FP16 data " "types according to the IP metric"); } } @@ -330,7 +323,7 @@ Status FieldSchema::validate() const { if (data_type_ != DataType::VECTOR_FP16 && data_type_ != DataType::VECTOR_FP32) { return Status::InvalidArgument( - "schema validate failed: cosine metric only supports FP32/FP16 " + "Invalid schema: cosine metric only supports FP32/FP16 " "data types, but field[", name_, "]'s data type is ", DataTypeCodeBook::AsString(data_type_)); @@ -341,14 +334,14 @@ Status FieldSchema::validate() const { if (index_params_) { if (index_params_->is_vector_index_type()) { return Status::InvalidArgument( - "schema validate failed: scalar field[", name_, + "Invalid schema: scalar field[", name_, "] does not support vector index params, but got index_type ", IndexTypeCodeBook::AsString(index_params_->type())); } if (index_params_->type() == IndexType::FTS && data_type_ != DataType::STRING) { return Status::InvalidArgument( - "schema validate failed: FTS index only supports STRING data type, " + "Invalid schema: FTS index only supports STRING data type, " "but field[", name_, "]'s data_type is ", DataTypeCodeBook::AsString(data_type_)); } @@ -412,32 +405,39 @@ std::string FieldSchema::to_string_formatted(int indent_level) const { } Status CollectionSchema::validate() const { - if (name_.empty()) { - return Status::InvalidArgument("schema validate failed: name is empty"); - } - if (!std::regex_match(name_, COLLECTION_NAME_REGEX)) { - return Status::InvalidArgument( - "schema validate failed: collection[", name_, - "]'s name cannot pass the regex verification"); + auto name_status = ValidateCollectionName(name_); + CHECK_RETURN_STATUS(name_status); + std::unordered_set names; + for (const auto &field : fields_) { + if (!field) { + return Status::InvalidArgument( + "Invalid schema: field schema must not be null"); + } + if (!names.insert(field->name()).second) { + return Status::InvalidArgument("Invalid schema: duplicate field name [", + FormatNameForError(field->name()), + "]; field names must be unique"); + } } if (forward_fields().size() > kMaxScalarFieldSize) { return Status::InvalidArgument( - "schema validate failed: collection[", name_, + "Invalid schema: collection[", FormatNameForError(name_), "]'s field size must <= ", kMaxScalarFieldSize); } if (max_doc_count_per_segment_ < MAX_DOC_COUNT_PER_SEGMENT_MIN_THRESHOLD) { return Status::InvalidArgument( - "schema validate failed: max_doc_count_per_segment must >= ", + "Invalid schema: max_doc_count_per_segment must >= ", MAX_DOC_COUNT_PER_SEGMENT_MIN_THRESHOLD); } if (fields_.empty()) { - return Status::InvalidArgument("schema validate failed: collection[", name_, + return Status::InvalidArgument("Invalid schema: collection[", + FormatNameForError(name_), "] has no fields"); } auto v_fields = vector_fields(); if (v_fields.size() > kMaxVectorFieldSize) { return Status::InvalidArgument( - "schema validate failed: collection[", name_, + "Invalid schema: collection[", FormatNameForError(name_), "]'s vector field size must <= ", kMaxVectorFieldSize); } for (auto &field : fields_) { @@ -485,9 +485,14 @@ std::string CollectionSchema::to_string_formatted(int indent_level) const { } Status CollectionSchema::add_field(FieldSchema::Ptr column_schema) { + if (!column_schema) { + return Status::InvalidArgument( + "Invalid schema: field schema must not be null"); + } // Check if field already exists if (has_field(column_schema->name())) { - return Status::AlreadyExists("field[", column_schema->name(), + return Status::AlreadyExists("field[", + FormatNameForError(column_schema->name()), "] already exists in schema"); } @@ -507,16 +512,21 @@ Status CollectionSchema::add_field(FieldSchema::Ptr column_schema) { Status CollectionSchema::alter_field( const std::string &column_name, const FieldSchema::Ptr &new_column_options) { + if (!new_column_options) { + return Status::InvalidArgument( + "Invalid schema: field schema must not be null"); + } // Check if field exists if (!has_field(column_name)) { - return Status::NotFound("field[", column_name, "] not found in schema"); + return Status::NotFound("field[", FormatNameForError(column_name), + "] not found in schema"); } std::string new_column_name = new_column_options->name(); // If renaming to an existing field name (and it's not the same field) if (new_column_name != column_name && has_field(new_column_name)) { - return Status::AlreadyExists("field[", new_column_name, + return Status::AlreadyExists("field[", FormatNameForError(new_column_name), "] already exists in schema"); } @@ -540,7 +550,8 @@ Status CollectionSchema::alter_field( Status CollectionSchema::drop_field(const std::string &column_name) { // Check if field exists if (!has_field(column_name)) { - return Status::NotFound("field[", column_name, "] not found in schema"); + return Status::NotFound("field[", FormatNameForError(column_name), + "] not found in schema"); } // Remove from map @@ -723,7 +734,8 @@ Status CollectionSchema::add_index(const std::string &column, if (field) { field->set_index_params(index_params); } else { - return Status::NotFound("field[", column, "] not found in schema"); + return Status::NotFound("field[", FormatNameForError(column), + "] not found in schema"); } return Status::OK(); @@ -739,7 +751,8 @@ Status CollectionSchema::drop_index(const std::string &column) { field->set_index_params(nullptr); } } else { - return Status::NotFound("field[", column, "] not found in schema"); + return Status::NotFound("field[", FormatNameForError(column), + "] not found in schema"); } return Status::OK(); diff --git a/src/db/index/segment/segment.cc b/src/db/index/segment/segment.cc index 8d5fc4297..26742c8d7 100644 --- a/src/db/index/segment/segment.cc +++ b/src/db/index/segment/segment.cc @@ -317,10 +317,12 @@ class SegmentImpl : public Segment, Status insert_vector_indexer(Doc &doc); Status internal_insert(Doc &doc); Status internal_update(Doc &doc); - Status internal_upsert(Doc &doc); Status internal_delete(const Doc &doc); Status recover(); + Result> find_legacy_upsert_predecessors( + const std::unordered_set &keys, + uint64_t first_replay_id) const; Status open_wal_file(); Status append_wal(const Doc &doc); Status update_version(uint32_t delete_snapshot_path_suffix); @@ -915,7 +917,7 @@ Status SegmentImpl::internal_insert(Doc &doc) { } // write idmap - auto s = id_map_->upsert(doc.pk(), g_doc_id); + auto s = id_map_->upsert(doc.pk_ref(), g_doc_id); CHECK_RETURN_STATUS(s); // write forward @@ -950,26 +952,17 @@ Status SegmentImpl::internal_update(Doc &doc) { return internal_insert(doc); } -Status SegmentImpl::internal_upsert(Doc &doc) { - uint64_t g_doc_id; - bool exist = id_map_->has(doc.pk(), &g_doc_id); - if (exist) { - delete_store_->mark_deleted(g_doc_id); - } - return internal_insert(doc); -} - Status SegmentImpl::internal_delete(const Doc &doc) { delete_store_->mark_deleted(doc.doc_id()); - id_map_->remove(doc.pk()); + id_map_->remove(doc.pk_ref()); return Status::OK(); } Status SegmentImpl::Insert(Doc &doc) { std::lock_guard lock(seg_mtx_); - if (id_map_ && id_map_->has(doc.pk())) { - return Status::AlreadyExists("insert failed: doc_id[", doc.pk(), + if (id_map_ && id_map_->has(doc.pk_ref())) { + return Status::AlreadyExists("insert failed: doc_id[", doc.pk_ref(), "] already exists in collection"); } @@ -985,8 +978,8 @@ Status SegmentImpl::Insert(Doc &doc) { Status SegmentImpl::Update(Doc &doc) { std::lock_guard lock(seg_mtx_); uint64_t g_doc_id; - if (!id_map_->has(doc.pk(), &g_doc_id)) { - return Status::NotFound("update failed: doc_id[", doc.pk(), + if (!id_map_->has(doc.pk_ref(), &g_doc_id)) { + return Status::NotFound("update failed: doc_id[", doc.pk_ref(), "] not found in collection"); } @@ -1003,13 +996,28 @@ Status SegmentImpl::Update(Doc &doc) { Status SegmentImpl::Upsert(Doc &doc) { std::lock_guard lock(seg_mtx_); - doc.set_operator(Operator::UPSERT); - - // append WAL - auto s = append_wal(doc); - CHECK_RETURN_STATUS(s); + // Persist the predecessor explicitly. RocksDB may flush the new ID mapping + // before the deletion snapshot is committed, so recovery cannot safely + // infer the superseded document from the current ID map. + const auto original_doc_id = doc.doc_id(); + uint64_t previous_id; + const bool exists = id_map_->has(doc.pk_ref(), &previous_id); + if (exists) { + doc.set_doc_id(previous_id); + doc.set_operator(Operator::UPDATE); + } else { + doc.set_operator(Operator::INSERT); + } - return internal_upsert(doc); + auto status = append_wal(doc); + // Preserve the public operation on the caller's document; only the WAL + // uses the already-supported INSERT/UPDATE representation. + doc.set_operator(Operator::UPSERT); + if (!status.ok()) { + doc.set_doc_id(original_doc_id); + return status; + } + return exists ? internal_update(doc) : internal_insert(doc); } Status SegmentImpl::Delete(const std::string &pk) { @@ -4266,10 +4274,82 @@ Status SegmentImpl::init_memory_components() { return Status::OK(); } +Result> SegmentImpl::find_legacy_upsert_predecessors( + const std::unordered_set &keys, + uint64_t first_replay_id) const { + std::vector predecessors; + const auto version = version_manager_->get_current_version(); + auto segments = version.persisted_segment_metas(); + if (auto writing = version.writing_segment_meta()) { + segments.push_back(std::move(writing)); + } + for (const auto &segment : segments) { + for (const auto &block : segment->persisted_blocks()) { + if (block.type() != BlockType::SCALAR || + !block.contain_column(GLOBAL_DOC_ID) || + !block.contain_column(USER_ID)) { + continue; + } + const auto forward_path = FileHelper::MakeForwardBlockPath( + path_, segment->id(), block.id(), !options_.enable_mmap_); + BaseForwardStore::Ptr store; + // BufferPoolForwardStore only supports Parquet; IPC always uses the + // mapped store. Release each store and its buffers before the next block. + if (options_.enable_mmap_ || + InferFileFormat(forward_path) == FileFormat::IPC) { + store = std::make_shared(forward_path); + } else { + store = std::make_shared(forward_path); + } + auto status = store->Open(); + if (!status.ok()) { + return tl::make_unexpected(Status::InternalError( + "Failed to open committed rows for legacy WAL recovery: path[", + forward_path, "], reason[", status.message(), "]")); + } + auto reader = store->scan({GLOBAL_DOC_ID, USER_ID}); + if (!reader) { + return tl::make_unexpected(Status::InternalError( + "Failed to scan committed rows for legacy WAL recovery: ", + forward_path)); + } + while (true) { + std::shared_ptr batch; + const auto read_status = reader->ReadNext(&batch); + if (!read_status.ok()) { + return tl::make_unexpected(Status::InternalError( + "Failed to read committed rows for legacy WAL recovery: path[", + forward_path, "], reason[", read_status.message(), "]")); + } + if (!batch) break; + const auto ids = std::dynamic_pointer_cast( + batch->GetColumnByName(GLOBAL_DOC_ID)); + const auto pks = std::dynamic_pointer_cast( + batch->GetColumnByName(USER_ID)); + if (!ids || !pks || ids->length() != pks->length() || + ids->null_count() != 0 || pks->null_count() != 0) { + return tl::make_unexpected(Status::InternalError( + "Invalid committed identity columns during legacy WAL recovery: ", + forward_path)); + } + for (int64_t row = 0; row < ids->length(); ++row) { + const auto doc_id = ids->Value(row); + if (doc_id < first_replay_id && !delete_store_->is_deleted(doc_id) && + keys.find(pks->GetString(row)) != keys.end()) { + predecessors.push_back(doc_id); + } + } + } + } + } + return predecessors; +} + Status SegmentImpl::recover() { // recover mem block meta auto &mem_block = segment_meta_->writing_forward_block().value(); - doc_id_allocator_.store(mem_block.min_doc_id()); + const auto first_replay_id = mem_block.min_doc_id(); + doc_id_allocator_.store(first_replay_id); std::string wal_file_path = FileHelper::MakeWalPath(path_, segment_meta_->id(), mem_block.id_); @@ -4286,99 +4366,141 @@ Status SegmentImpl::recover() { 0) { LOG_ERROR("WAL recovery failed: unable to open WAL file [%s]", wal_file_path.c_str()); - return Status::OK(); + return Status::InternalError("Failed to open WAL for recovery: ", + wal_file_path); } std::array(Operator::DELETE) + 1> recovered_doc_count{}; uint64_t total_recovered_doc_count{0}; - - int ret = recover_wal_file->prepare_for_read(); - if (ret != 0) { - LOG_ERROR( - "WAL recovery failed: unable to prepare file for reading, path[%s], " - "segment[%d], ret[%d]", - wal_file_path.c_str(), id(), ret); - return Status::InternalError( - "Failed to prepare WAL file for reading: path[", wal_file_path, - "], segment[", id(), "], ret[", ret, "]"); - } + std::unordered_set legacy_upsert_keys; LOG_INFO("WAL recovery started: path[%s], segment[%d]", wal_file_path.c_str(), id()); std::lock_guard lock(seg_mtx_); - while (true) { - std::string buf = recover_wal_file->next(); - if (buf.empty()) { - break; - } - total_recovered_doc_count++; - auto doc = Doc::deserialize(reinterpret_cast(buf.data()), - buf.size()); - if (doc == nullptr) { + // Validate the complete stream before changing the ID map or indexes. A + // failed open can close and flush those stores, so discovering corruption + // after applying a prefix would otherwise persist a partial recovery. + // Read one WAL record at a time. Legacy recovery additionally holds unique + // UPSERT keys, matched predecessor IDs, and the existing forward-store + // buffers for one committed block. New WAL records need no committed scan. + for (int pass = 0; pass < 2; ++pass) { + const bool replay = pass == 1; + total_recovered_doc_count = 0; + int ret = recover_wal_file->prepare_for_read(); + if (ret != 0) { LOG_ERROR( - "WAL record recovery failed: path[%s], segment[%d], record[%zu], " - "reason[deserialization failed]", - wal_file_path.c_str(), id(), (size_t)total_recovered_doc_count); - continue; + "WAL recovery failed: unable to prepare file for reading, path[%s], " + "segment[%d], ret[%d]", + wal_file_path.c_str(), id(), ret); + return Status::InternalError( + "Failed to prepare WAL file for reading: path[", wal_file_path, + "], segment[", id(), "], ret[", ret, "]"); + } + + if (replay && !legacy_upsert_keys.empty()) { + // Old UPSERT records did not include their predecessor ID. Reconstruct + // it from committed rows even if an interrupted replay already replaced + // its ID-map entry. Finish the entire scan before changing tombstones. + auto predecessors = + find_legacy_upsert_predecessors(legacy_upsert_keys, first_replay_id); + if (!predecessors.has_value()) return predecessors.error(); + for (const auto doc_id : predecessors.value()) { + delete_store_->mark_deleted(doc_id); + } } - Status status; - switch (doc->get_operator()) { - case Operator::INSERT: { - internal_insert(*doc); - break; - } - case Operator::UPDATE: { - internal_update(*doc); - break; + while (true) { + auto record = recover_wal_file->next(); + if (!record.has_value()) { + return Status::InternalError( + "Failed to read WAL during recovery: path[", wal_file_path, + "], segment[", id(), "], reason[", record.error().message(), "]"); } - case Operator::UPSERT: { - internal_upsert(*doc); + if (!record.value().has_value()) { break; } - case Operator::DELETE: { - internal_delete(*doc); - break; + if (options_.read_only_) { + return Status::FailedPrecondition( + "WAL recovery is required; open the collection in read-write mode " + "once to recover before opening it read-only"); } - default: + const auto &buf = record.value().value(); + total_recovered_doc_count++; + auto doc = Doc::deserialize(reinterpret_cast(buf.data()), + buf.size()); + if (doc == nullptr) { LOG_ERROR( "WAL record recovery failed: path[%s], segment[%d], record[%zu], " - "operator[%d], reason[unknown operator]", - wal_file_path.c_str(), id(), (size_t)total_recovered_doc_count, - static_cast(doc->get_operator())); - break; - } + "reason[deserialization failed]", + wal_file_path.c_str(), id(), (size_t)total_recovered_doc_count); + return Status::InternalError( + "Corrupt WAL document: path[", wal_file_path, "], segment[", id(), + "], record[", total_recovered_doc_count, "]"); + } - if (!status.ok()) { - LOG_ERROR( - "WAL record recovery failed: path[%s], segment[%d], record[%zu], " - "operator[%d], reason[%s]", - wal_file_path.c_str(), id(), (size_t)total_recovered_doc_count, - static_cast(doc->get_operator()), status.message().c_str()); - continue; - } + if (!replay) { + if (doc->get_operator() == Operator::UPSERT) { + legacy_upsert_keys.insert(doc->pk_ref()); + } + continue; + } - recovered_doc_count[static_cast(doc->get_operator())]++; - } + Status status; + switch (doc->get_operator()) { + case Operator::INSERT: { + status = internal_insert(*doc); + break; + } + case Operator::UPDATE: { + status = internal_update(*doc); + break; + } + case Operator::UPSERT: { + // A previous interrupted replay may already have persisted this + // record's ID (or a later ID for the same key) in RocksDB. Only an + // older document is superseded; marking this/later replay ID deleted + // would hide a successfully recovered document on retry. + uint64_t previous_id; + if (id_map_->has(doc->pk_ref(), &previous_id) && + previous_id < doc_id_allocator_.load()) { + delete_store_->mark_deleted(previous_id); + } + status = internal_insert(*doc); + break; + } + case Operator::DELETE: { + status = internal_delete(*doc); + break; + } + default: + LOG_ERROR( + "WAL record recovery failed: path[%s], segment[%d], record[%zu], " + "operator[%d], reason[unknown operator]", + wal_file_path.c_str(), id(), (size_t)total_recovered_doc_count, + static_cast(doc->get_operator())); + return Status::InternalError("Unknown WAL document operator: path[", + wal_file_path, "], record[", + total_recovered_doc_count, "]"); + } - const auto added_docs = recovered_doc_count[0] + // INSERT - recovered_doc_count[1] + // UPSERT - recovered_doc_count[2]; // UPDATE - mem_block.max_doc_id_ += added_docs; + if (!status.ok()) { + LOG_ERROR( + "WAL record recovery failed: path[%s], segment[%d], record[%zu], " + "operator[%d], reason[%s]", + wal_file_path.c_str(), id(), (size_t)total_recovered_doc_count, + static_cast(doc->get_operator()), status.message().c_str()); + return Status(status.code(), + ailego::StringHelper::Concat( + "Failed to apply WAL record: path[", wal_file_path, + "], record[", total_recovered_doc_count, "], reason[", + status.message(), "]")); + } - ret = recover_wal_file->close(); - if (ret != 0) { - LOG_ERROR( - "WAL recovery failed: unable to close file, path[%s], " - "segment[%d], ret[%d]", - wal_file_path.c_str(), id(), ret); - return Status::InternalError("Failed to close recovered WAL file: path[", - wal_file_path, "], segment[", id(), "], ret[", - ret, "]"); + recovered_doc_count[static_cast(doc->get_operator())]++; + } } - recover_wal_file.reset(); LOG_INFO( "WAL recovery completed: path[%s], segment[%d], total[%zu], " @@ -4394,7 +4516,10 @@ Status SegmentImpl::recover() { // optimize() flush the writing segment before sealing it; without an open // member WAL, flush() treats the recovered memory components as empty and // returns without persisting them. - return open_wal_file(); + // Retain the reader's valid-tail position so a later append can discard an + // incomplete crash record. Opening read-only does not truncate the WAL. + wal_file_ = std::move(recover_wal_file); + return Status::OK(); } Status SegmentImpl::open_wal_file() { @@ -4430,10 +4555,10 @@ Status SegmentImpl::append_wal(const Doc &doc) { auto ret = wal_file_->append(std::string(buf.begin(), buf.end())); if (ret != 0) { LOG_ERROR("WAL append failed: segment[%d], pk[%s], operator[%d], ret[%d]", - id(), doc.pk().c_str(), static_cast(doc.get_operator()), + id(), doc.pk_ref().c_str(), static_cast(doc.get_operator()), ret); return Status::InternalError("Failed to append WAL: segment[", id(), - "], pk[", doc.pk(), "], operator[", + "], pk[", doc.pk_ref(), "], operator[", static_cast(doc.get_operator()), "], ret[", ret, "]"); } diff --git a/src/db/index/storage/memory_forward_store.cc b/src/db/index/storage/memory_forward_store.cc index 61e9f9666..a532a3c9b 100644 --- a/src/db/index/storage/memory_forward_store.cc +++ b/src/db/index/storage/memory_forward_store.cc @@ -105,7 +105,7 @@ arrow::Status MemForwardStore::append_doc_to_builder( // user id(pk) auto uid_builder = dynamic_cast(rb_builder->GetField(1)); - ARROW_RETURN_NOT_OK(uid_builder->Append(doc.pk())); + ARROW_RETURN_NOT_OK(uid_builder->Append(doc.pk_ref())); // other fields for (size_t idx = 2; idx < fields.size(); ++idx) { diff --git a/src/db/index/storage/wal/local_wal_file.cc b/src/db/index/storage/wal/local_wal_file.cc index 6a50e5c34..e9d0493c5 100644 --- a/src/db/index/storage/wal/local_wal_file.cc +++ b/src/db/index/storage/wal/local_wal_file.cc @@ -13,6 +13,9 @@ // limitations under the License. #include "local_wal_file.h" +#include +#include +#include #ifndef _MSC_VER #include #endif @@ -22,49 +25,73 @@ #include "db/common/file_helper.h" #include "db/common/typedef.h" -#define MAX_RECORD_SIZE 4194304 // 4Mb - namespace zvec { int LocalWalFile::append(std::string &&data) { + if (data.empty() || data.size() > std::numeric_limits::max()) { + WLOG_ERROR("Wal record length is not representable: %zu", data.size()); + return -1; + } + WalRecord record; - record.length_ = data.size(); - record.crc_ = ailego::Crc32c::Hash( - reinterpret_cast(data.data()), record.length_, 0); - record.content_ = std::forward(data); + record.length_ = static_cast(data.size()); + record.crc_ = ailego::Crc32c::Hash(data.data(), data.size(), 0); + record.content_ = std::move(data); + std::lock_guard lock(file_mutex_); + if (!opened_ || failed_) { + return -1; + } + if (incomplete_tail_offset_) { + if (!file_.truncate(*incomplete_tail_offset_)) { + WLOG_ERROR("Wal incomplete tail truncation failed"); + failed_ = true; + return -1; + } + incomplete_tail_offset_.reset(); + } + if (!file_.seek(0, ailego::File::Origin::End)) { + return -1; + } if (write_record(record) < 0) { - WLOG_ERROR("Wal write record error. record.length_[%zu]", - (size_t)record.length_); return -1; } - // if max_docs_wal_flush_ is 0, no need flush + // Keep the flush counter and flush in the same critical section as writes. if (max_docs_wal_flush_ != 0 && docs_count_ >= max_docs_wal_flush_) { if (!file_.flush()) { WLOG_ERROR("Wal flush error. docs_count_[%zu] max_docs_wal_flush_[%zu]", (size_t)docs_count_, (size_t)max_docs_wal_flush_); + failed_ = true; + return -1; } docs_count_ = 0; } return 0; } -std::string LocalWalFile::next() { +Result> LocalWalFile::next() { + std::lock_guard lock(file_mutex_); + if (!opened_ || failed_) { + return tl::make_unexpected( + Status::InternalError("WAL is not open for reading or has failed")); + } WalRecord record; - if (read_record(record) > 0) { - uint32_t tmp_crc = ailego::Crc32c::Hash( - reinterpret_cast(record.content_.data()), record.length_, - 0); - if (tmp_crc == record.crc_) { - return std::move(record.content_); - } else { - WLOG_ERROR( - "Wal next error. record.length_[%zu] crc_[%zu] != tmp_crc[%zu]", - (size_t)record.length_, (size_t)record.crc_, (size_t)tmp_crc); - } + auto result = read_record(record); + if (!result.has_value()) { + failed_ = true; + return tl::make_unexpected(result.error()); + } + if (!result.value()) { + return std::nullopt; + } + const uint32_t crc = + ailego::Crc32c::Hash(record.content_.data(), record.content_.size(), 0); + if (crc != record.crc_) { + failed_ = true; + return tl::make_unexpected( + Status::InternalError("WAL record CRC mismatch")); } - // end of file or read error - return std::string(); + return std::optional(std::move(record.content_)); } int LocalWalFile::open(const WalOptions &wal_option) { @@ -82,7 +109,7 @@ int LocalWalFile::open(const WalOptions &wal_option) { } // write wal header - int write_size = file_.write((const void *)&header_, sizeof(header_)); + size_t write_size = file_.write((const void *)&header_, sizeof(header_)); if (write_size != sizeof(header_)) { WLOG_ERROR("Wal write header error. create_new[%d]", wal_option.create_new); @@ -102,11 +129,16 @@ int LocalWalFile::open(const WalOptions &wal_option) { } // open default for write - file_.seek(0, ailego::File::Origin::End); + if (!file_.seek(0, ailego::File::Origin::End)) { + return -1; + } } max_docs_wal_flush_ = wal_option.max_docs_wal_flush; opened_ = true; + failed_ = false; + incomplete_tail_offset_.reset(); + docs_count_ = 0; WLOG_INFO("Wal open success. create_new[%d]", wal_option.create_new); return 0; @@ -142,106 +174,98 @@ int LocalWalFile::flush() { int LocalWalFile::prepare_for_read() { CHECK_STATUS(opened_, true); - if (!file_.seek(0, ailego::File::Origin::Begin)) { + incomplete_tail_offset_.reset(); + if (failed_ || !file_.seek(0, ailego::File::Origin::Begin)) { return -1; } - int read_size = file_.read((void *)&header_, sizeof(header_)); + size_t read_size = file_.read((void *)&header_, sizeof(header_)); if (read_size != sizeof(header_)) { WLOG_ERROR("Wal read header error."); + failed_ = true; return -1; } if (header_.wal_version != 0UL) { WLOG_ERROR("Wal version not support error."); + failed_ = true; return -1; } return 0; } -//! Return 1 if success or -1 if write error +// Caller holds file_mutex_. A failed write must not strand future successful +// appends behind its incomplete record. int LocalWalFile::write_record(WalRecord &record) { - CHECK_STATUS(opened_, true); - - int write_size = 0; - int ret = -1; - - std::lock_guard lock(file_mutex_); - do { - write_size = file_.write((const void *)&record.length_, LENGTH_SIZE); - if (write_size != LENGTH_SIZE) { - WLOG_ERROR("Wal write error. record.length_ error write_size[%d]", - write_size); - break; - } - - write_size = file_.write((const void *)&record.crc_, CRC_SIZE); - if (write_size != CRC_SIZE) { - WLOG_ERROR("Wal write error. record.crc_ error write_size[%d]", - write_size); - break; - } - - write_size = - file_.write((const void *)record.content_.data(), record.length_); - if (write_size != (int)record.length_) { - WLOG_ERROR("Wal write error. record.content_ error write_size[%d]", - write_size); - break; + const auto start = file_.offset(); + if (start < static_cast(sizeof(header_))) { + failed_ = true; + return -1; + } + if (file_.write(&record.length_, LENGTH_SIZE) != LENGTH_SIZE || + file_.write(&record.crc_, CRC_SIZE) != CRC_SIZE || + file_.write(record.content_.data(), record.content_.size()) != + record.content_.size()) { + WLOG_ERROR("Wal write record failed. record.length_[%zu]", + record.content_.size()); + if (!file_.truncate(static_cast(start)) || + !file_.seek(start, ailego::File::Origin::Begin)) { + failed_ = true; } - ret = 1; // write one record success - docs_count_++; - } while (false); - - return ret; + return -1; + } + ++docs_count_; + return 1; } -//! Return 1 if success or 0 if eof or -1 if read error -int LocalWalFile::read_record(WalRecord &record) { - CHECK_STATUS(opened_, true); - - int read_size = 0; - std::string err_msg; - int ret = -1; - - do { - read_size = - file_.read(reinterpret_cast(&record.length_), LENGTH_SIZE); - if (read_size == 0) { - ret = 0; - WLOG_INFO("Wal read finished. end of file"); - break; - } - - if (read_size != LENGTH_SIZE) { - WLOG_ERROR("Wal read error. record.length_ error read_size[%d]", - read_size); - break; - } - - read_size = file_.read(reinterpret_cast(&record.crc_), CRC_SIZE); - if (read_size != CRC_SIZE) { - WLOG_ERROR("Wal read error. record.crc_ error read_size[%d]", read_size); - break; - } - - // resize may crash if record.length_ very large - if (record.length_ <= 0 || record.length_ > MAX_RECORD_SIZE) { - WLOG_ERROR("Wal read error. record.length_ value error read_size[%d]", - read_size); - break; - } - +Result LocalWalFile::read_record(WalRecord &record) { + if (incomplete_tail_offset_) { + return false; + } + // File::read reports bytes read for both EOF and I/O failures. Check the + // physical extent first: a short read within that extent is an I/O error, + // whereas a final frame that does not fit is a tolerated interrupted write. + const auto start = file_.offset(); + const size_t file_size = file_.size(); + if (!file_.is_valid() || start < static_cast(sizeof(header_)) || + file_size < sizeof(header_) || static_cast(start) > file_size) { + return tl::make_unexpected( + Status::InternalError("Failed to determine WAL read position or size")); + } + const size_t remaining = file_size - static_cast(start); + if (remaining == 0) { + return false; + } + if (remaining < LENGTH_SIZE + CRC_SIZE) { + incomplete_tail_offset_ = static_cast(start); + return false; + } + if (file_.read(&record.length_, LENGTH_SIZE) != LENGTH_SIZE || + file_.read(&record.crc_, CRC_SIZE) != CRC_SIZE) { + return tl::make_unexpected( + Status::InternalError("Failed to read WAL record header")); + } + if (record.length_ == 0) { + return tl::make_unexpected( + Status::InternalError("WAL record has zero length")); + } + if (record.length_ > remaining - LENGTH_SIZE - CRC_SIZE) { + incomplete_tail_offset_ = static_cast(start); + return false; + } + try { record.content_.resize(record.length_); - read_size = file_.read((void *)const_cast(record.content_.data()), - record.length_); - if (read_size != (int)record.length_) { - WLOG_ERROR("Wal read error. record.content_ error read_size[%d]", - read_size); - break; - } - ret = 1; // read one record success - } while (false); - - return ret; + } catch (const std::bad_alloc &) { + return tl::make_unexpected(Status(StatusCode::RESOURCE_EXHAUSTED, + "Unable to allocate WAL record buffer")); + } catch (const std::length_error &) { + return tl::make_unexpected(Status(StatusCode::RESOURCE_EXHAUSTED, + "WAL record exceeds string capacity")); + } + if (file_.read(record.content_.data(), record.content_.size()) != + record.content_.size()) { + return tl::make_unexpected( + Status::InternalError("Failed to read WAL record payload")); + } + return true; } -}; // namespace zvec \ No newline at end of file +} // namespace zvec diff --git a/src/db/index/storage/wal/local_wal_file.h b/src/db/index/storage/wal/local_wal_file.h index d4f392fa8..c2a8dcadf 100644 --- a/src/db/index/storage/wal/local_wal_file.h +++ b/src/db/index/storage/wal/local_wal_file.h @@ -14,12 +14,7 @@ #pragma once #include -#include -#include -#include #include -#include -#include #include #include "wal_file.h" @@ -30,7 +25,7 @@ namespace zvec { */ struct WalHeader { uint64_t wal_version{0U}; - uint64_t reserved_[7]; + uint64_t reserved_[7]{}; }; static_assert(sizeof(WalHeader) % 64 == 0, @@ -61,7 +56,7 @@ class LocalWalFile : public WalFile { public: int append(std::string &&data) override; int prepare_for_read() override; - std::string next() override; + Result> next() override; public: int open(const WalOptions &wal_option) override; @@ -78,7 +73,7 @@ class LocalWalFile : public WalFile { private: int write_record(WalRecord &record); - int read_record(WalRecord &record); + Result read_record(WalRecord &record); private: ailego::File file_; @@ -93,6 +88,10 @@ class LocalWalFile : public WalFile { WalHeader header_; bool opened_{false}; + bool failed_{false}; + // Preserve the complete prefix and remove a torn final record before the + // next append. Merely reading a WAL must not modify it. + std::optional incomplete_tail_offset_; }; diff --git a/src/db/index/storage/wal/wal_file.h b/src/db/index/storage/wal/wal_file.h index 6e17f7384..34c2570b3 100644 --- a/src/db/index/storage/wal/wal_file.h +++ b/src/db/index/storage/wal/wal_file.h @@ -13,8 +13,11 @@ // limitations under the License. #pragma once +#include #include +#include #include +#include namespace zvec { @@ -46,7 +49,9 @@ class WalFile { public: virtual int append(std::string &&data) = 0; virtual int prepare_for_read() = 0; - virtual std::string next() = 0; + // A successful empty optional means EOF or an incomplete final crash record. + // Read failures and complete but corrupt records return an error. + virtual Result> next() = 0; public: //! Open and initialize WalFile @@ -64,4 +69,4 @@ class WalFile { virtual bool has_record() = 0; }; -}; // namespace zvec \ No newline at end of file +}; // namespace zvec diff --git a/src/db/reranker/reranker.cc b/src/db/reranker/reranker.cc index 684130123..0eeacb12d 100644 --- a/src/db/reranker/reranker.cc +++ b/src/db/reranker/reranker.cc @@ -45,7 +45,7 @@ Result score_based_rerank(const ScoreFn &score_fn, const auto &docs = results[field_idx]; for (size_t rank = 0; rank < docs.size(); ++rank) { const auto &doc = docs[rank]; - const std::string &doc_id = doc->pk(); + const std::string &doc_id = doc->pk_ref(); auto rs = score_fn(static_cast(doc->score()), static_cast(rank), field_idx); if (!rs.has_value()) { diff --git a/src/db/sqlengine/common/util.h b/src/db/sqlengine/common/util.h index 84edde3cd..0f2012ed8 100644 --- a/src/db/sqlengine/common/util.h +++ b/src/db/sqlengine/common/util.h @@ -16,15 +16,14 @@ #include #include #include +#include "db/common/constants.h" namespace zvec::sqlengine { -static const constexpr char *kFieldScore = "_zvec_score"; static const constexpr char *kFieldVector = "_zvec_vector"; static const constexpr char *kFieldSparseIndices = "_zvec_sindices"; static const constexpr char *kFieldSparseValues = "_zvec_svalues"; static const constexpr char *kFieldIsValid = "_zvec_is_valid"; -static const constexpr char *kFieldGroupId = "_zvec_group_id"; static const inline std::string kCheckNotFiltered = "check_not_filtered"; static const inline std::string kFetchVector = "fetch_vector"; diff --git a/src/include/zvec/db/doc.h b/src/include/zvec/db/doc.h index 785d9c1d8..2366708dd 100644 --- a/src/include/zvec/db/doc.h +++ b/src/include/zvec/db/doc.h @@ -310,14 +310,9 @@ class ZVEC_API Doc { private: static void serialize_value(std::vector &buffer, const Value &value); - static Value deserialize_value(const uint8_t *&data, uint8_t type); - static Value deserialize_value(const uint8_t *&data); - static void write_to_buffer(std::vector &buffer, const void *src, size_t size); - static void read_from_buffer(const uint8_t *&data, void *dest, size_t size); - struct ValueEqual; private: diff --git a/src/include/zvec/db/schema.h b/src/include/zvec/db/schema.h index 1b2219337..856d5aa59 100644 --- a/src/include/zvec/db/schema.h +++ b/src/include/zvec/db/schema.h @@ -15,6 +15,7 @@ #include #include +#include #include #include #include @@ -407,6 +408,14 @@ class ZVEC_API CollectionSchema { private: void copy_fields(const FieldSchemaPtrList &fields) { + // Constructors cannot return a Status. Reject missing field objects here + // instead of dereferencing them or silently omitting part of the schema. + for (const auto &field : fields) { + if (!field) { + throw std::invalid_argument( + "Invalid schema: field schema must not be null"); + } + } for (auto &field : fields) { auto c = std::make_shared(*field); fields_.push_back(c); diff --git a/tests/c/c_api_test.c b/tests/c/c_api_test.c index 9b7ea432d..6b0666b35 100644 --- a/tests/c/c_api_test.c +++ b/tests/c/c_api_test.c @@ -85,6 +85,23 @@ static int current_test_passed = 1; // Track if current test function passes } \ } while (0) +static void check_last_error(zvec_error_code_t code, const char *reason) { + char *message = NULL; + TEST_ASSERT(zvec_get_last_error(&message) == ZVEC_OK); + TEST_ASSERT(message != NULL); + if (message) { + TEST_ASSERT(strstr(message, reason) != NULL); + } + zvec_error_details_t details = {0}; + TEST_ASSERT(zvec_get_last_error_details(&details) == ZVEC_OK); + TEST_ASSERT(details.code == code); + TEST_ASSERT(details.message != NULL); + if (message && details.message) { + TEST_ASSERT(strcmp(message, details.message) == 0); + } + zvec_free(message); +} + // ============================================================================= // Helper functions tests // ============================================================================= @@ -1026,6 +1043,269 @@ void test_collection_basic_operations(void) { TEST_END(); } +void test_relaxed_name_validation(void) { + TEST_START(); + + const char *temp_dir = "./zvec_test_relaxed_name_validation"; + cleanup_temp_directory(temp_dir); + const char *collection_name = "\xe9\x9b\x86\xe5\x90\x88 / 2026"; + char field_name[65]; + memset(field_name, 'f', sizeof(field_name) - 1); + field_name[sizeof(field_name) - 1] = '\0'; + zvec_collection_schema_t *schema = + zvec_collection_schema_create(collection_name); + zvec_field_schema_t *field = + zvec_field_schema_create(field_name, ZVEC_DATA_TYPE_INT32, false, 0); + TEST_ASSERT(schema != NULL); + TEST_ASSERT(field != NULL); + TEST_ASSERT(zvec_collection_schema_add_field(schema, field) == ZVEC_OK); + zvec_field_schema_destroy(field); + + zvec_collection_t *collection = NULL; + zvec_error_code_t err = + zvec_collection_create_and_open(temp_dir, schema, NULL, &collection); + TEST_ASSERT(err == ZVEC_OK); + TEST_ASSERT(collection != NULL); + if (collection) { + char long_id[1025]; + memset(long_id, 'x', sizeof(long_id) - 1); + long_id[sizeof(long_id) - 1] = '\0'; + const char *ids[] = {"\xe8\xae\xa2\xe5\x8d\x95:2026", + "https://example.com/articles/42", long_id}; + zvec_doc_t *doc = zvec_doc_create(); + int32_t value = 42; + TEST_ASSERT(zvec_doc_add_field_by_value(doc, field_name, + ZVEC_DATA_TYPE_INT32, &value, + sizeof(value)) == ZVEC_OK); + const zvec_doc_t *docs[] = {doc}; + for (size_t i = 0; i < sizeof(ids) / sizeof(ids[0]); ++i) { + zvec_doc_set_pk(doc, ids[i]); + size_t success_count = 0, error_count = 0; + err = zvec_collection_insert(collection, docs, 1, &success_count, + &error_count); + TEST_ASSERT(err == ZVEC_OK); + TEST_ASSERT(success_count == 1); + TEST_ASSERT(error_count == 0); + + zvec_doc_t **fetched = NULL; + size_t found_count = 0; + err = zvec_collection_fetch(collection, &ids[i], 1, NULL, 0, false, + &fetched, &found_count); + TEST_ASSERT(err == ZVEC_OK); + TEST_ASSERT(found_count == 1); + if (found_count == 1) { + TEST_ASSERT(strcmp(zvec_doc_get_pk_pointer(fetched[0]), ids[i]) == 0); + } + zvec_docs_free(fetched, found_count); + } + + char oversized_id[1026]; + memset(oversized_id, 'x', sizeof(oversized_id) - 1); + oversized_id[sizeof(oversized_id) - 1] = '\0'; + const char *invalid_ids[] = {"\xff", oversized_id, "doc\nid"}; + const char *reasons[] = {"not valid UTF-8", "exceeds 1024 bytes (got 1025)", + "newline"}; + for (size_t i = 0; i < sizeof(invalid_ids) / sizeof(invalid_ids[0]); ++i) { + zvec_doc_set_pk(doc, invalid_ids[i]); + size_t success_count = 0, error_count = 0; + err = zvec_collection_insert(collection, docs, 1, &success_count, + &error_count); + TEST_ASSERT(err == ZVEC_ERROR_INVALID_ARGUMENT); + TEST_ASSERT(success_count == 0); + TEST_ASSERT(error_count == 1); + char *error_msg = NULL; + zvec_get_last_error(&error_msg); + TEST_ASSERT(error_msg != NULL); + if (error_msg) { + TEST_ASSERT(strncmp(error_msg, "Invalid doc:", 12) == 0); + TEST_ASSERT(strstr(error_msg, reasons[i]) != NULL); + TEST_ASSERT(strstr(error_msg, "offset") == NULL); + zvec_free(error_msg); + } + } + zvec_doc_destroy(doc); + zvec_collection_destroy(collection); + } + zvec_collection_schema_destroy(schema); + cleanup_temp_directory(temp_dir); + + TEST_END(); +} + +void test_collection_name_encoding_validation(void) { + TEST_START(); + + char max_name[257]; + memset(max_name, 'c', sizeof(max_name) - 1); + max_name[sizeof(max_name) - 1] = '\0'; + char oversized_name[258]; + memset(oversized_name, 'c', sizeof(oversized_name) - 1); + oversized_name[sizeof(oversized_name) - 1] = '\0'; + const char *names[] = {"x", max_name, "\xff", oversized_name}; + const char *reasons[] = {NULL, NULL, "not valid UTF-8", + "exceeds 256 bytes (got 257)"}; + for (size_t i = 0; i < sizeof(names) / sizeof(names[0]); ++i) { + zvec_collection_schema_t *schema = zvec_collection_schema_create(names[i]); + zvec_field_schema_t *field = + zvec_field_schema_create("value", ZVEC_DATA_TYPE_INT32, false, 0); + TEST_ASSERT(schema != NULL); + TEST_ASSERT(field != NULL); + TEST_ASSERT(zvec_collection_schema_add_field(schema, field) == ZVEC_OK); + zvec_field_schema_destroy(field); + + zvec_string_t *error_msg = NULL; + zvec_error_code_t err = zvec_collection_schema_validate(schema, &error_msg); + if (reasons[i]) { + TEST_ASSERT(err == ZVEC_ERROR_INVALID_ARGUMENT); + TEST_ASSERT(error_msg != NULL); + if (error_msg) { + const char *message = zvec_string_c_str(error_msg); + TEST_ASSERT(strncmp(message, "Invalid schema:", 15) == 0); + TEST_ASSERT(strstr(message, "collection name") != NULL); + TEST_ASSERT(strstr(message, reasons[i]) != NULL); + TEST_ASSERT(strstr(message, "offset") == NULL); + } + } else { + TEST_ASSERT(err == ZVEC_OK); + TEST_ASSERT(error_msg == NULL); + } + zvec_free_string(error_msg); + zvec_collection_schema_destroy(schema); + } + + TEST_END(); +} + +void test_validation_last_error(void) { + TEST_START(); + + zvec_collection_schema_t *schema = zvec_collection_schema_create("\xff"); + zvec_field_schema_t *field = + zvec_field_schema_create("bad name", ZVEC_DATA_TYPE_INT32, false, 0); + TEST_ASSERT(schema != NULL); + TEST_ASSERT(field != NULL); + + zvec_clear_error(); + TEST_ASSERT(zvec_collection_schema_validate(schema, NULL) == + ZVEC_ERROR_INVALID_ARGUMENT); + check_last_error(ZVEC_ERROR_INVALID_ARGUMENT, + "Invalid schema: collection name is not valid UTF-8"); + // Replacing a previous error must update both text and code. + TEST_ASSERT(zvec_field_schema_validate(field, NULL) == + ZVEC_ERROR_INVALID_ARGUMENT); + check_last_error(ZVEC_ERROR_INVALID_ARGUMENT, + "Invalid schema: field[bad name] contains a space"); + + zvec_string_t *error = NULL; + TEST_ASSERT(zvec_collection_schema_validate(schema, &error) == + ZVEC_ERROR_INVALID_ARGUMENT); + TEST_ASSERT(error != NULL); + if (error) { + check_last_error(ZVEC_ERROR_INVALID_ARGUMENT, zvec_string_c_str(error)); + } + zvec_free_string(error); + + error = (zvec_string_t *)(uintptr_t)1; + TEST_ASSERT(zvec_collection_schema_validate(NULL, &error) == + ZVEC_ERROR_INVALID_ARGUMENT); + TEST_ASSERT(error == NULL); + check_last_error(ZVEC_ERROR_INVALID_ARGUMENT, "cannot be null"); + error = (zvec_string_t *)(uintptr_t)1; + TEST_ASSERT(zvec_field_schema_validate(NULL, &error) == + ZVEC_ERROR_INVALID_ARGUMENT); + TEST_ASSERT(error == NULL); + + zvec_collection_t *collection = (zvec_collection_t *)(uintptr_t)1; + TEST_ASSERT(zvec_collection_create_and_open("./zvec_test_invalid_utf8_name", + schema, NULL, &collection) == + ZVEC_ERROR_INVALID_ARGUMENT); + TEST_ASSERT(collection == NULL); + check_last_error(ZVEC_ERROR_INVALID_ARGUMENT, + "Invalid schema: collection name is not valid UTF-8"); + + zvec_field_schema_destroy(field); + zvec_collection_schema_destroy(schema); + TEST_END(); +} + +void test_batch_validation_errors(void) { + TEST_START(); + + typedef zvec_error_code_t(ZVEC_CALL * count_write_fn)( + zvec_collection_t *, const zvec_doc_t **, size_t, size_t *, size_t *); + typedef zvec_error_code_t(ZVEC_CALL * result_write_fn)( + zvec_collection_t *, const zvec_doc_t **, size_t, zvec_write_result_t **, + size_t *); + count_write_fn count_ops[] = {zvec_collection_insert, zvec_collection_update, + zvec_collection_upsert}; + result_write_fn result_ops[] = {zvec_collection_insert_with_results, + zvec_collection_update_with_results, + zvec_collection_upsert_with_results}; + for (size_t op = 0; op < 3; ++op) { + zvec_write_result_t *results = (zvec_write_result_t *)(uintptr_t)1; + size_t result_count = 123; + TEST_ASSERT(result_ops[op](NULL, NULL, 1, &results, &result_count) == + ZVEC_ERROR_INVALID_ARGUMENT); + TEST_ASSERT(results == NULL); + TEST_ASSERT(result_count == 0); + check_last_error(ZVEC_ERROR_INVALID_ARGUMENT, "Invalid arguments:"); + } + const char *path = "./zvec_test_batch_validation_errors"; + cleanup_temp_directory(path); + zvec_collection_schema_t *schema = zvec_collection_schema_create("batch"); + zvec_field_schema_t *field = + zvec_field_schema_create("value", ZVEC_DATA_TYPE_INT32, true, 0); + TEST_ASSERT(zvec_collection_schema_add_field(schema, field) == ZVEC_OK); + zvec_field_schema_destroy(field); + zvec_collection_t *collection = NULL; + TEST_ASSERT(zvec_collection_create_and_open(path, schema, NULL, + &collection) == ZVEC_OK); + TEST_ASSERT(collection != NULL); + if (collection) { + zvec_doc_t *valid_doc = zvec_doc_create(); + zvec_doc_t *invalid_doc = zvec_doc_create(); + zvec_doc_set_pk(valid_doc, "valid_before_error"); + zvec_doc_set_pk(invalid_doc, "\xff"); + const zvec_doc_t *invalid_inputs[] = {NULL, invalid_doc}; + const char *reasons[] = {"document must not be null", + "id is not valid UTF-8"}; + for (size_t i = 0; i < 2; ++i) { + const zvec_doc_t *docs[] = {valid_doc, invalid_inputs[i]}; + for (size_t op = 0; op < 3; ++op) { + size_t success_count = 123, error_count = 456; + TEST_ASSERT(count_ops[op](collection, docs, 2, &success_count, + &error_count) == ZVEC_ERROR_INVALID_ARGUMENT); + TEST_ASSERT(success_count == 0); + TEST_ASSERT(error_count == 2); + check_last_error(ZVEC_ERROR_INVALID_ARGUMENT, reasons[i]); + check_last_error(ZVEC_ERROR_INVALID_ARGUMENT, "document at index 1"); + + zvec_write_result_t *results = (zvec_write_result_t *)(uintptr_t)1; + size_t result_count = 123; + TEST_ASSERT( + result_ops[op](collection, docs, 2, &results, &result_count) == + ZVEC_ERROR_INVALID_ARGUMENT); + TEST_ASSERT(results == NULL); + TEST_ASSERT(result_count == 0); + check_last_error(ZVEC_ERROR_INVALID_ARGUMENT, reasons[i]); + } + } + const char *ids[] = {"valid_before_error"}; + zvec_doc_t **fetched = NULL; + size_t found_count = 0; + TEST_ASSERT(zvec_collection_fetch(collection, ids, 1, NULL, 0, false, + &fetched, &found_count) == ZVEC_OK); + TEST_ASSERT(found_count == 0); + zvec_docs_free(fetched, found_count); + zvec_doc_destroy(valid_doc); + zvec_doc_destroy(invalid_doc); + zvec_collection_destroy(collection); + } + zvec_collection_schema_destroy(schema); + cleanup_temp_directory(path); + TEST_END(); +} + void test_collection_edge_cases(void) { TEST_START(); @@ -3365,6 +3645,31 @@ void test_doc_serialization(void) { TEST_ASSERT(err == ZVEC_OK); TEST_ASSERT(deserialized_int32 == -2147483648); + const size_t truncated_sizes[] = {1, data_size / 2, data_size - 1}; + for (size_t i = 0; i < sizeof(truncated_sizes) / sizeof(truncated_sizes[0]); + ++i) { + zvec_doc_t *invalid_doc = (zvec_doc_t *)(uintptr_t)1; + TEST_ASSERT(zvec_doc_deserialize(serialized_data, truncated_sizes[i], + &invalid_doc) == + ZVEC_ERROR_INVALID_ARGUMENT); + TEST_ASSERT(invalid_doc == NULL); + check_last_error(ZVEC_ERROR_INVALID_ARGUMENT, + "Invalid doc: serialized data is incomplete or invalid"); + } + zvec_doc_t *invalid_doc = (zvec_doc_t *)(uintptr_t)1; + TEST_ASSERT(zvec_doc_deserialize(NULL, data_size, &invalid_doc) == + ZVEC_ERROR_INVALID_ARGUMENT); + TEST_ASSERT(invalid_doc == NULL); + invalid_doc = (zvec_doc_t *)(uintptr_t)1; + TEST_ASSERT(zvec_doc_deserialize(serialized_data, 0, &invalid_doc) == + ZVEC_ERROR_INVALID_ARGUMENT); + TEST_ASSERT(invalid_doc == NULL); + TEST_ASSERT(zvec_doc_deserialize(serialized_data, data_size, NULL) == + ZVEC_ERROR_INVALID_ARGUMENT); + check_last_error( + ZVEC_ERROR_INVALID_ARGUMENT, + "Invalid doc: data, size and document output must be provided"); + zvec_free_uint8_array(serialized_data); free(string_field.value.string_value.data); zvec_doc_destroy(deserialized_doc); @@ -6791,6 +7096,10 @@ int main(void) { // Collection-related tests test_collection_basic_operations(); + test_relaxed_name_validation(); + test_collection_name_encoding_validation(); + test_validation_last_error(); + test_batch_validation_errors(); test_collection_edge_cases(); test_collection_delete_by_filter(); test_collection_stats(); diff --git a/tests/db/crash_recovery/relaxed_validation_recovery_test.cc b/tests/db/crash_recovery/relaxed_validation_recovery_test.cc new file mode 100644 index 000000000..c33cf820c --- /dev/null +++ b/tests/db/crash_recovery/relaxed_validation_recovery_test.cc @@ -0,0 +1,646 @@ +// Copyright 2025-present the zvec project +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include "db/index/common/id_map.h" +#include "db/index/common/version_manager.h" +#include "db/index/storage/wal/wal_file.h" + +namespace zvec { +namespace { + +std::string ReadFileBytes(const std::string &path) { + std::ifstream file(path, std::ios::binary); + return std::string(std::istreambuf_iterator(file), {}); +} + +std::string FindWal(const std::string &path) { + for (const auto &entry : + std::filesystem::recursive_directory_iterator(path)) { + if (entry.path().extension() == ".wal") return entry.path().string(); + } + return {}; +} + +std::map ReadManifests(const std::string &path) { + std::map result; + for (const auto &entry : + std::filesystem::recursive_directory_iterator(path)) { + if (entry.path().filename().string().rfind("manifest", 0) == 0) { + result.emplace(entry.path().string(), + ReadFileBytes(entry.path().string())); + } + } + return result; +} + +// Only called in ASSERT_EXIT children. Intentionally skip collection cleanup. +void WriteStringDocsAndExit(const std::string &path, + const std::vector &ids, + const std::vector &values, + bool malformed_record = false) { + CollectionSchema schema("wal_recovery"); + if (!schema + .add_field( + std::make_shared("text", DataType::STRING, false)) + .ok()) { + std::_Exit(1); + } + auto created = Collection::CreateAndOpen(path, schema, CollectionOptions{}); + if (!created.has_value()) std::_Exit(2); + auto collection = std::move(created).value(); + std::vector docs; + for (size_t i = 0; i < ids.size(); ++i) { + Doc doc; + doc.set_pk(ids[i]); + doc.set("text", values[i]); + docs.push_back(std::move(doc)); + } + auto result = collection->insert(docs); + if (!result.has_value()) std::_Exit(3); + for (const auto &status : result.value()) { + if (!status.ok()) std::_Exit(4); + } + if (malformed_record) { + auto wal = WalFile::Create(FindWal(path)); + if (wal->open(WalOptions{}) != 0 || + wal->append("invalid encoded document") != 0) { + std::_Exit(5); + } + } + std::_Exit(0); +} + +std::string FindIdMap(const std::string &path) { + for (const auto &entry : std::filesystem::directory_iterator(path)) { + if (entry.is_directory() && + entry.path().filename().string().rfind("idmap", 0) == 0) { + return entry.path().string(); + } + } + return {}; +} + +void WriteUpsertsAndExit(const std::string &path, bool persisted_base, + bool append_suffix, bool separate_segment = false) { + CollectionSchema schema("wal_recovery"); + if (!schema + .add_field( + std::make_shared("text", DataType::STRING, false)) + .ok()) { + std::_Exit(1); + } + auto created = Collection::CreateAndOpen(path, schema, CollectionOptions{}); + if (!created.has_value()) std::_Exit(2); + auto collection = std::move(created).value(); + Doc doc; + doc.set_pk("target"); + if (persisted_base) { + doc.set("text", "original"); + std::vector docs{doc}; + auto inserted = collection->insert(docs); + if (!inserted.has_value() || !inserted.value().front().ok() || + !collection->flush().ok()) + std::_Exit(3); + if (separate_segment && !collection->optimize().ok()) std::_Exit(6); + } + for (const auto &value : {"first", "second"}) { + doc.set("text", value); + std::vector docs{doc}; + auto updated = collection->upsert(docs); + if (!updated.has_value() || !updated.value().front().ok()) std::_Exit(4); + } + if (append_suffix) { + doc.set_pk("broken"); + std::vector docs{doc}; + auto inserted = collection->insert(docs); + if (!inserted.has_value() || !inserted.value().front().ok()) std::_Exit(5); + } + std::_Exit(0); +} + +void ReadWalDocuments(const std::string &path, std::vector *docs) { + const auto wal_path = FindWal(path); + ASSERT_FALSE(wal_path.empty()); + auto wal = WalFile::Create(wal_path); + ASSERT_EQ(wal->open(WalOptions{}), 0); + ASSERT_EQ(wal->prepare_for_read(), 0); + while (true) { + auto record = wal->next(); + ASSERT_TRUE(record.has_value()) << record.error().message(); + if (!record.value().has_value()) break; + const auto &bytes = record.value().value(); + auto doc = Doc::deserialize(reinterpret_cast(bytes.data()), + bytes.size()); + ASSERT_NE(doc, nullptr); + docs->push_back(std::move(doc)); + } + ASSERT_EQ(wal->close(), 0); +} + +// Recreate historical UPSERT records through the serializer and WAL writer so +// the framing and checksums remain valid. New public writes use INSERT/UPDATE. +void RewriteWalAsLegacyUpserts(const std::string &path) { + std::vector docs; + ASSERT_NO_FATAL_FAILURE(ReadWalDocuments(path, &docs)); + ASSERT_FALSE(docs.empty()); + auto wal = WalFile::Create(FindWal(path)); + ASSERT_EQ(wal->remove(), 0); + WalOptions options; + options.create_new = true; + ASSERT_EQ(wal->open(options), 0); + for (auto &doc : docs) { + if (doc->pk_ref() == "target") { + doc->set_operator(Operator::UPSERT); + // Legacy UPSERT did not record its predecessor's ID. + doc->set_doc_id(0); + } + auto bytes = doc->serialize(); + ASSERT_EQ(wal->append(std::string(bytes.begin(), bytes.end())), 0); + } + ASSERT_EQ(wal->flush(), 0); + ASSERT_EQ(wal->close(), 0); +} + +void ExpectOnlyTarget(const Collection::Ptr &collection, + const std::string &value) { + auto fetched = collection->fetch({"target"}); + ASSERT_TRUE(fetched.has_value()) << fetched.error().message(); + ASSERT_NE(fetched.value().at("target"), nullptr); + EXPECT_EQ(fetched.value().at("target")->get("text"), value); + + SearchQuery query; + query.topk_ = 10; + query.filter_ = "text != ''"; + auto matches = collection->query(query); + ASSERT_TRUE(matches.has_value()) << matches.error().message(); + ASSERT_EQ(matches.value().size(), 1u); + EXPECT_EQ(matches.value().front()->pk_ref(), "target"); + EXPECT_EQ(matches.value().front()->get("text"), value); + auto stats = collection->stats(); + ASSERT_TRUE(stats.has_value()) << stats.error().message(); + EXPECT_EQ(stats.value().doc_count, 1u); +} + +class RelaxedValidationDeathTest : public ::testing::Test { + protected: + void SetUp() override { + // Re-exec children start with the default style; select threadsafe before + // InDeathTestChild() interprets their internal death-test flag. + ::testing::FLAGS_gtest_death_test_style = "threadsafe"; + // Threadsafe death tests re-exec the fixture. A second crash child must + // reopen the first child's database instead of deleting it in SetUp. + if (!::testing::internal::InDeathTestChild()) { + ailego::FileHelper::RemovePath(path_.c_str()); + } + } + + void TearDown() override { + ailego::FileHelper::RemovePath(path_.c_str()); + } + + void CheckLegacyCommittedPredecessor(bool separate_segment) { + ASSERT_EXIT(WriteUpsertsAndExit(path_, true, false, separate_segment), + ::testing::ExitedWithCode(0), ""); + if (!::testing::internal::InDeathTestChild()) { + auto recovered_version = VersionManager::Recovery(path_); + ASSERT_TRUE(recovered_version.has_value()); + const auto version = recovered_version.value()->get_current_version(); + EXPECT_EQ(version.persisted_segment_metas().empty(), !separate_segment); + if (!separate_segment) { + EXPECT_FALSE( + version.writing_segment_meta()->persisted_blocks().empty()); + } + ASSERT_EQ( + version.writing_segment_meta()->writing_forward_block()->min_doc_id(), + 1u); + + // Pin the new writer's format before constructing a historical WAL. + std::vector docs; + ASSERT_NO_FATAL_FAILURE(ReadWalDocuments(path_, &docs)); + ASSERT_EQ(docs.size(), 2u); + EXPECT_EQ(docs[0]->get_operator(), Operator::UPDATE); + EXPECT_EQ(docs[0]->doc_id(), 0u); + EXPECT_EQ(docs[1]->get_operator(), Operator::UPDATE); + EXPECT_EQ(docs[1]->doc_id(), 1u); + ASSERT_NO_FATAL_FAILURE(RewriteWalAsLegacyUpserts(path_)); + + auto map = + IDMap::CreateAndOpen("wal_recovery", FindIdMap(path_), false, false); + ASSERT_NE(map, nullptr); + // The map can reach disk ahead of the manifest's deletion snapshot. + // Its latest replay ID no longer identifies committed predecessor ID 0. + ASSERT_TRUE(map->upsert("target", 2).ok()); + ASSERT_TRUE(map->flush().ok()); + } + + // Neither child flushes its recovered deletion bitmap. Both retries must + // rediscover ID 0, including when it belongs to another persisted segment. + for (int attempt = 0; attempt < 2; ++attempt) { + ASSERT_EXIT( + { + auto opened = Collection::Open(path_, CollectionOptions{}); + if (!opened.has_value()) { + std::cerr << opened.error() << std::endl; + std::_Exit(1); + } + ExpectOnlyTarget(opened.value(), "second"); + std::_Exit(::testing::Test::HasFailure() ? 2 : 0); + }, + ::testing::ExitedWithCode(0), ""); + } + + { + auto opened = Collection::Open(path_, CollectionOptions{}); + ASSERT_TRUE(opened.has_value()) << opened.error().message(); + ASSERT_NO_FATAL_FAILURE(ExpectOnlyTarget(opened.value(), "second")); + ASSERT_TRUE(opened.value()->flush().ok()); + } + CollectionOptions options; + options.read_only_ = true; + auto reopened = Collection::Open(path_, options); + ASSERT_TRUE(reopened.has_value()) << reopened.error().message(); + ASSERT_NO_FATAL_FAILURE(ExpectOnlyTarget(reopened.value(), "second")); + auto iterator = reopened.value()->create_iterator(); + ASSERT_TRUE(iterator.has_value()) << iterator.error().message(); + auto first = iterator.value()->next(); + ASSERT_TRUE(first.has_value()); + ASSERT_NE(first.value(), nullptr); + EXPECT_EQ(first.value()->pk_ref(), "target"); + EXPECT_EQ(first.value()->get("text"), "second"); + auto end = iterator.value()->next(); + ASSERT_TRUE(end.has_value()); + EXPECT_EQ(end.value(), nullptr); + } + + const std::string path_{"relaxed_validation_recovery_db"}; +}; + +TEST_F(RelaxedValidationDeathTest, Utf8AndLongIdsRecoverFromUnflushedWal) { + const std::vector ids{u8"订单:😀", + std::string(1021, 'x') + u8"中", " doc ", + u8"café", u8"cafe\u0301"}; + // Re-exec the child before starting collection threads. Exit without stack + // unwinding so Collection destruction cannot flush the writing segment. + ::testing::FLAGS_gtest_death_test_style = "threadsafe"; + ASSERT_EXIT( + { + CollectionSchema schema(u8"恢复 集合"); + if (!schema + .add_field(std::make_shared( + "value", DataType::INT32, false)) + .ok()) { + std::_Exit(1); + } + auto created = + Collection::CreateAndOpen(path_, schema, CollectionOptions{}); + if (!created.has_value()) std::_Exit(2); + auto collection = std::move(created).value(); + std::vector docs; + for (size_t i = 0; i < ids.size(); ++i) { + Doc doc; + doc.set_pk(ids[i]); + doc.set("value", static_cast(i)); + docs.push_back(std::move(doc)); + } + auto inserted = collection->insert(docs); + if (!inserted.has_value()) std::_Exit(3); + for (const auto &status : inserted.value()) { + if (!status.ok()) std::_Exit(4); + } + std::_Exit(0); + }, + ::testing::ExitedWithCode(0), ""); + + auto opened = Collection::Open(path_, CollectionOptions{}); + ASSERT_TRUE(opened.has_value()) << opened.error().message(); + auto collection = std::move(opened).value(); + EXPECT_EQ(collection->schema().value().name(), u8"恢复 集合"); + auto fetched = collection->fetch(ids); + ASSERT_TRUE(fetched.has_value()) << fetched.error().message(); + ASSERT_EQ(fetched.value().size(), ids.size()); + for (size_t i = 0; i < ids.size(); ++i) { + const auto found = fetched.value().find(ids[i]); + ASSERT_NE(found, fetched.value().end()); + ASSERT_NE(found->second, nullptr); + EXPECT_EQ(found->second->pk_ref(), ids[i]); + EXPECT_EQ(found->second->get("value"), static_cast(i)); + } + auto status = collection->flush(); + ASSERT_TRUE(status.ok()) << status.message(); + collection.reset(); + auto reopened = Collection::Open(path_, CollectionOptions{}); + ASSERT_TRUE(reopened.has_value()) << reopened.error().message(); + EXPECT_EQ(reopened.value()->stats().value().doc_count, ids.size()); +} + +TEST_F(RelaxedValidationDeathTest, LargeStringAndTrailingDocRecoverFromWal) { + const std::vector ids{"prefix", std::string(1024, 'x'), + "suffix"}; + // The same value fits below 4MiB with the old 64-byte ID limit. A 1024-byte + // ID takes its serialized WAL record over that former reader-only limit. + const std::vector values{"before", std::string(4193700, 'v'), + "after"}; + ::testing::FLAGS_gtest_death_test_style = "threadsafe"; + ASSERT_EXIT(WriteStringDocsAndExit(path_, ids, values), + ::testing::ExitedWithCode(0), ""); + for (int reopen = 0; reopen < 2; ++reopen) { + auto opened = Collection::Open(path_, CollectionOptions{}); + ASSERT_TRUE(opened.has_value()) << opened.error().message(); + auto collection = std::move(opened).value(); + auto fetched = collection->fetch(ids); + ASSERT_TRUE(fetched.has_value()) << fetched.error().message(); + for (size_t i = 0; i < ids.size(); ++i) { + auto found = fetched.value().find(ids[i]); + ASSERT_NE(found, fetched.value().end()); + ASSERT_NE(found->second, nullptr); + EXPECT_EQ(found->second->pk_ref(), ids[i]); + EXPECT_EQ(found->second->get("text"), values[i]); + } + EXPECT_EQ(collection->stats().value().doc_count, ids.size()); + ASSERT_TRUE(collection->flush().ok()); + } +} + +TEST_F(RelaxedValidationDeathTest, IncompleteTailAllowsLaterCrashRecovery) { + ::testing::FLAGS_gtest_death_test_style = "threadsafe"; + ASSERT_EXIT( + WriteStringDocsAndExit(path_, {"prefix", "torn"}, {"before", "tail"}), + ::testing::ExitedWithCode(0), ""); + if (!::testing::internal::InDeathTestChild()) { + const auto wal_path = FindWal(path_); + ASSERT_FALSE(wal_path.empty()); + std::filesystem::resize_file(wal_path, + std::filesystem::file_size(wal_path) - 1); + const auto tail_bytes = ReadFileBytes(wal_path); + const auto manifests = ReadManifests(path_); + ASSERT_FALSE(manifests.empty()); + { + CollectionOptions options; + options.read_only_ = true; + auto opened = Collection::Open(path_, options); + // Recovery needs writable index stores. A read-only attempt must fail + // explicitly and leave the WAL and committed manifest untouched. + ASSERT_FALSE(opened.has_value()); + EXPECT_EQ(opened.error().code(), StatusCode::FAILED_PRECONDITION); + EXPECT_NE(opened.error().message().find("read-write mode once"), + std::string::npos); + } + EXPECT_EQ(ReadFileBytes(wal_path), tail_bytes); + EXPECT_EQ(ReadManifests(path_), manifests); + } + + ASSERT_EXIT( + { + auto opened = Collection::Open(path_, CollectionOptions{}); + if (!opened.has_value()) { + std::cerr << opened.error() << std::endl; + std::_Exit(1); + } + Doc doc; + doc.set_pk("suffix"); + doc.set("text", "after"); + std::vector docs{doc}; + auto inserted = opened.value()->insert(docs); + if (!inserted.has_value() || !inserted.value().front().ok()) + std::_Exit(2); + std::_Exit(0); + }, + ::testing::ExitedWithCode(0), ""); + auto opened = Collection::Open(path_, CollectionOptions{}); + ASSERT_TRUE(opened.has_value()) << opened.error().message(); + auto fetched = opened.value()->fetch({"prefix", "torn", "suffix"}); + ASSERT_TRUE(fetched.has_value()); + ASSERT_NE(fetched.value().at("prefix"), nullptr); + ASSERT_NE(fetched.value().at("suffix"), nullptr); + EXPECT_TRUE(fetched.value().find("torn") == fetched.value().end() || + fetched.value().at("torn") == nullptr); + EXPECT_EQ(opened.value()->stats().value().doc_count, 2); +} + +TEST_F(RelaxedValidationDeathTest, CorruptWalFailsOpenWithoutReplacingFiles) { + ::testing::FLAGS_gtest_death_test_style = "threadsafe"; + ASSERT_EXIT( + WriteStringDocsAndExit(path_, {"prefix", "broken"}, {"before", "after"}), + ::testing::ExitedWithCode(0), ""); + const auto wal_path = FindWal(path_); + ASSERT_FALSE(wal_path.empty()); + { + std::fstream file(wal_path, + std::ios::in | std::ios::out | std::ios::binary); + ASSERT_TRUE(file.is_open()); + uint32_t first_length; + file.seekg(64); + file.read(reinterpret_cast(&first_length), sizeof(first_length)); + ASSERT_TRUE(file.good()); + // Damage the second record CRC, keeping its complete framing intact. + file.seekp(64 + 8 + first_length + 4); + const uint32_t bad_crc = 0; + file.write(reinterpret_cast(&bad_crc), sizeof(bad_crc)); + ASSERT_TRUE(file.good()); + } + const auto bytes = ReadFileBytes(wal_path); + const auto manifests = ReadManifests(path_); + ASSERT_FALSE(manifests.empty()); + for (int attempt = 0; attempt < 2; ++attempt) { + auto opened = Collection::Open(path_, CollectionOptions{}); + ASSERT_FALSE(opened.has_value()); + EXPECT_NE(opened.error().message().find("CRC mismatch"), std::string::npos); + EXPECT_EQ(ReadFileBytes(wal_path), bytes); + EXPECT_EQ(ReadManifests(path_), manifests); + } +} + +TEST_F(RelaxedValidationDeathTest, InvalidDocumentPayloadFailsRecovery) { + ::testing::FLAGS_gtest_death_test_style = "threadsafe"; + ASSERT_EXIT(WriteStringDocsAndExit(path_, {"prefix"}, {"before"}, true), + ::testing::ExitedWithCode(0), ""); + const auto wal_path = FindWal(path_); + ASSERT_FALSE(wal_path.empty()); + const auto bytes = ReadFileBytes(wal_path); + const auto manifests = ReadManifests(path_); + auto opened = Collection::Open(path_, CollectionOptions{}); + ASSERT_FALSE(opened.has_value()); + EXPECT_NE(opened.error().message().find("Corrupt WAL document"), + std::string::npos); + EXPECT_EQ(ReadFileBytes(wal_path), bytes); + EXPECT_EQ(ReadManifests(path_), manifests); +} + +TEST_F(RelaxedValidationDeathTest, CorruptTailDoesNotApplyUpsertPrefix) { + ::testing::FLAGS_gtest_death_test_style = "threadsafe"; + ASSERT_EXIT(WriteUpsertsAndExit(path_, true, true), + ::testing::ExitedWithCode(0), ""); + ASSERT_NO_FATAL_FAILURE(RewriteWalAsLegacyUpserts(path_)); + const auto wal_path = FindWal(path_); + ASSERT_FALSE(wal_path.empty()); + auto bytes = ReadFileBytes(wal_path); + size_t prefix_end = 64; + for (int i = 0; i < 2; ++i) { + uint32_t length; + ASSERT_GE(bytes.size() - prefix_end, 8u); + std::memcpy(&length, bytes.data() + prefix_end, sizeof(length)); + prefix_end += 8 + length; + ASSERT_LE(prefix_end, bytes.size()); + } + ASSERT_GE(bytes.size() - prefix_end, 8u); + bytes[prefix_end + 4] ^= 1; // Corrupt only the trailing record's CRC. + { + std::ofstream file(wal_path, std::ios::binary | std::ios::trunc); + file.write(bytes.data(), bytes.size()); + ASSERT_TRUE(file.good()); + } + const auto idmap_path = FindIdMap(path_); + ASSERT_FALSE(idmap_path.empty()); + { + auto map = IDMap::CreateAndOpen("wal_recovery", idmap_path, false, false); + ASSERT_NE(map, nullptr); + // The original row is committed as ID 0. Make the pre-recovery mapping + // deterministic even if RocksDB flushed uncommitted writes before exit. + ASSERT_TRUE(map->upsert("target", 0).ok()); + map->remove("broken"); + ASSERT_TRUE(map->flush().ok()); + } + const auto manifests = ReadManifests(path_); + for (int attempt = 0; attempt < 2; ++attempt) { + auto opened = Collection::Open(path_, CollectionOptions{}); + ASSERT_FALSE(opened.has_value()); + EXPECT_NE(opened.error().message().find("CRC mismatch"), std::string::npos); + EXPECT_EQ(ReadFileBytes(wal_path), bytes); + EXPECT_EQ(ReadManifests(path_), manifests); + auto map = IDMap::CreateAndOpen("wal_recovery", idmap_path, false, true); + ASSERT_NE(map, nullptr); + uint64_t original_id; + ASSERT_TRUE(map->has("target", &original_id)); + EXPECT_EQ(original_id, 0u); + EXPECT_FALSE(map->has("broken")); + } + + // Remove the damaged last record and retry the intact UPSERT prefix. + std::filesystem::resize_file(wal_path, prefix_end); + { + auto opened = Collection::Open(path_, CollectionOptions{}); + ASSERT_TRUE(opened.has_value()) << opened.error().message(); + auto fetched = opened.value()->fetch({"target"}); + ASSERT_TRUE(fetched.has_value()); + ASSERT_NE(fetched.value().at("target"), nullptr); + EXPECT_EQ(fetched.value().at("target")->get("text"), "second"); + EXPECT_EQ(opened.value()->stats().value().doc_count, 1); + Doc update; + update.set_pk("target"); + update.set("text", "third"); + std::vector updates{update}; + auto updated = opened.value()->upsert(updates); + ASSERT_TRUE(updated.has_value()); + ASSERT_TRUE(updated.value().front().ok()); + ASSERT_TRUE(opened.value()->flush().ok()); + } + auto reopened = Collection::Open(path_, CollectionOptions{}); + ASSERT_TRUE(reopened.has_value()); + auto fetched = reopened.value()->fetch({"target"}); + ASSERT_TRUE(fetched.has_value()); + ASSERT_NE(fetched.value().at("target"), nullptr); + EXPECT_EQ(fetched.value().at("target")->get("text"), "third"); + EXPECT_EQ(reopened.value()->stats().value().doc_count, 1); +} + +TEST_F(RelaxedValidationDeathTest, UpsertReplayIgnoresSameAndLaterReplayIds) { + ::testing::FLAGS_gtest_death_test_style = "threadsafe"; + ASSERT_EXIT(WriteUpsertsAndExit(path_, false, false), + ::testing::ExitedWithCode(0), ""); + if (!::testing::internal::InDeathTestChild()) { + ASSERT_NO_FATAL_FAILURE(RewriteWalAsLegacyUpserts(path_)); + const auto idmap_path = FindIdMap(path_); + ASSERT_FALSE(idmap_path.empty()); + auto map = IDMap::CreateAndOpen("wal_recovery", idmap_path, false, false); + ASSERT_NE(map, nullptr); + // Simulate a partial previous recovery that persisted the second UPSERT's + // ID. It is later than the first replay ID and equal to the second. + ASSERT_TRUE(map->upsert("target", 1).ok()); + ASSERT_TRUE(map->flush().ok()); + } + ASSERT_EXIT( + { + auto opened = Collection::Open(path_, CollectionOptions{}); + if (!opened.has_value()) { + std::cerr << opened.error() << std::endl; + std::_Exit(1); + } + auto fetched = opened.value()->fetch({"target"}); + if (!fetched.has_value() || !fetched.value().at("target") || + fetched.value().at("target")->get("text") != + "second" || + opened.value()->stats().value().doc_count != 1) + std::_Exit(2); + Doc update; + update.set_pk("target"); + update.set("text", "third"); + std::vector updates{update}; + auto updated = opened.value()->upsert(updates); + if (!updated.has_value() || !updated.value().front().ok()) + std::_Exit(3); + std::_Exit(0); + }, + ::testing::ExitedWithCode(0), ""); + auto reopened = Collection::Open(path_, CollectionOptions{}); + ASSERT_TRUE(reopened.has_value()) << reopened.error().message(); + auto fetched = reopened.value()->fetch({"target"}); + ASSERT_TRUE(fetched.has_value()); + ASSERT_NE(fetched.value().at("target"), nullptr); + EXPECT_EQ(fetched.value().at("target")->get("text"), "third"); + EXPECT_EQ(reopened.value()->stats().value().doc_count, 1); +} + +TEST_F(RelaxedValidationDeathTest, + LegacyUpsertsRecoverCommittedPredecessorInWritingSegment) { + ASSERT_NO_FATAL_FAILURE(CheckLegacyCommittedPredecessor(false)); +} + +TEST_F(RelaxedValidationDeathTest, + LegacyUpsertsRecoverCommittedPredecessorInPersistedSegment) { + ASSERT_NO_FATAL_FAILURE(CheckLegacyCommittedPredecessor(true)); +} + +TEST_F(RelaxedValidationDeathTest, UpsertWalRecordsInsertOrUpdatePredecessor) { + ASSERT_EXIT(WriteUpsertsAndExit(path_, false, false), + ::testing::ExitedWithCode(0), ""); + std::vector docs; + ASSERT_NO_FATAL_FAILURE(ReadWalDocuments(path_, &docs)); + ASSERT_EQ(docs.size(), 2u); + EXPECT_EQ(docs[0]->get_operator(), Operator::INSERT); + EXPECT_EQ(docs[0]->pk_ref(), "target"); + EXPECT_EQ(docs[0]->get("text"), "first"); + EXPECT_EQ(docs[1]->get_operator(), Operator::UPDATE); + EXPECT_EQ(docs[1]->doc_id(), 0u); + EXPECT_EQ(docs[1]->get("text"), "second"); +} + +} // namespace +} // namespace zvec diff --git a/tests/db/index/common/doc_test.cc b/tests/db/index/common/doc_test.cc index e47174a43..f8bfe6a9c 100644 --- a/tests/db/index/common/doc_test.cc +++ b/tests/db/index/common/doc_test.cc @@ -14,6 +14,7 @@ #include "zvec/db/doc.h" #include +#include #include #include #include @@ -770,7 +771,7 @@ TEST_F(DocDetailedTest, ValidateAndSanitization) { ASSERT_TRUE(s.ok()); } - // pk with characters inside the allowed set is accepted + // Previously valid ASCII IDs remain accepted, including the old boundary. { auto schema = test::TestHelper::CreateNormalSchema(false); std::vector valid_names = { @@ -805,8 +806,10 @@ TEST_F(DocDetailedTest, ValidateAndSanitization) { "file-name_v1.2", // -, _, . allowed "a-b_c.d!@#$%+=.", // all specials in one - // Max length = 64 + // Former and new length boundaries std::string(64, 'a'), + std::string(65, 'a'), + std::string(1024, 'a'), std::string(63, 'a') + "_", "_" + std::string(62, 'x') + ".", "!" + std::string(62, '0') + "@", @@ -819,16 +822,11 @@ TEST_F(DocDetailedTest, ValidateAndSanitization) { } } - // pk that is too long or uses disallowed characters is rejected + // External IDs can contain punctuation, spaces and UTF-8 text. { auto schema = test::TestHelper::CreateNormalSchema(false); - std::vector invalid_names = { - // Too long (>64) - std::string(65, 'a'), std::string(64, 'a') + "_", - - // Illegal characters - "a b", // space - "a&b", // & not in set + std::vector valid_names = { + " ", " padded ", "a b", "a&b", "a*b", // * "a(b)", // ( ) "a:b", // : @@ -851,12 +849,35 @@ TEST_F(DocDetailedTest, ValidateAndSanitization) { "a,b", // , "用户", // non-ASCII (Chinese) "αβγ", // Greek - "résumé", // accented chars (é not in [a-zA-Z]) + "résumé", // accented characters }; - for (auto pk : invalid_names) { + for (const auto &pk : valid_names) { auto doc = test::TestHelper::CreateDoc(1, *schema, pk); auto s = doc.validate_and_sanitize(schema); - ASSERT_FALSE(s.ok()) << "expected invalid pk: " << pk; + ASSERT_TRUE(s.ok()) << "expected valid pk: " << pk << ": " << s.message(); + } + } + + // Invalid text must be rejected before it can enter storage. + { + auto schema = test::TestHelper::CreateNormalSchema(false); + const std::vector invalid_ids = { + "", + std::string(1025, 'a'), + std::string("a\0b", 3), + "a\nb", + "a\tb", + "a\rb", + std::string("\xff", 1), + std::string("\xe4\xb8", 2), + }; + for (const auto &pk : invalid_ids) { + auto doc = test::TestHelper::CreateDoc(1, *schema, pk); + doc.set_pk(pk); // The helper generates a default ID for an empty input. + auto s = doc.validate_and_sanitize(schema); + ASSERT_EQ(s.code(), StatusCode::INVALID_ARGUMENT); + EXPECT_EQ(s.message().find("Invalid doc: "), 0u); + EXPECT_EQ(s.message().find("offset"), std::string::npos); } } } @@ -1005,6 +1026,14 @@ TEST_F(DocDetailedTest, SerializeValueCoverage) { auto buffer = doc.serialize(); EXPECT_FALSE(buffer.empty()); + // Exercise every truncated header, scalar, vector and nested string using + // independently allocated buffers, so ASan also catches accidental overreads. + for (size_t size = 0; size < buffer.size(); ++size) { + SCOPED_TRACE(size); + std::vector truncated(buffer.begin(), buffer.begin() + size); + EXPECT_EQ(Doc::deserialize(truncated.data(), truncated.size()), nullptr); + } + auto deserialized_doc = Doc::deserialize(buffer.data(), buffer.size()); EXPECT_NE(deserialized_doc, nullptr); @@ -1623,3 +1652,122 @@ TEST_F(DocDetailedTest, FieldExistenceChecks) { auto type_mismatch_opt = doc.get("existent"); EXPECT_FALSE(type_mismatch_opt.has_value()); } + + +TEST_F(DocDetailedTest, DeserializeRejectsMalformedLengthsAndTags) { + Doc doc; + doc.set_pk("id"); + doc.set("field", "payload"); + const auto valid = doc.serialize(); + const size_t operation = + sizeof(uint32_t) + 2 + sizeof(float) + sizeof(uint64_t); + const size_t fields_count = operation + sizeof(uint32_t); + const size_t field = fields_count + sizeof(uint32_t); + const size_t tag = field + sizeof(uint32_t) + 5; + const auto maximum = std::numeric_limits::max(); + for (size_t offset : {size_t{0}, fields_count, field, tag + 1}) { + SCOPED_TRACE(offset); + auto invalid = valid; + std::memcpy(invalid.data() + offset, &maximum, sizeof(maximum)); + EXPECT_EQ(Doc::deserialize(invalid.data(), invalid.size()), nullptr); + } + auto invalid = valid; + std::memcpy(invalid.data() + operation, &maximum, sizeof(maximum)); + EXPECT_EQ(Doc::deserialize(invalid.data(), invalid.size()), nullptr); + invalid = valid; + invalid[tag] = 255; + EXPECT_EQ(Doc::deserialize(invalid.data(), invalid.size()), nullptr); + invalid = valid; + invalid.push_back(0); + EXPECT_EQ(Doc::deserialize(invalid.data(), invalid.size()), nullptr); + invalid = valid; + const uint32_t two = 2; + std::memcpy(invalid.data() + fields_count, &two, sizeof(two)); + invalid.insert(invalid.end(), valid.begin() + field, valid.end()); + EXPECT_EQ(Doc::deserialize(invalid.data(), invalid.size()), nullptr); + EXPECT_EQ(Doc::deserialize(nullptr, valid.size()), nullptr); + EXPECT_EQ(Doc::deserialize(valid.data(), 1), nullptr); +} + +TEST_F(DocDetailedTest, DeserializeChecksNestedCountsAndBooleanRepresentation) { + const size_t tag = sizeof(uint32_t) + 2 + sizeof(float) + sizeof(uint64_t) + + sizeof(uint32_t) + sizeof(uint32_t) + sizeof(uint32_t) + 1; + const std::vector values{ + std::vector{1, 2}, std::vector{"hello", "world"}, + std::vector{true, false}, + std::pair, std::vector>{{1, 2}, {1.f, 2.f}}, + std::pair, std::vector>{ + {1, 2}, {zvec::float16_t(1.f), zvec::float16_t(2.f)}}}; + for (const auto &value : values) { + Doc doc; + doc.set_pk("id"); + std::visit( + [&](const auto &item) { + if constexpr (std::is_same_v, + std::monostate>) { + doc.set_null("v"); + } else { + doc.set("v", item); + } + }, + value); + auto buffer = doc.serialize(); + ASSERT_NE(Doc::deserialize(buffer.data(), buffer.size()), nullptr); + // A UINT32_MAX element count cannot fit into any of these payloads. + std::fill_n(buffer.begin() + tag + 1, sizeof(uint32_t), uint8_t{0xff}); + EXPECT_EQ(Doc::deserialize(buffer.data(), buffer.size()), nullptr); + } + Doc sparse; + sparse.set_pk("id"); + sparse.set("v", std::pair, std::vector>{ + {1, 2}, {1.f, 2.f}}); + auto buffer = sparse.serialize(); + std::fill_n( + buffer.begin() + tag + 1 + sizeof(uint32_t) + 2 * sizeof(uint32_t), + sizeof(uint32_t), uint8_t{0xff}); + EXPECT_EQ(Doc::deserialize(buffer.data(), buffer.size()), nullptr); + Doc boolean; + boolean.set_pk("id"); + boolean.set("v", true); + buffer = boolean.serialize(); + buffer[tag + 1] = 2; + EXPECT_EQ(Doc::deserialize(buffer.data(), buffer.size()), nullptr); + boolean.set("v", std::vector{true}); + buffer = boolean.serialize(); + buffer[tag + 1 + sizeof(uint32_t)] = 2; + EXPECT_EQ(Doc::deserialize(buffer.data(), buffer.size()), nullptr); +} + +TEST_F(DocDetailedTest, DeserializePreservesHistoricalTextWithoutRevalidation) { + Doc doc; + doc.set_pk(std::string("old\0id", 6)); + doc.set("old\nfield", std::string("\xff\0value", 7)); + const auto buffer = doc.serialize(); + const auto restored = Doc::deserialize(buffer.data(), buffer.size()); + ASSERT_NE(restored, nullptr); + EXPECT_EQ(restored->pk_ref(), doc.pk_ref()); + EXPECT_EQ(restored->get("old\nfield"), + doc.get("old\nfield")); +} + +TEST_F(DocDetailedTest, ValidationErrorsEscapeAndBoundDocumentAndFieldNames) { + auto schema = std::make_shared( + "test", FieldSchemaPtrList{ + std::make_shared("value", DataType::INT32)}); + for (const auto &name : + std::vector{std::string("bad\0field", 9), "bad\nfield", + "\xff", std::string(10000, 'x')}) { + Doc doc; + doc.set_pk(std::string(1024, 'd')); + doc.set(name, int32_t{42}); + const auto status = doc.validate_and_sanitize(schema); + EXPECT_EQ(status.code(), StatusCode::INVALID_ARGUMENT); + EXPECT_EQ(status.message().find("Invalid doc:"), 0u); + EXPECT_NE(status.message().find("does not exist"), std::string::npos); + EXPECT_LT(status.message().size(), 256u); + for (unsigned char byte : status.message()) { + EXPECT_GE(byte, 0x20); + EXPECT_LE(byte, 0x7e); + } + } +} diff --git a/tests/db/index/common/manifest_codec_golden_test.cc b/tests/db/index/common/manifest_codec_golden_test.cc index 329f6199e..7846022e5 100644 --- a/tests/db/index/common/manifest_codec_golden_test.cc +++ b/tests/db/index/common/manifest_codec_golden_test.cc @@ -1081,6 +1081,100 @@ TEST(ManifestCodecGolden, IndexParamsLastBranchWins) { EXPECT_EQ(decoded->type(), IndexType::HNSW); } +TEST(ManifestCodecGolden, LegacyFieldNamesSurviveWithoutRevalidation) { + // Older alter_column paths could persist names outside the schema rules. + // Build the wire data without schema helpers so this also catches a decoder + // silently dropping fields after a new validation check is added. + const std::vector names{"user name", std::string(65, 'f'), + u8"历史字段", "_zvec_uid_"}; + std::string encoded_schema; + pbwire::Writer schema_writer(&encoded_schema); + schema_writer.PutString(1, "legacy_collection"); + for (const auto &name : names) { + std::string encoded_field; + pbwire::Writer field_writer(&encoded_field); + field_writer.PutString(1, name); + field_writer.PutVarint(2, 2); // Persisted STRING data type. + schema_writer.PutMessage(2, encoded_field); + } + schema_writer.PutVarint(3, 10000); + std::string encoded_manifest; + pbwire::Writer(&encoded_manifest).PutMessage(2, encoded_schema); + + ManifestData restored; + auto status = ManifestCodec::Decode(encoded_manifest, &restored); + ASSERT_TRUE(status.ok()) << status.message(); + ASSERT_NE(restored.schema, nullptr); + ASSERT_EQ(restored.schema->fields().size(), names.size()); + for (size_t i = 0; i < names.size(); ++i) { + SCOPED_TRACE(i); + const auto *field = restored.schema->get_field(names[i]); + ASSERT_NE(field, nullptr); + EXPECT_EQ(field->name(), names[i]); + EXPECT_EQ(field->data_type(), DataType::STRING); + EXPECT_EQ(restored.schema->fields()[i]->name(), names[i]); + } + EXPECT_EQ(restored.schema->validate().code(), StatusCode::INVALID_ARGUMENT); + + std::string reencoded; + status = ManifestCodec::Encode(restored, &reencoded); + ASSERT_TRUE(status.ok()) << status.message(); + EXPECT_EQ(reencoded, encoded_manifest); +} + +TEST(ManifestCodecGolden, DuplicateFieldsFailInsteadOfBeingSilentlyDropped) { + // This malformed schema could previously be created through the C++ list + // constructor. Restoring only its first field silently changes its meaning. + CollectionSchema schema( + "legacy", {std::make_shared("duplicate", DataType::INT32), + std::make_shared("duplicate", DataType::INT64)}); + std::string schema_bytes; + ManifestCodec::EncodeCollectionSchema(schema, &schema_bytes); + auto decoded_schema = ManifestCodec::DecodeCollectionSchema(schema_bytes); + ASSERT_FALSE(decoded_schema.has_value()); + EXPECT_EQ(decoded_schema.error().code(), StatusCode::INTERNAL_ERROR); + EXPECT_NE(decoded_schema.error().message().find("duplicate"), + std::string::npos); + + std::string manifest_bytes; + pbwire::Writer(&manifest_bytes).PutMessage(2, schema_bytes); + ManifestData restored; + auto status = ManifestCodec::Decode(manifest_bytes, &restored); + EXPECT_EQ(status.code(), StatusCode::INTERNAL_ERROR); + EXPECT_EQ(restored.schema, nullptr); +} + +TEST(ManifestCodecGolden, MalformedNestedSchemaFailsExplicitly) { + const std::string malformed_schema("\x0a\x05x", 3); + std::string manifest_bytes; + pbwire::Writer(&manifest_bytes).PutMessage(2, malformed_schema); + ManifestData restored; + auto status = ManifestCodec::Decode(manifest_bytes, &restored); + EXPECT_EQ(status.code(), StatusCode::INTERNAL_ERROR); + EXPECT_EQ(restored.schema, nullptr); +} + +TEST(ManifestCodecGolden, StructuralSchemaHelpersPreserveLegacyNames) { + CollectionSchema schema("legacy_collection"); + auto status = schema.add_field( + std::make_shared("user name", DataType::STRING)); + ASSERT_TRUE(status.ok()) << status.message(); + status = schema.alter_field("user name", std::make_shared( + u8"历史字段", DataType::STRING)); + ASSERT_TRUE(status.ok()) << status.message(); + EXPECT_FALSE(schema.has_field("user name")); + ASSERT_TRUE(schema.has_field(u8"历史字段")); + + // Structural mutation is also used during recovery. Explicit validation is + // kept separate, and legacy names can still be replaced with legal names. + EXPECT_EQ(schema.validate().code(), StatusCode::INVALID_ARGUMENT); + status = schema.alter_field( + u8"历史字段", std::make_shared("renamed", DataType::STRING)); + ASSERT_TRUE(status.ok()) << status.message(); + EXPECT_TRUE(schema.validate().ok()); + EXPECT_TRUE(schema.has_field("renamed")); +} + TEST(ManifestCodecGolden, UnknownFieldsAreIgnored) { // Forward compatibility: a manifest written by a newer zvec may carry fields // this build does not know about. They must be skipped silently. diff --git a/tests/db/index/common/name_validation_test.cc b/tests/db/index/common/name_validation_test.cc new file mode 100644 index 000000000..def1424e6 --- /dev/null +++ b/tests/db/index/common/name_validation_test.cc @@ -0,0 +1,289 @@ +// Copyright 2025-present the zvec project +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include "db/index/common/name_validation.h" +#include +#include +#include +#include +#include + +namespace zvec { +namespace { + +struct Utf8NameValidator { + Status (*validate)(std::string_view); + size_t max_bytes; + const char *prefix; +}; + +const std::array kUtf8NameValidators{{ + {ValidateDocumentId, kMaxDocumentIdBytes, "Invalid doc: id"}, + {ValidateCollectionName, kMaxCollectionNameBytes, + "Invalid schema: collection name"}, +}}; + +void ExpectInvalid(const Status &status, const std::string &message) { + EXPECT_EQ(status.code(), StatusCode::INVALID_ARGUMENT); + EXPECT_EQ(status.message(), message); + EXPECT_EQ(status.message().find("offset"), std::string::npos); +} + +std::string Repeat(std::string_view text, size_t count) { + std::string result; + result.reserve(text.size() * count); + for (size_t i = 0; i < count; ++i) { + result.append(text.data(), text.size()); + } + return result; +} + +TEST(NameValidationTest, AcceptsUnicodePunctuationAndSpaces) { + const std::vector values{ + "a", + "A_b-9", + "https://example.com/docs/1?lang=zh#section", + "a/b.c:d@e+f", + "'quoted' \"text\" \\ value", + " ", + " ", + " leading and trailing ", + u8"\u00A0\u3000", + u8"中文", + u8"€", + u8"😀", + u8"👩\u200D💻", + u8"é", + u8"e\u0301", + u8"\U0010FFFF", + }; + for (const auto &validator : kUtf8NameValidators) { + SCOPED_TRACE(validator.prefix); + for (const auto &value : values) { + EXPECT_TRUE(validator.validate(value).ok()); + } + } +} + +TEST(NameValidationTest, RejectsEmptyNames) { + for (const auto &validator : kUtf8NameValidators) { + ExpectInvalid(validator.validate(std::string_view{}), + std::string(validator.prefix) + " must not be empty"); + } + ExpectInvalid(ValidateFieldName(""), + "Invalid schema: field name must not be empty"); +} + +TEST(NameValidationTest, MeasuresLimitsInUtf8Bytes) { + for (const auto &validator : kUtf8NameValidators) { + SCOPED_TRACE(validator.prefix); + const auto max_bytes = validator.max_bytes; + EXPECT_TRUE(validator.validate(std::string(max_bytes, 'a')).ok()); + auto emoji = Repeat(u8"😀", max_bytes / 4); + ASSERT_EQ(emoji.size(), max_bytes); + EXPECT_TRUE(validator.validate(emoji).ok()); + EXPECT_TRUE( + validator.validate(std::string(max_bytes - 3, 'a') + u8"中").ok()); + + const auto expected = std::string(validator.prefix) + " exceeds " + + std::to_string(max_bytes) + " bytes (got " + + std::to_string(max_bytes + 1) + ")"; + ExpectInvalid(validator.validate(std::string(max_bytes + 1, 'a')), + expected); + ExpectInvalid(validator.validate(emoji + "a"), expected); + ExpectInvalid(validator.validate(std::string(max_bytes - 2, 'a') + u8"中"), + expected); + } +} + +TEST(NameValidationTest, RejectsMalformedUtf8) { + const std::vector malformed{ + "\x80", // Isolated continuation byte. + "\xBF", + "\xC2", // Truncated sequences. + "\xE4\xB8", + "\xF0\x9F\x98", + "\xC2" + "A", // Invalid continuation byte. + "\xE2" + "A" + "\xAC", + "\xC0\x80", // Overlong NUL. + "\xC1\xBF", // Overlong two-byte sequence. + "\xE0\x80\xAF", // Overlong three-byte sequence. + "\xF0\x80\x80\xAF", // Overlong four-byte sequence. + "\xED\xA0\x80", // UTF-16 surrogate U+D800. + "\xED\xBF\xBF", // UTF-16 surrogate U+DFFF. + "\xF4\x90\x80\x80", // Above U+10FFFF. + "\xF5\x80\x80\x80", // Invalid lead byte. + "\xFE", + "\xFF", + }; + for (const auto &validator : kUtf8NameValidators) { + SCOPED_TRACE(validator.prefix); + for (const auto &value : malformed) { + ExpectInvalid(validator.validate(value), + std::string(validator.prefix) + " is not valid UTF-8"); + ExpectInvalid(validator.validate("prefix" + value), + std::string(validator.prefix) + " is not valid UTF-8"); + } + } +} + +TEST(NameValidationTest, HonorsStringViewLengthAndEmbeddedNulls) { + const std::string backing = std::string(u8"中文") + "\xFF"; + for (const auto &validator : kUtf8NameValidators) { + SCOPED_TRACE(validator.prefix); + EXPECT_TRUE(validator.validate(std::string_view(backing.data(), 6)).ok()); + ExpectInvalid(validator.validate(std::string_view(backing.data(), 5)), + std::string(validator.prefix) + " is not valid UTF-8"); + ExpectInvalid(validator.validate(std::string("a\0b", 3)), + std::string(validator.prefix) + " contains a null character"); + } +} + +TEST(NameValidationTest, RejectsEveryC0AndC1ControlByCodepoint) { + for (const auto &validator : kUtf8NameValidators) { + SCOPED_TRACE(validator.prefix); + for (unsigned int codepoint = 0; codepoint <= 0x9F; ++codepoint) { + if (codepoint >= 0x20 && codepoint < 0x7F) { + continue; + } + SCOPED_TRACE(codepoint); + std::string value; + if (codepoint >= 0x80) { + value += '\xC2'; + } + value += static_cast(codepoint); + std::string reason = "contains a control character"; + if (codepoint == 0) { + reason = "contains a null character"; + } else if (codepoint == '\n' || codepoint == '\r') { + reason = "contains a newline"; + } else if (codepoint == '\t') { + reason = "contains a tab"; + } + ExpectInvalid(validator.validate("a" + value + "b"), + std::string(validator.prefix) + " " + reason); + } + // These continuation bytes overlap the C1 byte range, but their decoded + // codepoints are ordinary letters/symbols and must not be rejected. + EXPECT_TRUE(validator.validate(u8"中文€😀").ok()); + } +} + +TEST(NameValidationTest, DistinguishesUnicodeLineAndParagraphSeparators) { + for (const auto &validator : kUtf8NameValidators) { + ExpectInvalid(validator.validate(u8"a\u2028b"), + std::string(validator.prefix) + " contains a line separator"); + ExpectInvalid( + validator.validate(u8"a\u2029b"), + std::string(validator.prefix) + " contains a paragraph separator"); + } +} + +TEST(NameValidationTest, RetainsTheFieldAsciiCharacterSet) { + for (const std::string name : + {"a", "Z", "0", "_", "-", "a_b-c1", "123_test", "_zvec_custom"}) { + EXPECT_TRUE(ValidateFieldName(name).ok()); + } + EXPECT_TRUE(ValidateFieldName("ABCDEFGHIJKLMNOPQRSTUVWXYZ" + "abcdefghijklmnopqrstuvwxyz0123456789_-") + .ok()); + EXPECT_TRUE(ValidateFieldName(std::string(kMaxFieldNameBytes, 'a')).ok()); + ExpectInvalid(ValidateFieldName(std::string(kMaxFieldNameBytes + 1, 'a')), + "Invalid schema: field name exceeds 64 bytes (got 65)"); + ExpectInvalid(ValidateFieldName(std::string(10000, 'a')), + "Invalid schema: field name exceeds 64 bytes (got 10000)"); +} + +TEST(NameValidationTest, RejectsExactInternalFieldNames) { + for (const std::string name : + {"_zvec_row_id_", "_zvec_g_doc_id_", "_zvec_uid_", "_zvec_score", + "_zvec_group_id"}) { + SCOPED_TRACE(name); + ExpectInvalid(ValidateFieldName(name), + "Invalid schema: field[" + name + + "] is reserved; use a different name"); + // The restriction is an exact match, not a new prefix or case policy. + EXPECT_TRUE(ValidateFieldName(name + "_custom").ok()); + EXPECT_TRUE(ValidateDocumentId(name).ok()); + EXPECT_TRUE(ValidateCollectionName(name).ok()); + } + EXPECT_TRUE(ValidateFieldName("_ZVEC_UID_").ok()); + for (const std::string name : + {"_zvec_vector", "_zvec_sindices", "_zvec_svalues", "_zvec_is_valid"}) { + EXPECT_TRUE(ValidateFieldName(name).ok()); + } +} + +TEST(NameValidationTest, SharedErrorPreviewIsEscapedAndBounded) { + EXPECT_EQ(FormatNameForError(""), ""); + EXPECT_EQ(FormatNameForError(std::string("a\0\n\r\t[]\\", 8)), + "a\\0\\n\\r\\t\\[\\]\\\\"); + EXPECT_EQ(FormatNameForError(u8"中"), "\\xE4\\xB8\\xAD"); + EXPECT_EQ(FormatNameForError(std::string(10000, '\xff')), + Repeat("\\xFF", 32) + "..."); + EXPECT_EQ(FormatNameForError(std::string(10000, 'x')), + std::string(32, 'x') + "..."); +} + +TEST(NameValidationTest, DescribesInvalidFieldCharactersWithSafePreviews) { + const std::string rule = + "; use letters (A-Z, a-z), digits, underscores (_) or hyphens (-)"; + ExpectInvalid(ValidateFieldName("user name"), + "Invalid schema: field[user name] contains a space" + rule); + ExpectInvalid( + ValidateFieldName("a.b"), + "Invalid schema: field[a.b] contains an unsupported character" + rule); + ExpectInvalid(ValidateFieldName(u8"中"), + "Invalid schema: field[\\xE4\\xB8\\xAD] contains a non-ASCII " + "character" + + rule); + ExpectInvalid( + ValidateFieldName("\x80"), + "Invalid schema: field[\\x80] contains a non-ASCII character" + rule); + ExpectInvalid( + ValidateFieldName(std::string("a\0b", 3)), + "Invalid schema: field[a\\0b] contains a null character" + rule); + ExpectInvalid(ValidateFieldName("a\nb"), + "Invalid schema: field[a\\nb] contains a newline" + rule); + ExpectInvalid(ValidateFieldName("a\tb"), + "Invalid schema: field[a\\tb] contains a tab" + rule); + ExpectInvalid( + ValidateFieldName("a\x1B" + "b"), + "Invalid schema: field[a\\x1Bb] contains a control character" + rule); + ExpectInvalid(ValidateFieldName("][\\\n"), + "Invalid schema: field[\\]\\[\\\\\\n] contains an unsupported " + "character" + + rule); +} + +TEST(NameValidationTest, BoundsInvalidFieldPreviewsAndNeverEchoesRawBytes) { + const std::string name(64, '\xFF'); + const auto status = ValidateFieldName(name); + ExpectInvalid(status, + "Invalid schema: field[" + Repeat("\\xFF", 32) + + "...] contains a non-ASCII character; use letters " + "(A-Z, a-z), digits, underscores (_) or hyphens (-)"); + EXPECT_LT(status.message().size(), 256u); + for (unsigned char byte : status.message()) { + EXPECT_GE(byte, 0x20); + EXPECT_LE(byte, 0x7E); + } +} + +} // namespace +} // namespace zvec diff --git a/tests/db/index/common/schema_test.cc b/tests/db/index/common/schema_test.cc index 84260d02d..897f6382f 100644 --- a/tests/db/index/common/schema_test.cc +++ b/tests/db/index/common/schema_test.cc @@ -19,6 +19,54 @@ using namespace zvec; +TEST(CollectionSchemaTest, RejectsNullFieldObjectsWithoutDroppingThem) { + auto field = std::make_shared("valid", DataType::INT32); + EXPECT_THROW(CollectionSchema("schema", {nullptr}), std::invalid_argument); + EXPECT_THROW(CollectionSchema("schema", {field, nullptr}), + std::invalid_argument); + + CollectionSchema schema("schema", {field}); + const CollectionSchema before(schema); + EXPECT_EQ(schema.add_field(nullptr).code(), StatusCode::INVALID_ARGUMENT); + EXPECT_EQ(schema.alter_field("valid", nullptr).code(), + StatusCode::INVALID_ARGUMENT); + EXPECT_EQ(schema, before); + EXPECT_TRUE(schema.validate().ok()); +} + +TEST(CollectionSchemaTest, ValidatesDuplicateNamesFromConstructorsAndCopies) { + CollectionSchema schema( + "schema", {std::make_shared("duplicate", DataType::INT32), + std::make_shared("duplicate", DataType::INT64)}); + // Retain the invalid input until explicit validation, rather than hiding one + // field and allowing a collection whose schema changes after reopening. + ASSERT_EQ(schema.fields().size(), 2u); + auto status = schema.validate(); + EXPECT_EQ(status.code(), StatusCode::INVALID_ARGUMENT); + EXPECT_EQ(status.message(), + "Invalid schema: duplicate field name [duplicate]; field names " + "must be unique"); + CollectionSchema copied(schema); + CollectionSchema assigned; + assigned = schema; + EXPECT_EQ(copied.validate(), status); + EXPECT_EQ(assigned.validate(), status); + EXPECT_EQ(copied.fields().size(), 2u); + EXPECT_EQ(assigned.fields().size(), 2u); +} + +TEST(CollectionSchemaTest, CopyOwnsIndependentFieldObjects) { + auto original_field = + std::make_shared("value", DataType::INT32, true); + CollectionSchema schema("schema", {original_field}); + original_field->set_name("invalid name"); + original_field->set_data_type(DataType::STRING); + EXPECT_TRUE(schema.validate().ok()); + ASSERT_NE(schema.get_field("value"), nullptr); + EXPECT_EQ(schema.get_field("value")->data_type(), DataType::INT32); + EXPECT_FALSE(schema.has_field("invalid name")); +} + TEST(FieldSchemaTest, DefaultConstructor) { FieldSchema field; EXPECT_EQ(field.name(), ""); @@ -554,7 +602,9 @@ TEST(FieldSchemaTest, Validate) { "user_name", "test-123", "aBc123_-", - std::string(32, 'a'), // max len = 32 + std::string(32, 'a'), + std::string(33, 'a'), + std::string(64, 'a'), // max len = 64 "a_b-c1", "__test__", "123_test"}; @@ -571,7 +621,7 @@ TEST(FieldSchemaTest, Validate) { { std::vector invalid_names = { "", // empty — len < 1 - std::string(33, 'a'), // len > 32 + std::string(65, 'a'), // len > 64 "a b", // space "a.b", "a@b", @@ -888,7 +938,7 @@ TEST(CollectionSchemaTest, Validate) { ASSERT_EQ(s.code(), StatusCode::INVALID_ARGUMENT); CollectionSchema c2("c2", {}); - s = c1.validate(); + s = c2.validate(); ASSERT_FALSE(s.ok()); ASSERT_EQ(s.code(), StatusCode::INVALID_ARGUMENT); @@ -901,8 +951,7 @@ TEST(CollectionSchemaTest, Validate) { auto f2 = std::make_shared("f2", DataType::INT32); CollectionSchema c4("c4", {f2}); s = c4.validate(); - ASSERT_FALSE(s.ok()); - ASSERT_EQ(s.code(), StatusCode::INVALID_ARGUMENT); + ASSERT_TRUE(s.ok()); auto f3 = std::make_shared("f3", DataType::VECTOR_FP16); CollectionSchema c5("c5", {f3}); @@ -910,19 +959,16 @@ TEST(CollectionSchemaTest, Validate) { ASSERT_FALSE(s.ok()); ASSERT_EQ(s.code(), StatusCode::INVALID_ARGUMENT); - // validate collection name regex "^[a-zA-Z0-9_-]{3,32}$" + // Collection names are bounded UTF-8 strings, independent of the path. { std::vector invalid_names = { - "", // empty - "ab", // too short (<3) - std::string(65, 'a'), // too long (>64) - "a b", // space not allowed - "a.b", // dot not allowed - "a$b", // $ not allowed - "中文", // non-ASCII - "a\nb", // newline not allowed - "a\tb", // tab not allowed - "a\rb", // carriage return not allowed + "", // empty + std::string(257, 'a'), // too long (>256 bytes) + std::string("a\0b", 3), // embedded NUL + std::string("\xff", 1), // invalid UTF-8 + "a\nb", // newline not allowed + "a\tb", // tab not allowed + "a\rb", // carriage return not allowed }; for (const auto &name : invalid_names) { @@ -937,14 +983,24 @@ TEST(CollectionSchemaTest, Validate) { std::vector valid_names = { "test_collection_supported_vectors", + "a", + "ab", std::string(64, 'a'), + std::string(65, 'a'), + std::string(256, 'a'), + "a b", + "a.b", + "a$b", + "中文", + " ", + " padded ", "a_b", // underscore allowed "a-b", // dash allowed "a_1", // underscore and digit allowed "a-1", // dash and digit allowed "a_1b", // underscore, digit and letter allowed "a-1b", // dash, digit and letter allowed - "-start", // allowed! (regex permits leading -/_) + "-start", // leading -/_ remains allowed "_start", // also allowed "end-", "end_", // trailing -/_ allowed diff --git a/tests/db/index/storage/wal_file_test.cc b/tests/db/index/storage/wal_file_test.cc index 50cd122da..a29e80e3d 100644 --- a/tests/db/index/storage/wal_file_test.cc +++ b/tests/db/index/storage/wal_file_test.cc @@ -12,20 +12,14 @@ // See the License for the specific language governing permissions and // limitations under the License. -#ifdef _MSC_VER -#define _ALLOW_KEYWORD_MACROS -#endif -#define private public -#define protected public #include "db/index/storage/wal/wal_file.h" -#undef private -#undef protected - #include #include #include #include +#include #include +#include #include #include #include @@ -47,6 +41,16 @@ class WalFileTest : public testing::Test { } void TearDown() override {} + + // Legacy success-path tests use empty string as their loop sentinel, but + // still assert that the new API did not report an error. + std::string ReadRecord(const WalFilePtr &wal_file) { + auto result = wal_file->next(); + EXPECT_TRUE(result.has_value()) + << (result.has_value() ? "" : result.error().message()); + if (!result.has_value() || !result.value().has_value()) return {}; + return std::move(result.value().value()); + } }; TEST_F(WalFileTest, TestGeneral) { @@ -126,14 +130,14 @@ TEST_F(WalFileTest, TestGeneral) { uint32_t idx = 0; ret = wal_file->prepare_for_read(); ASSERT_EQ(ret, 0); - std::string record = wal_file->next(); + std::string record = ReadRecord(wal_file); while (!record.empty()) { if (idx < 100) { ASSERT_EQ(record, "hello"); } else { ASSERT_EQ(record, std::string("hello") + std::to_string(idx)); } - record = wal_file->next(); + record = ReadRecord(wal_file); idx++; } ASSERT_EQ(idx, 400); @@ -205,9 +209,9 @@ TEST_F(WalFileTest, TestMultiThread) { uint32_t idx = 0; ret = wal_file->prepare_for_read(); ASSERT_EQ(ret, 0); - std::string record = wal_file->next(); + std::string record = ReadRecord(wal_file); while (!record.empty()) { - record = wal_file->next(); + record = ReadRecord(wal_file); idx++; } ASSERT_EQ(idx, 30000); @@ -243,9 +247,9 @@ TEST_F(WalFileTest, TestBoundaryCondition) { ret = wal_file->open(wal_option); ASSERT_EQ(ret, 0); uint32_t idx = 0; - std::string record = wal_file->next(); + std::string record = ReadRecord(wal_file); while (!record.empty()) { - record = wal_file->next(); + record = ReadRecord(wal_file); idx++; } ASSERT_EQ(idx, 0); @@ -271,13 +275,13 @@ TEST_F(WalFileTest, TestBoundaryCondition) { idx = 0; ret = wal_file->prepare_for_read(); ASSERT_EQ(ret, 0); - record = wal_file->next(); + record = ReadRecord(wal_file); while (!record.empty()) { ASSERT_EQ(record.size(), 4); for (size_t i = 0; i < 4; i++) { ASSERT_EQ(record[i], i); } - record = wal_file->next(); + record = ReadRecord(wal_file); idx++; } ASSERT_EQ(idx, 1); @@ -312,13 +316,13 @@ TEST_F(WalFileTest, TestBoundaryCondition) { idx = 0; ret = wal_file->prepare_for_read(); ASSERT_EQ(ret, 0); - record = wal_file->next(); + record = ReadRecord(wal_file); while (!record.empty()) { ASSERT_EQ(record.size(), BIG_DATA_SIZE); for (size_t i = 0; i < BIG_DATA_SIZE; i++) { ASSERT_EQ((uint8_t)record[i], i % 256); } - record = wal_file->next(); + record = ReadRecord(wal_file); idx++; } ASSERT_EQ(idx, 1); @@ -349,10 +353,10 @@ TEST_F(WalFileTest, TestBoundaryCondition) { idx = 0; ret = wal_file->prepare_for_read(); ASSERT_EQ(ret, 0); - record = wal_file->next(); + record = ReadRecord(wal_file); while (!record.empty()) { ASSERT_EQ(record, std::string("hello") + std::to_string(idx)); - record = wal_file->next(); + record = ReadRecord(wal_file); idx++; } ASSERT_EQ(idx, 99); @@ -417,16 +421,15 @@ TEST_F(WalFileTest, TestFirstErrorCase) { ret = wal_file->open(wal_option); ASSERT_EQ(ret, 0); - uint32_t idx = 0; - ret = wal_file->prepare_for_read(); - ASSERT_EQ(ret, 0); - std::string record = wal_file->next(); - while (!record.empty()) { - ASSERT_EQ(record, "hello"); - record = wal_file->next(); - idx++; + ASSERT_EQ(wal_file->prepare_for_read(), 0); + for (size_t i = 0; i < 0; ++i) { + EXPECT_EQ(ReadRecord(wal_file), "hello"); } - ASSERT_EQ(idx, 0); + auto result = wal_file->next(); + ASSERT_FALSE(result.has_value()); + EXPECT_EQ(result.error().code(), StatusCode::INTERNAL_ERROR); + EXPECT_NE(result.error().message().find("CRC mismatch"), std::string::npos); + EXPECT_NE(wal_file->append("after corruption"), 0); // close ret = wal_file->close(); ASSERT_EQ(ret, 0); @@ -477,16 +480,15 @@ TEST_F(WalFileTest, TestMiddleErrorCase) { ret = wal_file->open(wal_option); ASSERT_EQ(ret, 0); - uint32_t idx = 0; - ret = wal_file->prepare_for_read(); - ASSERT_EQ(ret, 0); - std::string record = wal_file->next(); - while (!record.empty()) { - ASSERT_EQ(record, "hello"); - record = wal_file->next(); - idx++; + ASSERT_EQ(wal_file->prepare_for_read(), 0); + for (size_t i = 0; i < 5; ++i) { + EXPECT_EQ(ReadRecord(wal_file), "hello"); } - ASSERT_EQ(idx, 5); + auto result = wal_file->next(); + ASSERT_FALSE(result.has_value()); + EXPECT_EQ(result.error().code(), StatusCode::INTERNAL_ERROR); + EXPECT_NE(result.error().message().find("CRC mismatch"), std::string::npos); + EXPECT_NE(wal_file->append("after corruption"), 0); // close ret = wal_file->close(); ASSERT_EQ(ret, 0); @@ -538,10 +540,10 @@ TEST_F(WalFileTest, TestLastErrorCase) { uint32_t idx = 0; ret = wal_file->prepare_for_read(); ASSERT_EQ(ret, 0); - std::string record = wal_file->next(); + std::string record = ReadRecord(wal_file); while (!record.empty()) { ASSERT_EQ(record, "hello"); - record = wal_file->next(); + record = ReadRecord(wal_file); idx++; } ASSERT_EQ(idx, 9); @@ -594,16 +596,15 @@ TEST_F(WalFileTest, TestLengthSmallErrorCase) { ret = wal_file->open(wal_option); ASSERT_EQ(ret, 0); - uint32_t idx = 0; - ret = wal_file->prepare_for_read(); - ASSERT_EQ(ret, 0); - std::string record = wal_file->next(); - while (!record.empty()) { - ASSERT_EQ(record, "hello"); - record = wal_file->next(); - idx++; + ASSERT_EQ(wal_file->prepare_for_read(), 0); + for (size_t i = 0; i < 0; ++i) { + EXPECT_EQ(ReadRecord(wal_file), "hello"); } - ASSERT_EQ(idx, 0); + auto result = wal_file->next(); + ASSERT_FALSE(result.has_value()); + EXPECT_EQ(result.error().code(), StatusCode::INTERNAL_ERROR); + EXPECT_NE(result.error().message().find("CRC mismatch"), std::string::npos); + EXPECT_NE(wal_file->append("after corruption"), 0); // close ret = wal_file->close(); ASSERT_EQ(ret, 0); @@ -643,7 +644,7 @@ TEST_F(WalFileTest, TestLengthBigErrorCase) { dir_path, "data.wal.", std::to_string(segment_id)); int wal_fd = open(wal_path.c_str(), O_RDWR, 0644); ASSERT_GT(wal_fd, 0); - uint32_t err_length = 200; // exceed file size 130 + uint32_t err_length = std::numeric_limits::max(); lseek(wal_fd, 64, SEEK_SET); write(wal_fd, (const void *)&err_length, 4); @@ -657,10 +658,10 @@ TEST_F(WalFileTest, TestLengthBigErrorCase) { uint32_t idx = 0; ret = wal_file->prepare_for_read(); ASSERT_EQ(ret, 0); - std::string record = wal_file->next(); + std::string record = ReadRecord(wal_file); while (!record.empty()) { ASSERT_EQ(record, "hello"); - record = wal_file->next(); + record = ReadRecord(wal_file); idx++; } ASSERT_EQ(idx, 0); @@ -714,16 +715,15 @@ TEST_F(WalFileTest, TestCRCErrorCase) { ret = wal_file->open(wal_option); ASSERT_EQ(ret, 0); - uint32_t idx = 0; - ret = wal_file->prepare_for_read(); - ASSERT_EQ(ret, 0); - std::string record = wal_file->next(); - while (!record.empty()) { - ASSERT_EQ(record, "hello"); - record = wal_file->next(); - idx++; + ASSERT_EQ(wal_file->prepare_for_read(), 0); + for (size_t i = 0; i < 1; ++i) { + EXPECT_EQ(ReadRecord(wal_file), "hello"); } - ASSERT_EQ(idx, 1); + auto result = wal_file->next(); + ASSERT_FALSE(result.has_value()); + EXPECT_EQ(result.error().code(), StatusCode::INTERNAL_ERROR); + EXPECT_NE(result.error().message().find("CRC mismatch"), std::string::npos); + EXPECT_NE(wal_file->append("after corruption"), 0); // close ret = wal_file->close(); ASSERT_EQ(ret, 0); @@ -732,6 +732,115 @@ TEST_F(WalFileTest, TestCRCErrorCase) { ASSERT_EQ(ret, 0); } +TEST_F(WalFileTest, RecordLargerThanFourMiBPreservesFollowingRecord) { + const std::string path = "./data.wal.large"; + auto wal = WalFile::Create(path); + WalOptions options; + options.create_new = true; + ASSERT_EQ(wal->open(options), 0); + const std::string large(4 * 1024 * 1024 + 1024, 'x'); + ASSERT_EQ(wal->append("prefix"), 0); + ASSERT_EQ(wal->append(std::string(large)), 0); + ASSERT_EQ(wal->append("suffix"), 0); + ASSERT_EQ(wal->close(), 0); + options.create_new = false; + ASSERT_EQ(wal->open(options), 0); + ASSERT_EQ(wal->prepare_for_read(), 0); + EXPECT_EQ(ReadRecord(wal), "prefix"); + EXPECT_EQ(ReadRecord(wal), large); + EXPECT_EQ(ReadRecord(wal), "suffix"); + auto end = wal->next(); + ASSERT_TRUE(end.has_value()); + EXPECT_FALSE(end.value().has_value()); +} + +TEST_F(WalFileTest, IncompleteTailIsRemovedOnlyBeforeAppend) { + const std::string path = "./data.wal.tail"; + constexpr size_t prefix_end = 64 + 8 + 6; + // Exercise every partial header and partial payload boundary. + for (size_t tail_size = 1; tail_size < 8 + 4; ++tail_size) { + SCOPED_TRACE(tail_size); + auto wal = WalFile::Create(path); + WalOptions options; + options.create_new = true; + ASSERT_EQ(wal->open(options), 0); + ASSERT_EQ(wal->append("prefix"), 0); + ASSERT_EQ(wal->append("torn"), 0); + ASSERT_EQ(wal->close(), 0); + ailego::File file; + ASSERT_TRUE(file.open(path, false)); + ASSERT_TRUE(file.truncate(prefix_end + tail_size)); + file.close(); + + options.create_new = false; + ASSERT_EQ(wal->open(options), 0); + ASSERT_EQ(wal->prepare_for_read(), 0); + EXPECT_EQ(ReadRecord(wal), "prefix"); + auto end = wal->next(); + ASSERT_TRUE(end.has_value()); + EXPECT_FALSE(end.value().has_value()); + ASSERT_TRUE(file.open(path, true)); + EXPECT_EQ(file.size(), prefix_end + tail_size); + file.close(); + + ASSERT_EQ(wal->append("suffix"), 0); + ASSERT_EQ(wal->close(), 0); + ASSERT_EQ(wal->open(options), 0); + ASSERT_EQ(wal->prepare_for_read(), 0); + EXPECT_EQ(ReadRecord(wal), "prefix"); + EXPECT_EQ(ReadRecord(wal), "suffix"); + end = wal->next(); + ASSERT_TRUE(end.has_value()); + EXPECT_FALSE(end.value().has_value()); + ASSERT_EQ(wal->remove(), 0); + } +} + +TEST_F(WalFileTest, ClosedReaderReturnsError) { + auto wal = WalFile::Create("./data.wal.closed"); + auto result = wal->next(); + ASSERT_FALSE(result.has_value()); + EXPECT_EQ(result.error().code(), StatusCode::INTERNAL_ERROR); +} + +TEST_F(WalFileTest, ZeroLengthRecordIsCorruption) { + const std::string path = "./data.wal.zero"; + auto wal = WalFile::Create(path); + WalOptions options; + options.create_new = true; + ASSERT_EQ(wal->open(options), 0); + EXPECT_NE(wal->append(""), 0); + ASSERT_EQ(wal->append("payload"), 0); + ASSERT_EQ(wal->close(), 0); + ailego::File file; + ASSERT_TRUE(file.open(path, false)); + const uint32_t length = 0; + ASSERT_EQ(file.write(64, &length, sizeof(length)), sizeof(length)); + file.close(); + options.create_new = false; + ASSERT_EQ(wal->open(options), 0); + ASSERT_EQ(wal->prepare_for_read(), 0); + auto result = wal->next(); + ASSERT_FALSE(result.has_value()); + EXPECT_NE(result.error().message().find("zero length"), std::string::npos); +} + +TEST_F(WalFileTest, TruncatedFileHeaderReturnsError) { + const std::string path = "./data.wal.header"; + auto wal = WalFile::Create(path); + WalOptions options; + options.create_new = true; + ASSERT_EQ(wal->open(options), 0); + ASSERT_EQ(wal->close(), 0); + ailego::File file; + ASSERT_TRUE(file.open(path, false)); + ASSERT_TRUE(file.truncate(63)); + file.close(); + options.create_new = false; + ASSERT_EQ(wal->open(options), 0); + EXPECT_NE(wal->prepare_for_read(), 0); +} + #if defined(__GNUC__) || defined(__GNUG__) #pragma GCC diagnostic pop #endif \ No newline at end of file diff --git a/tests/db/relaxed_validation_test.cc b/tests/db/relaxed_validation_test.cc new file mode 100644 index 000000000..ef7978067 --- /dev/null +++ b/tests/db/relaxed_validation_test.cc @@ -0,0 +1,412 @@ +// Copyright 2025-present the zvec project +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include "db/common/constants.h" + +namespace zvec { +namespace { + +class RelaxedValidationTest : public ::testing::Test { + protected: + void SetUp() override { + ailego::MemoryLimitPool::get_instance().init(2 * 1024ll * 1024ll * 1024ll); + ailego::FileHelper::RemovePath(path_.c_str()); + } + + void TearDown() override { + collection_.reset(); + ailego::FileHelper::RemovePath(path_.c_str()); + } + + CollectionSchema MakeSchema(const std::string &name = "x", + const std::string &field = "value") { + CollectionSchema schema(name); + EXPECT_TRUE(schema + .add_field(std::make_shared( + field, DataType::INT32, false)) + .ok()); + return schema; + } + + void Create(const CollectionSchema &schema) { + auto result = Collection::CreateAndOpen(path_, schema, options_); + ASSERT_TRUE(result.has_value()) << result.error().message(); + collection_ = std::move(result).value(); + } + + void Reopen() { + collection_.reset(); + auto result = Collection::Open(path_, options_); + ASSERT_TRUE(result.has_value()) << result.error().message(); + collection_ = std::move(result).value(); + } + + Doc MakeDoc(const std::string &id, int32_t value, + const std::string &field = "value") { + Doc doc; + doc.set_pk(id); + EXPECT_TRUE(doc.set(field, value)); + return doc; + } + + void ExpectWrite(const Result &result, size_t count) { + ASSERT_TRUE(result.has_value()) << result.error().message(); + ASSERT_EQ(result.value().size(), count); + for (const auto &status : result.value()) { + ASSERT_TRUE(status.ok()) << status.message(); + } + } + + void ExpectValue(const std::string &id, int32_t expected, + const std::string &field = "value") { + auto result = collection_->fetch({id}); + ASSERT_TRUE(result.has_value()) << result.error().message(); + ASSERT_EQ(result.value().size(), 1u); + auto found = result.value().find(id); + ASSERT_NE(found, result.value().end()); + ASSERT_NE(found->second, nullptr); + EXPECT_EQ(found->second->pk(), id); + EXPECT_EQ(found->second->get(field), expected); + } + + const std::string path_{"relaxed_validation_test_db"}; + CollectionOptions options_; + Collection::Ptr collection_; +}; + +TEST_F(RelaxedValidationTest, Utf8IdsKeepTheirIdentityAcrossCrudAndReopen) { + const std::string name = u8"测试 集合/v1"; + ASSERT_NO_FATAL_FAILURE(Create(MakeSchema(name))); + const std::vector ids = { + "user:123", "https://example.com/document/42", + u8"订单-😀", std::string(1021, 'i') + u8"中", + "doc", " doc", + "doc ", " ", + u8"café", u8"cafe\u0301", + "DOC"}; + std::vector docs; + for (size_t i = 0; i < ids.size(); ++i) { + docs.push_back(MakeDoc(ids[i], static_cast(i))); + } + ASSERT_NO_FATAL_FAILURE(ExpectWrite(collection_->insert(docs), docs.size())); + + std::vector updates{MakeDoc(ids[0], 100)}; + ASSERT_NO_FATAL_FAILURE(ExpectWrite(collection_->update(updates), 1)); + std::vector upserts{MakeDoc(ids[3], 103), MakeDoc(u8"新增:文档", 200)}; + ASSERT_NO_FATAL_FAILURE(ExpectWrite(collection_->upsert(upserts), 2)); + ASSERT_NO_FATAL_FAILURE(ExpectWrite(collection_->delete_({ids[1]}), 1)); + + auto flush_status = collection_->flush(); + ASSERT_TRUE(flush_status.ok()) << flush_status.message(); + ASSERT_NO_FATAL_FAILURE(Reopen()); + auto schema = collection_->schema(); + ASSERT_TRUE(schema.has_value()) << schema.error().message(); + EXPECT_EQ(schema.value().name(), name); + for (size_t i = 0; i < ids.size(); ++i) { + if (i == 1) { + continue; + } + int32_t expected = static_cast(i); + if (i == 0) expected = 100; + if (i == 3) expected = 103; + ASSERT_NO_FATAL_FAILURE(ExpectValue(ids[i], expected)); + } + ASSERT_NO_FATAL_FAILURE(ExpectValue(u8"新增:文档", 200)); + auto deleted = collection_->fetch({ids[1]}); + ASSERT_TRUE(deleted.has_value()) << deleted.error().message(); + ASSERT_EQ(deleted.value().size(), 1u); + EXPECT_EQ(deleted.value().at(ids[1]), nullptr); +} + +TEST_F(RelaxedValidationTest, ShortAndMaximumLengthCollectionNamesPersist) { + for (const auto &name : std::vector{ + "x", "xy", u8"集", std::string(253, 'n') + u8"集"}) { + SCOPED_TRACE(name); + ASSERT_NO_FATAL_FAILURE(Create(MakeSchema(name))); + std::vector docs{MakeDoc("id", 1)}; + ASSERT_NO_FATAL_FAILURE(ExpectWrite(collection_->insert(docs), 1)); + auto status = collection_->flush(); + ASSERT_TRUE(status.ok()) << status.message(); + ASSERT_NO_FATAL_FAILURE(Reopen()); + auto schema = collection_->schema(); + ASSERT_TRUE(schema.has_value()) << schema.error().message(); + EXPECT_EQ(schema.value().name(), name); + ASSERT_NO_FATAL_FAILURE(ExpectValue("id", 1)); + status = collection_->destroy(); + ASSERT_TRUE(status.ok()) << status.message(); + collection_.reset(); + } +} + +TEST_F(RelaxedValidationTest, MaximumLengthFieldsSupportIndexesAndFilters) { + const std::string scalar = "s" + std::string(63, 'a'); + const std::string vector = "v" + std::string(63, 'b'); + auto schema = MakeSchema("x", scalar); + ASSERT_TRUE(schema + .add_field(std::make_shared( + vector, DataType::VECTOR_FP32, 4, false, + std::make_shared(MetricType::L2))) + .ok()); + ASSERT_NO_FATAL_FAILURE(Create(schema)); + const std::vector values{1.0f, 2.0f, 3.0f, 4.0f}; + Doc doc = MakeDoc(u8"文档:1", 42, scalar); + ASSERT_TRUE(doc.set>(vector, values)); + std::vector docs{doc}; + ASSERT_NO_FATAL_FAILURE(ExpectWrite(collection_->insert(docs), 1)); + auto status = collection_->flush(); + ASSERT_TRUE(status.ok()) << status.message(); + status = + collection_->create_index(scalar, std::make_shared()); + ASSERT_TRUE(status.ok()) << status.message(); + status = collection_->create_index( + vector, std::make_shared(MetricType::L2)); + ASSERT_TRUE(status.ok()) << status.message(); + ASSERT_NO_FATAL_FAILURE(Reopen()); + + // Fetch supplies the query vector, covering lookup by a newly allowed ID. + auto fetched = collection_->fetch({doc.pk()}); + ASSERT_TRUE(fetched.has_value()) << fetched.error().message(); + ASSERT_EQ(fetched.value().size(), 1u); + const auto stored_vector = + fetched.value().at(doc.pk())->get>(vector); + ASSERT_TRUE(stored_vector.has_value()); + EXPECT_EQ(stored_vector.value(), values); + SearchQuery query; + query.topk_ = 1; + query.target_.field_name_ = vector; + query.target_.set_vector( + std::string(reinterpret_cast(stored_vector->data()), + stored_vector->size() * sizeof(float))); + query.filter_ = scalar + " = 42"; + query.output_fields_ = std::vector{scalar}; + auto matches = collection_->query(query); + ASSERT_TRUE(matches.has_value()) << matches.error().message(); + ASSERT_EQ(matches.value().size(), 1u); + EXPECT_EQ(matches.value()[0]->pk(), doc.pk()); + EXPECT_EQ(matches.value()[0]->get(scalar), 42); +} + +TEST_F(RelaxedValidationTest, InvalidRenameLeavesSchemaAndDataUnchanged) { + ASSERT_NO_FATAL_FAILURE(Create(MakeSchema())); + std::vector docs{MakeDoc("id", 42)}; + ASSERT_NO_FATAL_FAILURE(ExpectWrite(collection_->insert(docs), 1)); + const auto before = collection_->schema().value(); + for (const auto &name : std::vector{ + "user name", "../value", u8"字段", std::string(65, 'f')}) { + SCOPED_TRACE(name); + auto status = collection_->alter_column("value", name); + ASSERT_EQ(status.code(), StatusCode::INVALID_ARGUMENT); + EXPECT_EQ(status.message().find("Invalid schema:"), 0u); + EXPECT_EQ(status.message().find("offset"), std::string::npos); + EXPECT_EQ(collection_->schema().value(), before); + ASSERT_NO_FATAL_FAILURE(ExpectValue("id", 42)); + } + ASSERT_NO_FATAL_FAILURE(Reopen()); + EXPECT_EQ(collection_->schema().value(), before); + ASSERT_NO_FATAL_FAILURE(ExpectValue("id", 42)); + + const std::string renamed(64, 'r'); + auto status = collection_->alter_column("value", renamed); + ASSERT_TRUE(status.ok()) << status.message(); + ASSERT_NO_FATAL_FAILURE(Reopen()); + EXPECT_FALSE(collection_->schema().value().has_field("value")); + EXPECT_TRUE(collection_->schema().value().has_field(renamed)); + ASSERT_NO_FATAL_FAILURE(ExpectValue("id", 42, renamed)); +} + +TEST_F(RelaxedValidationTest, InvalidIdRejectsWholeBatchBeforeWriting) { + ASSERT_NO_FATAL_FAILURE(Create(MakeSchema())); + std::vector initial{MakeDoc("existing", 1)}; + ASSERT_NO_FATAL_FAILURE(ExpectWrite(collection_->insert(initial), 1)); + for (int operation = 0; operation < 3; ++operation) { + SCOPED_TRACE(operation); + const std::string first_id = operation == 0 ? "new:id" : "existing"; + std::vector batch{MakeDoc(first_id, 99), + MakeDoc(std::string("bad\0id", 6), 100)}; + auto result = operation == 0 ? collection_->insert(batch) + : operation == 1 ? collection_->update(batch) + : collection_->upsert(batch); + ASSERT_FALSE(result.has_value()); + EXPECT_EQ(result.error().code(), StatusCode::INVALID_ARGUMENT); + EXPECT_EQ(result.error().message().find("Invalid doc:"), 0u); + EXPECT_NE(result.error().message().find("null character"), + std::string::npos); + EXPECT_EQ(result.error().message().find("offset"), std::string::npos); + ASSERT_NO_FATAL_FAILURE(ExpectValue("existing", 1)); + auto missing = collection_->fetch({"new:id"}); + ASSERT_TRUE(missing.has_value()) << missing.error().message(); + ASSERT_EQ(missing.value().size(), 1u); + EXPECT_EQ(missing.value().at("new:id"), nullptr); + } + ASSERT_NO_FATAL_FAILURE(Reopen()); + ASSERT_NO_FATAL_FAILURE(ExpectValue("existing", 1)); + EXPECT_EQ(collection_->stats().value().doc_count, 1u); +} + +TEST_F(RelaxedValidationTest, FetchAndDeleteKeepLookupSemantics) { + ASSERT_NO_FATAL_FAILURE(Create(MakeSchema())); + // These are invalid for a new document, but lookup must retain its existing + // missing-key behavior rather than introducing input validation errors. + const std::vector absent_ids{"", std::string("bad\0id", 6), + std::string(1025, 'x'), "\xff"}; + auto fetched = collection_->fetch(absent_ids); + ASSERT_TRUE(fetched.has_value()) << fetched.error().message(); + ASSERT_EQ(fetched.value().size(), absent_ids.size()); + for (const auto &id : absent_ids) { + EXPECT_EQ(fetched.value().at(id), nullptr); + } + auto deleted = collection_->delete_(absent_ids); + ASSERT_TRUE(deleted.has_value()) << deleted.error().message(); + ASSERT_EQ(deleted.value().size(), absent_ids.size()); + for (const auto &status : deleted.value()) { + EXPECT_EQ(status.code(), StatusCode::NOT_FOUND); + } +} + +TEST_F(RelaxedValidationTest, ReservedNamesAndDuplicatesFailBeforeCreation) { + for (const std::string name : + {"_zvec_uid_", "_zvec_g_doc_id_", "_zvec_row_id_", "_zvec_score", + "_zvec_group_id"}) { + SCOPED_TRACE(name); + auto result = + Collection::CreateAndOpen(path_, MakeSchema("x", name), options_); + ASSERT_FALSE(result.has_value()); + EXPECT_EQ(result.error().code(), StatusCode::INVALID_ARGUMENT); + EXPECT_NE(result.error().message().find("is reserved"), std::string::npos); + EXPECT_FALSE(ailego::FileHelper::IsExist(path_.c_str())); + } + CollectionSchema duplicate( + "x", {std::make_shared("value", DataType::INT32), + std::make_shared("value", DataType::INT64)}); + auto result = Collection::CreateAndOpen(path_, duplicate, options_); + ASSERT_FALSE(result.has_value()); + EXPECT_EQ(result.error().code(), StatusCode::INVALID_ARGUMENT); + EXPECT_NE(result.error().message().find("duplicate field name"), + std::string::npos); + EXPECT_FALSE(ailego::FileHelper::IsExist(path_.c_str())); +} + +TEST_F(RelaxedValidationTest, ReservedDdlTargetsLeaveTheCollectionUnchanged) { + ASSERT_NO_FATAL_FAILURE(Create(MakeSchema())); + std::vector docs{MakeDoc("id", 42)}; + ASSERT_NO_FATAL_FAILURE(ExpectWrite(collection_->insert(docs), 1)); + const auto before = collection_->schema().value(); + for (const std::string name : + {"_zvec_uid_", "_zvec_g_doc_id_", "_zvec_row_id_", "_zvec_score", + "_zvec_group_id"}) { + SCOPED_TRACE(name); + auto status = collection_->add_column( + std::make_shared(name, DataType::INT32, true), ""); + EXPECT_EQ(status.code(), StatusCode::INVALID_ARGUMENT); + EXPECT_NE(status.message().find("is reserved"), std::string::npos); + status = collection_->alter_column("value", name); + EXPECT_EQ(status.code(), StatusCode::INVALID_ARGUMENT); + EXPECT_NE(status.message().find("is reserved"), std::string::npos); + status = collection_->alter_column( + "value", "", std::make_shared(name, DataType::INT32)); + EXPECT_EQ(status.code(), StatusCode::INVALID_ARGUMENT); + EXPECT_NE(status.message().find("is reserved"), std::string::npos); + EXPECT_EQ(collection_->schema().value(), before); + ASSERT_NO_FATAL_FAILURE(ExpectValue("id", 42)); + } + ASSERT_NO_FATAL_FAILURE(Reopen()); + EXPECT_EQ(collection_->schema().value(), before); + ASSERT_NO_FATAL_FAILURE(ExpectValue("id", 42)); +} + +TEST_F(RelaxedValidationTest, DdlRetainsFieldCountInvariants) { + CollectionSchema schema("x"); + for (uint32_t i = 0; i < kMaxScalarFieldSize; ++i) { + ASSERT_TRUE(schema + .add_field(std::make_shared( + "f" + std::to_string(i), DataType::INT32, true)) + .ok()); + } + ASSERT_NO_FATAL_FAILURE(Create(schema)); + auto status = collection_->add_column( + std::make_shared("excess", DataType::INT32, true), ""); + EXPECT_EQ(status.code(), StatusCode::INVALID_ARGUMENT); + EXPECT_NE(status.message().find("1024 scalar fields"), std::string::npos); + EXPECT_EQ(collection_->schema().value(), schema); + EXPECT_FALSE(collection_->schema().value().has_field("excess")); + status = collection_->destroy(); + ASSERT_TRUE(status.ok()) << status.message(); + collection_.reset(); + + ASSERT_NO_FATAL_FAILURE(Create(MakeSchema())); + std::vector docs{MakeDoc("id", 42)}; + ASSERT_NO_FATAL_FAILURE(ExpectWrite(collection_->insert(docs), 1)); + status = collection_->drop_column("value"); + EXPECT_EQ(status.code(), StatusCode::INVALID_ARGUMENT); + EXPECT_NE(status.message().find("last field"), std::string::npos); + ASSERT_NO_FATAL_FAILURE(ExpectValue("id", 42)); + ASSERT_NO_FATAL_FAILURE(Reopen()); + EXPECT_TRUE(collection_->schema().value().has_field("value")); + ASSERT_NO_FATAL_FAILURE(ExpectValue("id", 42)); +} + +TEST_F(RelaxedValidationTest, DdlSnapshotsCallerOwnedFieldSchemas) { + ASSERT_NO_FATAL_FAILURE(Create(MakeSchema())); + std::vector initial{MakeDoc("original", 1)}; + ASSERT_NO_FATAL_FAILURE(ExpectWrite(collection_->insert(initial), 1)); + auto added = std::make_shared("extra", DataType::INT32, true); + auto status = collection_->add_column(added, ""); + ASSERT_TRUE(status.ok()) << status.message(); + added->set_name("bad name"); + added->set_data_type(DataType::STRING); + added->set_nullable(false); + + auto current = collection_->schema().value(); + ASSERT_TRUE(current.has_field("extra")); + EXPECT_EQ(current.get_field("extra")->data_type(), DataType::INT32); + EXPECT_TRUE(current.get_field("extra")->nullable()); + EXPECT_FALSE(current.has_field("bad name")); + Doc doc = MakeDoc("new", 2); + ASSERT_TRUE(doc.set("extra", 7)); + std::vector docs{doc}; + ASSERT_NO_FATAL_FAILURE(ExpectWrite(collection_->insert(docs), 1)); + + auto altered = std::make_shared("extra", DataType::INT64, true); + status = collection_->alter_column("extra", "", altered); + ASSERT_TRUE(status.ok()) << status.message(); + altered->set_name("another bad name"); + altered->set_data_type(DataType::STRING); + current = collection_->schema().value(); + ASSERT_TRUE(current.has_field("extra")); + EXPECT_EQ(current.get_field("extra")->data_type(), DataType::INT64); + EXPECT_FALSE(current.has_field("another bad name")); + ASSERT_NO_FATAL_FAILURE(Reopen()); + ASSERT_NO_FATAL_FAILURE(ExpectValue("original", 1)); + ASSERT_NO_FATAL_FAILURE(ExpectValue("new", 2)); + auto fetched = collection_->fetch({"new"}); + ASSERT_TRUE(fetched.has_value()) << fetched.error().message(); + ASSERT_NE(fetched.value().at("new"), nullptr); + EXPECT_EQ(fetched.value().at("new")->get("extra"), 7); +} + +} // namespace +} // namespace zvec From 2231fcfb2f9682bf7a20b6adbb6eba8e9f02a7b3 Mon Sep 17 00:00:00 2001 From: Qinren Zhou Date: Wed, 16 Sep 2026 11:07:58 +0800 Subject: [PATCH 02/11] refactor --- src/db/collection.cc | 19 +- src/db/common/constants.h | 9 +- src/db/common/utils.cc | 46 +++++ src/db/common/utils.h | 8 +- src/db/index/common/doc.cc | 173 +++++++++--------- ...validation.cc => identifier_validation.cc} | 84 +++------ ...e_validation.h => identifier_validation.h} | 18 +- src/db/index/common/schema.cc | 28 ++- src/db/sqlengine/planner/fts_recall_node.cc | 4 +- src/db/sqlengine/planner/query_planner.cc | 4 +- .../sqlengine/planner/vector_recall_node.cc | 10 +- src/db/sqlengine/sqlengine_impl.cc | 6 +- ..._test.cc => identifier_validation_test.cc} | 92 +++++----- 13 files changed, 246 insertions(+), 255 deletions(-) rename src/db/index/common/{name_validation.cc => identifier_validation.cc} (59%) rename src/db/index/common/{name_validation.h => identifier_validation.h} (51%) rename tests/db/index/common/{name_validation_test.cc => identifier_validation_test.cc} (76%) diff --git a/src/db/collection.cc b/src/db/collection.cc index fab778299..65328c22f 100644 --- a/src/db/collection.cc +++ b/src/db/collection.cc @@ -42,11 +42,12 @@ #include "db/common/global_resource.h" #include "db/common/profiler.h" #include "db/common/typedef.h" +#include "db/common/utils.h" #include "db/doc_iterator_internal.h" #include "db/index/common/delete_store.h" #include "db/index/common/id_map.h" +#include "db/index/common/identifier_validation.h" #include "db/index/common/index_filter.h" -#include "db/index/common/name_validation.h" #include "db/index/common/type_helper.h" #include "db/index/common/version_manager.h" #include "db/index/segment/segment.h" @@ -1194,7 +1195,7 @@ Status CollectionImpl::validate(const std::string &column, field->data_type() > DataType::DOUBLE) { return Status::InvalidArgument( "Invalid schema: this operation requires a numeric field; field[", - FormatNameForError(field->name()), "] has type ", + format_name(field->name()), "] has type ", DataTypeCodeBook::AsString(field->data_type())); } return Status::OK(); @@ -1209,7 +1210,7 @@ Status CollectionImpl::validate(const std::string &column, if (schema_->has_field(schema->name())) { return Status::InvalidArgument("Invalid schema: field[", - FormatNameForError(schema->name()), + format_name(schema->name()), "] already exists"); } @@ -1227,7 +1228,7 @@ Status CollectionImpl::validate(const std::string &column, if (expression.empty() && !schema->nullable()) { return Status::InvalidArgument("Invalid schema: non-nullable field[", - FormatNameForError(schema->name()), + format_name(schema->name()), "] requires an expression when added"); } @@ -1241,8 +1242,7 @@ Status CollectionImpl::validate(const std::string &column, if (!schema_->has_field(column)) { return Status::InvalidArgument("Invalid schema: field[", - FormatNameForError(column), - "] not found"); + format_name(column), "] not found"); } if (!rename.empty() && schema) { @@ -1256,11 +1256,11 @@ Status CollectionImpl::validate(const std::string &column, if (!rename.empty()) { // rename case - s = ValidateFieldName(rename); + s = validate_field_name(rename); CHECK_RETURN_STATUS(s); if (schema_->has_field(rename)) { return Status::InvalidArgument("Invalid schema: field[", - FormatNameForError(rename), + format_name(rename), "] already exists"); } } else { @@ -1292,8 +1292,7 @@ Status CollectionImpl::validate(const std::string &column, case ColumnOp::DROP: { if (!schema_->has_field(column)) { return Status::InvalidArgument("Invalid schema: field[", - FormatNameForError(column), - "] not found"); + format_name(column), "] not found"); } if (schema_->fields().size() <= 1) { diff --git a/src/db/common/constants.h b/src/db/common/constants.h index 2e16cb7b4..2e023264a 100644 --- a/src/db/common/constants.h +++ b/src/db/common/constants.h @@ -31,12 +31,9 @@ const std::string GLOBAL_DOC_ID = "_zvec_g_doc_id_"; const std::string USER_ID = "_zvec_uid_"; -// Query result columns share a namespace with user fields. Keep these names -// available to validation without pulling in the Arrow query utilities. -namespace sqlengine { -inline constexpr const char *kFieldScore = "_zvec_score"; -inline constexpr const char *kFieldGroupId = "_zvec_group_id"; -} // namespace sqlengine +const std::string FIELD_SCORE = "_zvec_score"; + +const std::string FIELD_GROUP_ID = "_zvec_group_id"; const int kSparseMaxDimSize = 16384; diff --git a/src/db/common/utils.cc b/src/db/common/utils.cc index a1b83a114..047f01e94 100644 --- a/src/db/common/utils.cc +++ b/src/db/common/utils.cc @@ -13,6 +13,8 @@ // limitations under the License. #include "utils.h" +#include + namespace zvec { @@ -20,4 +22,48 @@ std::string indent(int level) { return std::string(level * 2, ' '); } +std::string format_name(std::string_view value) { + constexpr size_t kMaxPreviewBytes = 32; + constexpr char kHexDigits[] = "0123456789ABCDEF"; + auto length = std::min(value.size(), kMaxPreviewBytes); + std::string preview; + preview.reserve(length); + for (size_t i = 0; i < length; ++i) { + auto byte = static_cast(value[i]); + switch (byte) { + case '\0': + preview += "\\0"; + break; + case '\n': + preview += "\\n"; + break; + case '\r': + preview += "\\r"; + break; + case '\t': + preview += "\\t"; + break; + case '\\': + case '[': + case ']': + preview += '\\'; + preview += static_cast(byte); + break; + default: + if (byte >= 0x20 && byte <= 0x7E) { + preview += static_cast(byte); + } else { + preview += "\\x"; + preview += kHexDigits[byte >> 4]; + preview += kHexDigits[byte & 0x0F]; + } + break; + } + } + if (length < value.size()) { + preview += "..."; + } + return preview; +} + } // namespace zvec \ No newline at end of file diff --git a/src/db/common/utils.h b/src/db/common/utils.h index c238b928d..8baab3826 100644 --- a/src/db/common/utils.h +++ b/src/db/common/utils.h @@ -14,8 +14,14 @@ #pragma once #include +#include + namespace zvec { + std::string indent(int level); -} // namespace zvec \ No newline at end of file +// Format a name for clearer display. +std::string format_name(std::string_view value); + +} // namespace zvec diff --git a/src/db/index/common/doc.cc b/src/db/index/common/doc.cc index 4935414dd..99dceca30 100644 --- a/src/db/index/common/doc.cc +++ b/src/db/index/common/doc.cc @@ -17,13 +17,12 @@ #include #include #include -#include -#include #include #include #include #include "db/common/constants.h" -#include "db/index/common/name_validation.h" +#include "db/common/utils.h" +#include "db/index/common/identifier_validation.h" #include "db/index/common/type_helper.h" #if defined(__BYTE_ORDER__) && __BYTE_ORDER__ == __ORDER_BIG_ENDIAN__ @@ -151,7 +150,7 @@ class DocBufferReader { return remaining_; } - bool ReadBytes(void *destination, size_t size) { + bool read_bytes(void *destination, size_t size) { if (size > remaining_) return false; if (size != 0) { std::memcpy(destination, data_, size); @@ -162,11 +161,11 @@ class DocBufferReader { } template - bool ReadNative(T &value) { - return ReadBytes(&value, sizeof(T)); + bool read_native(T &value) { + return read_bytes(&value, sizeof(T)); } - bool ReadStringBytes(std::string &value, size_t size) { + bool read_string_bytes(std::string &value, size_t size) { if (size > remaining_) return false; value.assign(reinterpret_cast(data_), size); data_ += size; @@ -174,57 +173,57 @@ class DocBufferReader { return true; } - bool ReadValue(Doc::Value &value) { + bool read_value(Doc::Value &value) { uint8_t type; - if (!ReadNative(type)) return false; + if (!read_native(type)) return false; switch (type) { case TYPE_EMPTY: value = std::monostate{}; return true; case TYPE_BOOL: - return ReadAs(value); + return read_as(value); case TYPE_INT32: - return ReadAs(value); + return read_as(value); case TYPE_UINT32: - return ReadAs(value); + return read_as(value); case TYPE_INT64: - return ReadAs(value); + return read_as(value); case TYPE_UINT64: - return ReadAs(value); + return read_as(value); case TYPE_FLOAT: - return ReadAs(value); + return read_as(value); case TYPE_DOUBLE: - return ReadAs(value); + return read_as(value); case TYPE_STRING: - return ReadAs(value); + return read_as(value); case TYPE_VECTOR_BOOL: - return ReadAs>(value); + return read_as>(value); case TYPE_VECTOR_INT8: - return ReadAs>(value); + return read_as>(value); case TYPE_VECTOR_INT16: - return ReadAs>(value); + return read_as>(value); case TYPE_VECTOR_INT32: - return ReadAs>(value); + return read_as>(value); case TYPE_VECTOR_INT64: - return ReadAs>(value); + return read_as>(value); case TYPE_VECTOR_UINT32: - return ReadAs>(value); + return read_as>(value); case TYPE_VECTOR_UINT64: - return ReadAs>(value); + return read_as>(value); case TYPE_VECTOR_FLOAT16: - return ReadAs>(value); + return read_as>(value); case TYPE_VECTOR_FLOAT: - return ReadAs>(value); + return read_as>(value); case TYPE_VECTOR_DOUBLE: - return ReadAs>(value); + return read_as>(value); case TYPE_VECTOR_STRING: - return ReadAs>(value); + return read_as>(value); case TYPE_VECTOR_PAIR_INT_FLOAT: - return ReadAs, std::vector>>( + return read_as, std::vector>>( value); case TYPE_VECTOR_PAIR_INT_FLOAT16: - return ReadAs, std::vector>>( - value); + return read_as< + std::pair, std::vector>>(value); default: return false; } @@ -232,8 +231,8 @@ class DocBufferReader { private: template - bool ReadLittle(T &value) { - if (!ReadNative(value)) return false; + bool read_little_endian(T &value) { + if (!read_native(value)) return false; if (IS_BIG_ENDIAN) { auto *bytes = reinterpret_cast(&value); std::reverse(bytes, bytes + sizeof(T)); @@ -242,42 +241,42 @@ class DocBufferReader { } template - bool ReadAs(Doc::Value &out) { + bool read_as(Doc::Value &out) { T value; - if (!Read(value)) return false; + if (!read(value)) return false; out = std::move(value); return true; } template - bool Read(T &value) { - return ReadLittle(value); + bool read(T &value) { + return read_little_endian(value); } - bool Read(bool &value) { + bool read(bool &value) { static_assert(sizeof(bool) == sizeof(uint8_t)); uint8_t byte; - if (!ReadNative(byte) || byte > 1) return false; + if (!read_native(byte) || byte > 1) return false; value = byte != 0; return true; } - bool Read(std::string &value) { + bool read(std::string &value) { uint32_t size; - return ReadLittle(size) && ReadStringBytes(value, size); + return read_little_endian(size) && read_string_bytes(value, size); } template - bool Read(std::vector &values) { + bool read(std::vector &values) { uint32_t count; - if (!ReadLittle(count)) return false; + if (!read_little_endian(count)) return false; if constexpr (std::is_same_v) { // Each string contains at least its four-byte length prefix. if (count > remaining_ / sizeof(uint32_t)) return false; values.reserve(count); for (uint32_t i = 0; i < count; ++i) { std::string value; - if (!Read(value)) return false; + if (!read(value)) return false; values.push_back(std::move(value)); } } else if constexpr (std::is_same_v) { @@ -285,14 +284,14 @@ class DocBufferReader { values.reserve(count); for (uint32_t i = 0; i < count; ++i) { bool value; - if (!Read(value)) return false; + if (!read(value)) return false; values.push_back(value); } } else { // Division avoids overflow before checking the allocation/copy size. if (count > remaining_ / sizeof(T)) return false; values.resize(count); - if (!ReadBytes(values.data(), static_cast(count) * sizeof(T))) { + if (!read_bytes(values.data(), static_cast(count) * sizeof(T))) { return false; } if (IS_BIG_ENDIAN) { @@ -306,8 +305,8 @@ class DocBufferReader { } template - bool Read(std::pair, std::vector> &value) { - return Read(value.first) && Read(value.second); + bool read(std::pair, std::vector> &value) { + return read(value.first) && read(value.second); } const uint8_t *data_; @@ -610,12 +609,12 @@ Doc::Ptr Doc::deserialize(const uint8_t *data, size_t size) { uint32_t field_count; // The document header and field-name lengths retain their existing native // representation; value payloads use the existing little-endian encoding. - if (!reader.ReadNative(pk_length) || - !reader.ReadStringBytes(doc->pk_, pk_length) || - !reader.ReadNative(doc->score_) || !reader.ReadNative(doc->doc_id_) || - !reader.ReadNative(operation) || + if (!reader.read_native(pk_length) || + !reader.read_string_bytes(doc->pk_, pk_length) || + !reader.read_native(doc->score_) || !reader.read_native(doc->doc_id_) || + !reader.read_native(operation) || operation > static_cast(Operator::DELETE) || - !reader.ReadNative(field_count)) { + !reader.read_native(field_count)) { return nullptr; } doc->op_ = static_cast(operation); @@ -627,9 +626,9 @@ Doc::Ptr Doc::deserialize(const uint8_t *data, size_t size) { uint32_t name_length; std::string name; Value value; - if (!reader.ReadNative(name_length) || - !reader.ReadStringBytes(name, name_length) || - !reader.ReadValue(value) || + if (!reader.read_native(name_length) || + !reader.read_string_bytes(name, name_length) || + !reader.read_value(value) || !doc->fields_.emplace(std::move(name), std::move(value)).second) { return nullptr; } @@ -643,7 +642,7 @@ Status Doc::validate_and_sanitize(const CollectionSchema::Ptr &schema, return Status::InternalError("schema is null during doc validation"); } - auto id_status = ValidateDocumentId(pk_); + auto id_status = validate_document_id(pk_); if (!id_status.ok()) { return id_status; } @@ -652,8 +651,7 @@ Status Doc::validate_and_sanitize(const CollectionSchema::Ptr &schema, for (auto &[name, value] : fields_) { if (!schema->has_field(name)) { return Status::InvalidArgument( - "Invalid doc: doc[", FormatNameForError(pk_), "]: field[", - FormatNameForError(name), + "Invalid doc: doc[", format_name(pk_), "]: field[", format_name(name), "] does not exist in the collection schema"); } } @@ -666,17 +664,16 @@ Status Doc::validate_and_sanitize(const CollectionSchema::Ptr &schema, if (field_schema->nullable() || is_update) { continue; } - return Status::InvalidArgument( - "Invalid doc: doc[", FormatNameForError(pk_), "]: field[", - FormatNameForError(field_name), "] is required but not provided"); + return Status::InvalidArgument("Invalid doc: doc[", format_name(pk_), + "]: field[", format_name(field_name), + "] is required but not provided"); } else { if (std::holds_alternative(field_pair->second)) { if (field_schema->nullable()) { continue; } - return Status::InvalidArgument("Invalid doc: doc[", - FormatNameForError(pk_), "]: field[", - FormatNameForError(field_name), + return Status::InvalidArgument("Invalid doc: doc[", format_name(pk_), + "]: field[", format_name(field_name), "] is required but its value is null"); } } @@ -807,14 +804,14 @@ Status Doc::validate_and_sanitize(const CollectionSchema::Ptr &schema, field_value); if (sparse_values.size() != sparse_indices.size()) { return Status::InvalidArgument( - "Invalid doc: doc[", FormatNameForError(pk_), - "]: sparse vector field[", FormatNameForError(field_name), + "Invalid doc: doc[", format_name(pk_), + "]: sparse vector field[", format_name(field_name), "] has mismatched indices and values sizes"); } if (sparse_indices.size() > kSparseMaxDimSize) { return Status::InvalidArgument( - "Invalid doc: doc[", FormatNameForError(pk_), - "]: sparse vector field[", FormatNameForError(field_name), + "Invalid doc: doc[", format_name(pk_), + "]: sparse vector field[", format_name(field_name), "] exceeds the maximum number of sparse indices (", kSparseMaxDimSize, ")"); } @@ -822,8 +819,8 @@ Status Doc::validate_and_sanitize(const CollectionSchema::Ptr &schema, sparse_indices.size()); if (status == SparseIndicesStatus::kHasDuplicate) { return Status::InvalidArgument( - "Invalid doc: doc[", FormatNameForError(pk_), - "]: sparse vector field[", FormatNameForError(field_name), + "Invalid doc: doc[", format_name(pk_), + "]: sparse vector field[", format_name(field_name), "] contains duplicate indices"); } if (status == SparseIndicesStatus::kNeedSort) { @@ -832,8 +829,8 @@ Status Doc::validate_and_sanitize(const CollectionSchema::Ptr &schema, reinterpret_cast(sparse_values.data()), sparse_indices.size(), sizeof(float16_t))) { return Status::InvalidArgument( - "Invalid doc: doc[", FormatNameForError(pk_), - "]: sparse vector field[", FormatNameForError(field_name), + "Invalid doc: doc[", format_name(pk_), + "]: sparse vector field[", format_name(field_name), "] contains duplicate indices"); } } @@ -849,14 +846,14 @@ Status Doc::validate_and_sanitize(const CollectionSchema::Ptr &schema, field_value); if (sparse_values.size() != sparse_indices.size()) { return Status::InvalidArgument( - "Invalid doc: doc[", FormatNameForError(pk_), - "]: sparse vector field[", FormatNameForError(field_name), + "Invalid doc: doc[", format_name(pk_), + "]: sparse vector field[", format_name(field_name), "] has mismatched indices and values sizes"); } if (sparse_indices.size() > kSparseMaxDimSize) { return Status::InvalidArgument( - "Invalid doc: doc[", FormatNameForError(pk_), - "]: sparse vector field[", FormatNameForError(field_name), + "Invalid doc: doc[", format_name(pk_), + "]: sparse vector field[", format_name(field_name), "] exceeds the maximum number of sparse indices (", kSparseMaxDimSize, ")"); } @@ -864,8 +861,8 @@ Status Doc::validate_and_sanitize(const CollectionSchema::Ptr &schema, sparse_indices.size()); if (status == SparseIndicesStatus::kHasDuplicate) { return Status::InvalidArgument( - "Invalid doc: doc[", FormatNameForError(pk_), - "]: sparse vector field[", FormatNameForError(field_name), + "Invalid doc: doc[", format_name(pk_), + "]: sparse vector field[", format_name(field_name), "] contains duplicate indices"); } if (status == SparseIndicesStatus::kNeedSort) { @@ -874,8 +871,8 @@ Status Doc::validate_and_sanitize(const CollectionSchema::Ptr &schema, reinterpret_cast(sparse_values.data()), sparse_indices.size(), sizeof(float))) { return Status::InvalidArgument( - "Invalid doc: doc[", FormatNameForError(pk_), - "]: sparse vector field[", FormatNameForError(field_name), + "Invalid doc: doc[", format_name(pk_), + "]: sparse vector field[", format_name(field_name), "] contains duplicate indices"); } } @@ -883,24 +880,24 @@ Status Doc::validate_and_sanitize(const CollectionSchema::Ptr &schema, break; } default: - return Status::InvalidArgument( - "Invalid doc: doc[", FormatNameForError(pk_), "]: field[", - FormatNameForError(field_name), "] has unsupported data type"); + return Status::InvalidArgument("Invalid doc: doc[", format_name(pk_), + "]: field[", format_name(field_name), + "] has unsupported data type"); break; } if (!type_match) { return Status::InvalidArgument( - "Invalid doc: doc[", FormatNameForError(pk_), "]: field[", - FormatNameForError(field_name), "] type mismatch, expected ", + "Invalid doc: doc[", format_name(pk_), "]: field[", + format_name(field_name), "] type mismatch, expected ", DataTypeCodeBook::AsString(expected_type), " but got ", get_value_type_name(field_value, field_schema->is_vector_field())); } if (field_schema->is_dense_vector()) { if (value_dimension != field_schema->dimension()) { return Status::InvalidArgument( - "Invalid doc: doc[", FormatNameForError(pk_), "]: field[", - FormatNameForError(field_name), "] dimension mismatch, expected ", + "Invalid doc: doc[", format_name(pk_), "]: field[", + format_name(field_name), "] dimension mismatch, expected ", field_schema->dimension(), " but got ", value_dimension); } } diff --git a/src/db/index/common/name_validation.cc b/src/db/index/common/identifier_validation.cc similarity index 59% rename from src/db/index/common/name_validation.cc rename to src/db/index/common/identifier_validation.cc index 349cc4348..b0b109241 100644 --- a/src/db/index/common/name_validation.cc +++ b/src/db/index/common/identifier_validation.cc @@ -12,17 +12,20 @@ // See the License for the specific language governing permissions and // limitations under the License. -#include "name_validation.h" +#include "identifier_validation.h" #include #include #include -#include #include "db/common/constants.h" +#include "db/common/utils.h" + namespace zvec { + + namespace { -const char *ForbiddenCodepointReason(utf8proc_int32_t codepoint) { +const char *forbidden_codepoint_reason(utf8proc_int32_t codepoint) { if (codepoint == 0) { return "contains a null character"; } @@ -44,8 +47,8 @@ const char *ForbiddenCodepointReason(utf8proc_int32_t codepoint) { return nullptr; } -Status ValidateUtf8Name(std::string_view value, size_t max_bytes, - const char *prefix) { +Status validate_utf8_name(std::string_view value, size_t max_bytes, + const char *prefix) { if (value.empty()) { return Status::InvalidArgument(prefix, " must not be empty"); } @@ -64,7 +67,7 @@ Status ValidateUtf8Name(std::string_view value, size_t max_bytes, if (bytes <= 0) { return Status::InvalidArgument(prefix, " is not valid UTF-8"); } - if (const char *reason = ForbiddenCodepointReason(codepoint)) { + if (const char *reason = forbidden_codepoint_reason(codepoint)) { return Status::InvalidArgument(prefix, " ", reason); } position += static_cast(bytes); @@ -72,70 +75,26 @@ Status ValidateUtf8Name(std::string_view value, size_t max_bytes, return Status::OK(); } -bool IsReservedFieldName(std::string_view name) { +bool is_reserved_field_name(std::string_view name) { static const std::array reserved_names{ - LOCAL_ROW_ID, GLOBAL_DOC_ID, USER_ID, sqlengine::kFieldScore, - sqlengine::kFieldGroupId}; + LOCAL_ROW_ID, GLOBAL_DOC_ID, USER_ID, FIELD_SCORE, FIELD_GROUP_ID}; return std::find(reserved_names.begin(), reserved_names.end(), name) != reserved_names.end(); } } // namespace -std::string FormatNameForError(std::string_view name) { - constexpr size_t kMaxPreviewBytes = 32; - constexpr char kHexDigits[] = "0123456789ABCDEF"; - auto length = std::min(name.size(), kMaxPreviewBytes); - std::string preview; - preview.reserve(length); - for (size_t i = 0; i < length; ++i) { - auto byte = static_cast(name[i]); - switch (byte) { - case '\0': - preview += "\\0"; - break; - case '\n': - preview += "\\n"; - break; - case '\r': - preview += "\\r"; - break; - case '\t': - preview += "\\t"; - break; - case '\\': - case '[': - case ']': - preview += '\\'; - preview += static_cast(byte); - break; - default: - if (byte >= 0x20 && byte <= 0x7E) { - preview += static_cast(byte); - } else { - preview += "\\x"; - preview += kHexDigits[byte >> 4]; - preview += kHexDigits[byte & 0x0F]; - } - break; - } - } - if (length < name.size()) { - preview += "..."; - } - return preview; -} -Status ValidateDocumentId(std::string_view id) { - return ValidateUtf8Name(id, kMaxDocumentIdBytes, "Invalid doc: id"); +Status validate_document_id(std::string_view id) { + return validate_utf8_name(id, kMaxDocumentIdBytes, "Invalid doc: id"); } -Status ValidateCollectionName(std::string_view name) { - return ValidateUtf8Name(name, kMaxCollectionNameBytes, - "Invalid schema: collection name"); +Status validate_collection_name(std::string_view name) { + return validate_utf8_name(name, kMaxCollectionNameBytes, + "Invalid schema: collection name"); } -Status ValidateFieldName(std::string_view name) { +Status validate_field_name(std::string_view name) { if (name.empty()) { return Status::InvalidArgument( "Invalid schema: field name must not be empty"); @@ -151,18 +110,17 @@ Status ValidateFieldName(std::string_view name) { continue; } const char *reason = byte >= 0x80 ? "contains a non-ASCII character" - : ForbiddenCodepointReason(byte); + : forbidden_codepoint_reason(byte); if (!reason) { reason = byte == ' ' ? "contains a space" : "contains an unsupported character"; } return Status::InvalidArgument( - "Invalid schema: field[", FormatNameForError(name), "] ", reason, + "Invalid schema: field[", format_name(name), "] ", reason, "; use letters (A-Z, a-z), digits, underscores (_) or hyphens (-)"); } - if (IsReservedFieldName(name)) { - return Status::InvalidArgument("Invalid schema: field[", - FormatNameForError(name), + if (is_reserved_field_name(name)) { + return Status::InvalidArgument("Invalid schema: field[", format_name(name), "] is reserved; use a different name"); } return Status::OK(); diff --git a/src/db/index/common/name_validation.h b/src/db/index/common/identifier_validation.h similarity index 51% rename from src/db/index/common/name_validation.h rename to src/db/index/common/identifier_validation.h index 49adcf3b2..0342550d1 100644 --- a/src/db/index/common/name_validation.h +++ b/src/db/index/common/identifier_validation.h @@ -15,7 +15,6 @@ #pragma once #include -#include #include #include @@ -25,19 +24,8 @@ inline constexpr size_t kMaxDocumentIdBytes = 1024; inline constexpr size_t kMaxCollectionNameBytes = 256; inline constexpr size_t kMaxFieldNameBytes = 64; -// Validate new input without changing its bytes. Document IDs and collection -// names are nonempty UTF-8 strings without C0/C1 controls or line/paragraph -// separators. Other spaces, including strings consisting only of spaces, are -// allowed. These validators do not normalize, trim, or change case. -Status ValidateDocumentId(std::string_view id); -Status ValidateCollectionName(std::string_view name); - -// Field names retain the ASCII letters, digits, underscore, and hyphen set, -// excluding exact names used by storage and query execution. -Status ValidateFieldName(std::string_view name); - -// Bounded, escaped preview for errors. Never includes raw control characters -// or malformed UTF-8 bytes, even when the supplied name has not been validated. -std::string FormatNameForError(std::string_view name); +Status validate_document_id(std::string_view id); +Status validate_collection_name(std::string_view name); +Status validate_field_name(std::string_view name); } // namespace zvec diff --git a/src/db/index/common/schema.cc b/src/db/index/common/schema.cc index 7a2b302f9..5f3e91b78 100644 --- a/src/db/index/common/schema.cc +++ b/src/db/index/common/schema.cc @@ -25,7 +25,7 @@ #include "db/common/utils.h" #include "db/index/column/fts_column/fts_types.h" #include "db/index/column/fts_column/tokenizer/tokenizer_factory.h" -#include "db/index/common/name_validation.h" +#include "db/index/common/identifier_validation.h" #include "db/index/common/type_helper.h" namespace zvec { @@ -81,7 +81,7 @@ static Status validate_fts_index_params(const FieldSchema &field) { } Status FieldSchema::validate() const { - auto name_status = ValidateFieldName(name_); + auto name_status = validate_field_name(name_); CHECK_RETURN_STATUS(name_status); if (data_type_ == DataType::UNDEFINED) { @@ -405,7 +405,7 @@ std::string FieldSchema::to_string_formatted(int indent_level) const { } Status CollectionSchema::validate() const { - auto name_status = ValidateCollectionName(name_); + auto name_status = validate_collection_name(name_); CHECK_RETURN_STATUS(name_status); std::unordered_set names; for (const auto &field : fields_) { @@ -415,13 +415,13 @@ Status CollectionSchema::validate() const { } if (!names.insert(field->name()).second) { return Status::InvalidArgument("Invalid schema: duplicate field name [", - FormatNameForError(field->name()), + format_name(field->name()), "]; field names must be unique"); } } if (forward_fields().size() > kMaxScalarFieldSize) { return Status::InvalidArgument( - "Invalid schema: collection[", FormatNameForError(name_), + "Invalid schema: collection[", format_name(name_), "]'s field size must <= ", kMaxScalarFieldSize); } if (max_doc_count_per_segment_ < MAX_DOC_COUNT_PER_SEGMENT_MIN_THRESHOLD) { @@ -431,13 +431,12 @@ Status CollectionSchema::validate() const { } if (fields_.empty()) { return Status::InvalidArgument("Invalid schema: collection[", - FormatNameForError(name_), - "] has no fields"); + format_name(name_), "] has no fields"); } auto v_fields = vector_fields(); if (v_fields.size() > kMaxVectorFieldSize) { return Status::InvalidArgument( - "Invalid schema: collection[", FormatNameForError(name_), + "Invalid schema: collection[", format_name(name_), "]'s vector field size must <= ", kMaxVectorFieldSize); } for (auto &field : fields_) { @@ -491,8 +490,7 @@ Status CollectionSchema::add_field(FieldSchema::Ptr column_schema) { } // Check if field already exists if (has_field(column_schema->name())) { - return Status::AlreadyExists("field[", - FormatNameForError(column_schema->name()), + return Status::AlreadyExists("field[", format_name(column_schema->name()), "] already exists in schema"); } @@ -518,7 +516,7 @@ Status CollectionSchema::alter_field( } // Check if field exists if (!has_field(column_name)) { - return Status::NotFound("field[", FormatNameForError(column_name), + return Status::NotFound("field[", format_name(column_name), "] not found in schema"); } @@ -526,7 +524,7 @@ Status CollectionSchema::alter_field( // If renaming to an existing field name (and it's not the same field) if (new_column_name != column_name && has_field(new_column_name)) { - return Status::AlreadyExists("field[", FormatNameForError(new_column_name), + return Status::AlreadyExists("field[", format_name(new_column_name), "] already exists in schema"); } @@ -550,7 +548,7 @@ Status CollectionSchema::alter_field( Status CollectionSchema::drop_field(const std::string &column_name) { // Check if field exists if (!has_field(column_name)) { - return Status::NotFound("field[", FormatNameForError(column_name), + return Status::NotFound("field[", format_name(column_name), "] not found in schema"); } @@ -734,7 +732,7 @@ Status CollectionSchema::add_index(const std::string &column, if (field) { field->set_index_params(index_params); } else { - return Status::NotFound("field[", FormatNameForError(column), + return Status::NotFound("field[", format_name(column), "] not found in schema"); } @@ -751,7 +749,7 @@ Status CollectionSchema::drop_index(const std::string &column) { field->set_index_params(nullptr); } } else { - return Status::NotFound("field[", FormatNameForError(column), + return Status::NotFound("field[", format_name(column), "] not found in schema"); } diff --git a/src/db/sqlengine/planner/fts_recall_node.cc b/src/db/sqlengine/planner/fts_recall_node.cc index 45313d9e0..6a94ebeb3 100644 --- a/src/db/sqlengine/planner/fts_recall_node.cc +++ b/src/db/sqlengine/planner/fts_recall_node.cc @@ -32,7 +32,7 @@ FtsRecallNode::FtsRecallNode(Segment::Ptr segment, QueryInfo::Ptr query_info, auto table = segment_->fetch(fetched_columns_, std::vector{}); // Append BM25 score column so downstream fill_doc_score() surfaces it to // the Python Doc.score, matching the vector-recall path. - schema_ = Util::append_field(*table->schema(), kFieldScore, arrow::float32()); + schema_ = Util::append_field(*table->schema(), FIELD_SCORE, arrow::float32()); } arrow::AsyncGenerator> FtsRecallNode::gen() { @@ -94,7 +94,7 @@ arrow::AsyncGenerator> FtsRecallNode::gen() { } auto record_batch = std::move(batch.ValueUnsafe()); auto with_score = - record_batch->AddColumn(record_batch->num_columns(), kFieldScore, + record_batch->AddColumn(record_batch->num_columns(), FIELD_SCORE, score_array.MoveValueUnsafe()); if (!with_score.ok()) { return arrow::Future>::MakeFinished( diff --git a/src/db/sqlengine/planner/query_planner.cc b/src/db/sqlengine/planner/query_planner.cc index d539b76ae..4a6aee766 100644 --- a/src/db/sqlengine/planner/query_planner.cc +++ b/src/db/sqlengine/planner/query_planner.cc @@ -446,7 +446,7 @@ Result QueryPlanner::make_physical_plan( "order_by", {std::move(node)}, ac::OrderByNodeOptions{cp::Ordering{{cp::SortKey{ - kFieldScore, vector_is_reverse ? cp::SortOrder::Descending + FIELD_SCORE, vector_is_reverse ? cp::SortOrder::Descending : cp::SortOrder::Ascending}}}}}; } else if (has_fts) { // FTS uses BM25 where higher score = more relevant. Per-segment results @@ -455,7 +455,7 @@ Result QueryPlanner::make_physical_plan( node = ac::Declaration{"order_by", {std::move(node)}, ac::OrderByNodeOptions{cp::Ordering{{cp::SortKey{ - kFieldScore, cp::SortOrder::Descending}}}}}; + FIELD_SCORE, cp::SortOrder::Descending}}}}}; } // group by need to collect all docs diff --git a/src/db/sqlengine/planner/vector_recall_node.cc b/src/db/sqlengine/planner/vector_recall_node.cc index f58d02c1b..58e4e8d43 100644 --- a/src/db/sqlengine/planner/vector_recall_node.cc +++ b/src/db/sqlengine/planner/vector_recall_node.cc @@ -47,7 +47,7 @@ VectorRecallNode::VectorRecallNode(Segment::Ptr segment, : query_info_->get_selected_scalar_field_names()) { auto table = segment_->fetch(fetched_columns_, std::vector{}); schema_ = table->schema(); - schema_ = Util::append_field(*schema_, kFieldScore, arrow::float32()); + schema_ = Util::append_field(*schema_, FIELD_SCORE, arrow::float32()); if (query_info_->is_include_vector()) { for (auto &field : query_info_->selected_vector_fields()) { if (field.field_schema_ptr->is_dense_vector()) { @@ -60,14 +60,14 @@ VectorRecallNode::VectorRecallNode(Segment::Ptr segment, } } if (query_info_->group_by()) { - schema_ = Util::append_field(*schema_, kFieldGroupId, arrow::utf8()); + schema_ = Util::append_field(*schema_, FIELD_GROUP_ID, arrow::utf8()); } } arrow::AsyncGenerator> VectorRecallNode::gen() { auto state_ptr = std::make_shared(shared_from_this()); return [state_ptr = std::move(state_ptr)]() mutable - -> arrow::Future> { + -> arrow::Future> { auto &state = *state_ptr; if (!state.iter_) { auto vector_ret = state.self_->prepare(); @@ -245,7 +245,7 @@ VectorRecallNode::State::collect_batch() { auto record_batch = std::move(batch.ValueUnsafe()); ARROW_ASSIGN_OR_RAISE( record_batch, - record_batch->AddColumn(record_batch->num_columns(), kFieldScore, + record_batch->AddColumn(record_batch->num_columns(), FIELD_SCORE, score_array.MoveValueUnsafe())); if (self_->query_info_->is_include_vector()) { @@ -277,7 +277,7 @@ VectorRecallNode::State::collect_batch() { } ARROW_ASSIGN_OR_RAISE( record_batch, - record_batch->AddColumn(record_batch->num_columns(), kFieldGroupId, + record_batch->AddColumn(record_batch->num_columns(), FIELD_GROUP_ID, group_id_array.MoveValueUnsafe())); } diff --git a/src/db/sqlengine/sqlengine_impl.cc b/src/db/sqlengine/sqlengine_impl.cc index e74fbdc25..07530ca3f 100644 --- a/src/db/sqlengine/sqlengine_impl.cc +++ b/src/db/sqlengine/sqlengine_impl.cc @@ -521,7 +521,7 @@ Status record_batch_to_doc_list( doc_id_array != nullptr) { fill_doc_id(doc_id_array, doc_it); } - if (auto score_array = record_batch.GetColumnByName(kFieldScore); + if (auto score_array = record_batch.GetColumnByName(FIELD_SCORE); score_array != nullptr) { fill_doc_score(score_array, doc_it); } @@ -597,10 +597,10 @@ Result SQLEngineImpl::fill_group_by_result( if (!status.ok()) { return tl::make_unexpected(status); } - auto group_id_array = record_batch->GetColumnByName(kFieldGroupId); + auto group_id_array = record_batch->GetColumnByName(FIELD_GROUP_ID); if (!group_id_array) { return tl::make_unexpected(Status::InternalError( - "Column not found in record batch: [", kFieldGroupId, "]")); + "Column not found in record batch: [", FIELD_GROUP_ID, "]")); } arrow::StringArray *typed_arr = static_cast(group_id_array.get()); diff --git a/tests/db/index/common/name_validation_test.cc b/tests/db/index/common/identifier_validation_test.cc similarity index 76% rename from tests/db/index/common/name_validation_test.cc rename to tests/db/index/common/identifier_validation_test.cc index def1424e6..c3cc36864 100644 --- a/tests/db/index/common/name_validation_test.cc +++ b/tests/db/index/common/identifier_validation_test.cc @@ -12,12 +12,13 @@ // See the License for the specific language governing permissions and // limitations under the License. -#include "db/index/common/name_validation.h" +#include "db/index/common/identifier_validation.h" #include #include #include #include #include +#include "db/common/utils.h" namespace zvec { namespace { @@ -29,8 +30,8 @@ struct Utf8NameValidator { }; const std::array kUtf8NameValidators{{ - {ValidateDocumentId, kMaxDocumentIdBytes, "Invalid doc: id"}, - {ValidateCollectionName, kMaxCollectionNameBytes, + {validate_document_id, kMaxDocumentIdBytes, "Invalid doc: id"}, + {validate_collection_name, kMaxCollectionNameBytes, "Invalid schema: collection name"}, }}; @@ -49,7 +50,7 @@ std::string Repeat(std::string_view text, size_t count) { return result; } -TEST(NameValidationTest, AcceptsUnicodePunctuationAndSpaces) { +TEST(IdentifierValidationTest, AcceptsUnicodePunctuationAndSpaces) { const std::vector values{ "a", "A_b-9", @@ -76,16 +77,16 @@ TEST(NameValidationTest, AcceptsUnicodePunctuationAndSpaces) { } } -TEST(NameValidationTest, RejectsEmptyNames) { +TEST(IdentifierValidationTest, RejectsEmptyNames) { for (const auto &validator : kUtf8NameValidators) { ExpectInvalid(validator.validate(std::string_view{}), std::string(validator.prefix) + " must not be empty"); } - ExpectInvalid(ValidateFieldName(""), + ExpectInvalid(validate_field_name(""), "Invalid schema: field name must not be empty"); } -TEST(NameValidationTest, MeasuresLimitsInUtf8Bytes) { +TEST(IdentifierValidationTest, MeasuresLimitsInUtf8Bytes) { for (const auto &validator : kUtf8NameValidators) { SCOPED_TRACE(validator.prefix); const auto max_bytes = validator.max_bytes; @@ -107,7 +108,7 @@ TEST(NameValidationTest, MeasuresLimitsInUtf8Bytes) { } } -TEST(NameValidationTest, RejectsMalformedUtf8) { +TEST(IdentifierValidationTest, RejectsMalformedUtf8) { const std::vector malformed{ "\x80", // Isolated continuation byte. "\xBF", @@ -141,7 +142,7 @@ TEST(NameValidationTest, RejectsMalformedUtf8) { } } -TEST(NameValidationTest, HonorsStringViewLengthAndEmbeddedNulls) { +TEST(IdentifierValidationTest, HonorsStringViewLengthAndEmbeddedNulls) { const std::string backing = std::string(u8"中文") + "\xFF"; for (const auto &validator : kUtf8NameValidators) { SCOPED_TRACE(validator.prefix); @@ -153,7 +154,7 @@ TEST(NameValidationTest, HonorsStringViewLengthAndEmbeddedNulls) { } } -TEST(NameValidationTest, RejectsEveryC0AndC1ControlByCodepoint) { +TEST(IdentifierValidationTest, RejectsEveryC0AndC1ControlByCodepoint) { for (const auto &validator : kUtf8NameValidators) { SCOPED_TRACE(validator.prefix); for (unsigned int codepoint = 0; codepoint <= 0x9F; ++codepoint) { @@ -183,7 +184,7 @@ TEST(NameValidationTest, RejectsEveryC0AndC1ControlByCodepoint) { } } -TEST(NameValidationTest, DistinguishesUnicodeLineAndParagraphSeparators) { +TEST(IdentifierValidationTest, DistinguishesUnicodeLineAndParagraphSeparators) { for (const auto &validator : kUtf8NameValidators) { ExpectInvalid(validator.validate(u8"a\u2028b"), std::string(validator.prefix) + " contains a line separator"); @@ -193,87 +194,88 @@ TEST(NameValidationTest, DistinguishesUnicodeLineAndParagraphSeparators) { } } -TEST(NameValidationTest, RetainsTheFieldAsciiCharacterSet) { +TEST(IdentifierValidationTest, RetainsTheFieldAsciiCharacterSet) { for (const std::string name : {"a", "Z", "0", "_", "-", "a_b-c1", "123_test", "_zvec_custom"}) { - EXPECT_TRUE(ValidateFieldName(name).ok()); + EXPECT_TRUE(validate_field_name(name).ok()); } - EXPECT_TRUE(ValidateFieldName("ABCDEFGHIJKLMNOPQRSTUVWXYZ" - "abcdefghijklmnopqrstuvwxyz0123456789_-") + EXPECT_TRUE(validate_field_name("ABCDEFGHIJKLMNOPQRSTUVWXYZ" + "abcdefghijklmnopqrstuvwxyz0123456789_-") .ok()); - EXPECT_TRUE(ValidateFieldName(std::string(kMaxFieldNameBytes, 'a')).ok()); - ExpectInvalid(ValidateFieldName(std::string(kMaxFieldNameBytes + 1, 'a')), + EXPECT_TRUE(validate_field_name(std::string(kMaxFieldNameBytes, 'a')).ok()); + ExpectInvalid(validate_field_name(std::string(kMaxFieldNameBytes + 1, 'a')), "Invalid schema: field name exceeds 64 bytes (got 65)"); - ExpectInvalid(ValidateFieldName(std::string(10000, 'a')), + ExpectInvalid(validate_field_name(std::string(10000, 'a')), "Invalid schema: field name exceeds 64 bytes (got 10000)"); } -TEST(NameValidationTest, RejectsExactInternalFieldNames) { +TEST(IdentifierValidationTest, RejectsExactInternalFieldNames) { for (const std::string name : {"_zvec_row_id_", "_zvec_g_doc_id_", "_zvec_uid_", "_zvec_score", "_zvec_group_id"}) { SCOPED_TRACE(name); - ExpectInvalid(ValidateFieldName(name), + ExpectInvalid(validate_field_name(name), "Invalid schema: field[" + name + "] is reserved; use a different name"); // The restriction is an exact match, not a new prefix or case policy. - EXPECT_TRUE(ValidateFieldName(name + "_custom").ok()); - EXPECT_TRUE(ValidateDocumentId(name).ok()); - EXPECT_TRUE(ValidateCollectionName(name).ok()); + EXPECT_TRUE(validate_field_name(name + "_custom").ok()); + EXPECT_TRUE(validate_document_id(name).ok()); + EXPECT_TRUE(validate_collection_name(name).ok()); } - EXPECT_TRUE(ValidateFieldName("_ZVEC_UID_").ok()); + EXPECT_TRUE(validate_field_name("_ZVEC_UID_").ok()); for (const std::string name : {"_zvec_vector", "_zvec_sindices", "_zvec_svalues", "_zvec_is_valid"}) { - EXPECT_TRUE(ValidateFieldName(name).ok()); + EXPECT_TRUE(validate_field_name(name).ok()); } } -TEST(NameValidationTest, SharedErrorPreviewIsEscapedAndBounded) { - EXPECT_EQ(FormatNameForError(""), ""); - EXPECT_EQ(FormatNameForError(std::string("a\0\n\r\t[]\\", 8)), +TEST(IdentifierValidationTest, SharedErrorPreviewIsEscapedAndBounded) { + EXPECT_EQ(format_name(""), ""); + EXPECT_EQ(format_name(std::string("a\0\n\r\t[]\\", 8)), "a\\0\\n\\r\\t\\[\\]\\\\"); - EXPECT_EQ(FormatNameForError(u8"中"), "\\xE4\\xB8\\xAD"); - EXPECT_EQ(FormatNameForError(std::string(10000, '\xff')), + EXPECT_EQ(format_name(u8"中"), "\\xE4\\xB8\\xAD"); + EXPECT_EQ(format_name(std::string(10000, '\xff')), Repeat("\\xFF", 32) + "..."); - EXPECT_EQ(FormatNameForError(std::string(10000, 'x')), - std::string(32, 'x') + "..."); + EXPECT_EQ(format_name(std::string(10000, 'x')), std::string(32, 'x') + "..."); } -TEST(NameValidationTest, DescribesInvalidFieldCharactersWithSafePreviews) { +TEST(IdentifierValidationTest, + DescribesInvalidFieldCharactersWithSafePreviews) { const std::string rule = "; use letters (A-Z, a-z), digits, underscores (_) or hyphens (-)"; - ExpectInvalid(ValidateFieldName("user name"), + ExpectInvalid(validate_field_name("user name"), "Invalid schema: field[user name] contains a space" + rule); ExpectInvalid( - ValidateFieldName("a.b"), + validate_field_name("a.b"), "Invalid schema: field[a.b] contains an unsupported character" + rule); - ExpectInvalid(ValidateFieldName(u8"中"), + ExpectInvalid(validate_field_name(u8"中"), "Invalid schema: field[\\xE4\\xB8\\xAD] contains a non-ASCII " "character" + rule); ExpectInvalid( - ValidateFieldName("\x80"), + validate_field_name("\x80"), "Invalid schema: field[\\x80] contains a non-ASCII character" + rule); ExpectInvalid( - ValidateFieldName(std::string("a\0b", 3)), + validate_field_name(std::string("a\0b", 3)), "Invalid schema: field[a\\0b] contains a null character" + rule); - ExpectInvalid(ValidateFieldName("a\nb"), + ExpectInvalid(validate_field_name("a\nb"), "Invalid schema: field[a\\nb] contains a newline" + rule); - ExpectInvalid(ValidateFieldName("a\tb"), + ExpectInvalid(validate_field_name("a\tb"), "Invalid schema: field[a\\tb] contains a tab" + rule); ExpectInvalid( - ValidateFieldName("a\x1B" - "b"), + validate_field_name("a\x1B" + "b"), "Invalid schema: field[a\\x1Bb] contains a control character" + rule); - ExpectInvalid(ValidateFieldName("][\\\n"), + ExpectInvalid(validate_field_name("][\\\n"), "Invalid schema: field[\\]\\[\\\\\\n] contains an unsupported " "character" + rule); } -TEST(NameValidationTest, BoundsInvalidFieldPreviewsAndNeverEchoesRawBytes) { +TEST(IdentifierValidationTest, + BoundsInvalidFieldPreviewsAndNeverEchoesRawBytes) { const std::string name(64, '\xFF'); - const auto status = ValidateFieldName(name); + const auto status = validate_field_name(name); ExpectInvalid(status, "Invalid schema: field[" + Repeat("\\xFF", 32) + "...] contains a non-ASCII character; use letters " From 3c3a711dc5869468eb27ed2b0febcb5a05b0ecb3 Mon Sep 17 00:00:00 2001 From: Qinren Zhou Date: Wed, 16 Sep 2026 13:52:43 +0800 Subject: [PATCH 03/11] fix --- python/tests/test_name_validation.py | 41 +++++++++++- python/zvec/model/_validation.py | 4 +- src/db/collection.cc | 41 +++++------- src/db/index/common/identifier_validation.cc | 11 ++-- src/db/sqlengine/common/util.h | 1 - src/include/zvec/db/schema.h | 10 +-- tests/db/index/common/doc_test.cc | 19 ++++++ .../common/identifier_validation_test.cc | 66 ++++++++++++------- tests/db/index/common/schema_test.cc | 17 +++-- tests/db/relaxed_validation_test.cc | 28 +++++++- 10 files changed, 168 insertions(+), 70 deletions(-) diff --git a/python/tests/test_name_validation.py b/python/tests/test_name_validation.py index d0ad43a80..97329ef10 100644 --- a/python/tests/test_name_validation.py +++ b/python/tests/test_name_validation.py @@ -12,6 +12,8 @@ # See the License for the specific language governing permissions and # limitations under the License. +import re + import pytest import zvec @@ -96,8 +98,16 @@ def test_invalid_id_rejects_batch_before_writing(collection, operation, doc_id, message = str(exc_info.value) assert message.startswith("Invalid doc:") assert reason in message - assert "document at index 1" in message + if doc_id == "\ud800": + # Conversion fails in Python before native validation runs. + assert "document at index 1" in message + else: + assert "document at index" not in message assert "offset" not in message + if doc_id: + assert "id[" in message + assert "\n" not in message + assert "\0" not in message fetched = collection.fetch("valid") if operation == "update": assert fetched["valid"].field("text") == "before" @@ -149,6 +159,24 @@ def test_invalid_collection_names_report_the_reason(tmp_path, name, reason): assert "collection name" in message assert reason in message assert "offset" not in message + if name: + assert "collection name[" in message + assert "\n" not in message + assert "\0" not in message + + +@pytest.mark.parametrize( + "doc_id,preview", + [ + ("order\n[123]", r"order\n\[123\]"), + (b"\xff", r"\xFF"), + ("x" * 1025, "x" * 32 + "..."), + ], +) +def test_id_error_includes_safe_preview(collection, doc_id, preview): + with pytest.raises(ValueError) as exc_info: + collection.insert(zvec.Doc(doc_id, fields={"text": "value"})) + assert f"Invalid doc: id[{preview}]" in str(exc_info.value) def test_long_field_name_and_rejected_rename_preserve_data(tmp_path): @@ -175,7 +203,11 @@ def test_long_field_name_and_rejected_rename_preserve_data(tmp_path): def test_surrogate_id_has_a_readable_encoding_error(collection, operation): with pytest.raises( ValueError, - match=r"^Invalid doc: id is not valid UTF-8 \(document at index 0\)$", + match="^" + + re.escape( + r"Invalid doc: id['\ud800'] is not valid UTF-8 (document at index 0)" + ) + + "$", ): getattr(collection, operation)(zvec.Doc("\ud800", fields={"text": "value"})) assert collection.stats.doc_count == 0 @@ -183,13 +215,16 @@ def test_surrogate_id_has_a_readable_encoding_error(collection, operation): @pytest.mark.parametrize("kind", ["collection", "field", "vector"]) def test_surrogate_schema_name_has_a_readable_encoding_error(kind): - with pytest.raises(ValueError, match="^Invalid schema: .* is not valid UTF-8$"): + with pytest.raises( + ValueError, match="^Invalid schema: .* is not valid UTF-8$" + ) as exc_info: if kind == "collection": zvec.CollectionSchema("\ud800") elif kind == "field": zvec.FieldSchema("\ud800", zvec.DataType.INT32) else: zvec.VectorSchema("\ud800", zvec.DataType.VECTOR_FP32, dimension=2) + assert r"['\ud800']" in str(exc_info.value) @pytest.mark.parametrize("invalid", [0, False, [], {}, b""]) diff --git a/python/zvec/model/_validation.py b/python/zvec/model/_validation.py index 7fb461a58..1d2e399de 100644 --- a/python/zvec/model/_validation.py +++ b/python/zvec/model/_validation.py @@ -20,7 +20,9 @@ def explain_utf8_conversion_error(value: object, context: str) -> None: try: value.encode("utf-8") except UnicodeEncodeError: - raise ValueError(f"{context} is not valid UTF-8") from None + raise ValueError( + f"{context}[{format_name_for_error(value)}] is not valid UTF-8" + ) from None def format_name_for_error(name: str) -> str: diff --git a/src/db/collection.cc b/src/db/collection.cc index 65328c22f..3b3ca6483 100644 --- a/src/db/collection.cc +++ b/src/db/collection.cc @@ -1195,7 +1195,7 @@ Status CollectionImpl::validate(const std::string &column, field->data_type() > DataType::DOUBLE) { return Status::InvalidArgument( "Invalid schema: this operation requires a numeric field; field[", - format_name(field->name()), "] has type ", + field->name(), "] has type ", DataTypeCodeBook::AsString(field->data_type())); } return Status::OK(); @@ -1209,8 +1209,7 @@ Status CollectionImpl::validate(const std::string &column, } if (schema_->has_field(schema->name())) { - return Status::InvalidArgument("Invalid schema: field[", - format_name(schema->name()), + return Status::InvalidArgument("Invalid schema: field[", schema->name(), "] already exists"); } @@ -1228,7 +1227,7 @@ Status CollectionImpl::validate(const std::string &column, if (expression.empty() && !schema->nullable()) { return Status::InvalidArgument("Invalid schema: non-nullable field[", - format_name(schema->name()), + schema->name(), "] requires an expression when added"); } @@ -1259,8 +1258,7 @@ Status CollectionImpl::validate(const std::string &column, s = validate_field_name(rename); CHECK_RETURN_STATUS(s); if (schema_->has_field(rename)) { - return Status::InvalidArgument("Invalid schema: field[", - format_name(rename), + return Status::InvalidArgument("Invalid schema: field[", rename, "] already exists"); } } else { @@ -1325,18 +1323,16 @@ Status CollectionImpl::add_column(const FieldSchema::Ptr &column_schema, CHECK_DESTROY_RETURN_STATUS(destroyed_, false); CHECK_CLOSED_RETURN_STATUS(closed_, false); - // Keep caller-owned mutable objects out of the published schema. Validate - // and execute the operation using the same independent snapshot. - auto field_snapshot = + auto field_copy = column_schema ? std::make_shared(*column_schema) : nullptr; - auto s = validate("", field_snapshot, expression, "", ColumnOp::ADD); + auto s = validate("", field_copy, expression, "", ColumnOp::ADD); CHECK_RETURN_STATUS(s); // forbidden writing until index is ready std::lock_guard write_lock(write_mtx_); auto new_schema = std::make_shared(*schema_); - s = new_schema->add_field(field_snapshot); + s = new_schema->add_field(field_copy); CHECK_RETURN_STATUS(s); if (writing_segment_->has_record()) { @@ -1347,7 +1343,7 @@ Status CollectionImpl::add_column(const FieldSchema::Ptr &column_schema, Version new_version = version_manager_->get_current_version(); // add column on segment manager - s = segment_manager_->add_column(field_snapshot, expression, + s = segment_manager_->add_column(field_copy, expression, options.concurrency_); CHECK_RETURN_STATUS(s); @@ -1482,23 +1478,23 @@ Status CollectionImpl::alter_column(const std::string &column_name, CHECK_DESTROY_RETURN_STATUS(destroyed_, false); CHECK_CLOSED_RETURN_STATUS(closed_, false); - auto new_field_schema = + auto field_copy = new_column_schema ? std::make_shared(*new_column_schema) : nullptr; - auto s = validate(column_name, new_field_schema, "", rename, ColumnOp::ALTER); + auto s = validate(column_name, field_copy, "", rename, ColumnOp::ALTER); CHECK_RETURN_STATUS(s); // forbidden writing until index is ready std::lock_guard write_lock(write_mtx_); if (!rename.empty()) { - new_field_schema = + field_copy = std::make_shared(*schema_->get_field(column_name)); - new_field_schema->set_name(rename); + field_copy->set_name(rename); } auto new_schema = std::make_shared(*schema_); - s = new_schema->alter_field(column_name, new_field_schema); + s = new_schema->alter_field(column_name, field_copy); CHECK_RETURN_STATUS(s); if (writing_segment_->has_record()) { @@ -1509,7 +1505,7 @@ Status CollectionImpl::alter_column(const std::string &column_name, Version new_version = version_manager_->get_current_version(); // alter column on segment manager - s = segment_manager_->alter_column(column_name, new_field_schema, + s = segment_manager_->alter_column(column_name, field_copy, options.concurrency_); CHECK_RETURN_STATUS(s); @@ -1621,14 +1617,9 @@ Result CollectionImpl::write_impl(std::vector &docs, CHECK_DESTROY_RETURN_STATUS_EXPECTED(destroyed_, false); CHECK_CLOSED_RETURN_STATUS_EXPECTED(closed_, false); - for (size_t i = 0; i < docs.size(); ++i) { - auto &doc = docs[i]; + for (auto &&doc : docs) { auto s = doc.validate_and_sanitize(schema_, mode == WriteMode::UPDATE); - if (!s.ok()) { - return tl::make_unexpected(Status( - s.code(), - s.message() + " (document at index " + std::to_string(i) + ")")); - } + CHECK_RETURN_STATUS_EXPECTED(s); } // TODO: The granularity of the write_lock is too coarse. diff --git a/src/db/index/common/identifier_validation.cc b/src/db/index/common/identifier_validation.cc index b0b109241..1720c6278 100644 --- a/src/db/index/common/identifier_validation.cc +++ b/src/db/index/common/identifier_validation.cc @@ -53,8 +53,9 @@ Status validate_utf8_name(std::string_view value, size_t max_bytes, return Status::InvalidArgument(prefix, " must not be empty"); } if (value.size() > max_bytes) { - return Status::InvalidArgument(prefix, " exceeds ", max_bytes, - " bytes (got ", value.size(), ")"); + return Status::InvalidArgument(prefix, "[", format_name(value), + "] exceeds ", max_bytes, " bytes (got ", + value.size(), ")"); } const auto *data = reinterpret_cast(value.data()); @@ -65,10 +66,12 @@ Status validate_utf8_name(std::string_view value, size_t max_bytes, data + position, static_cast(value.size() - position), &codepoint); if (bytes <= 0) { - return Status::InvalidArgument(prefix, " is not valid UTF-8"); + return Status::InvalidArgument(prefix, "[", format_name(value), + "] is not valid UTF-8"); } if (const char *reason = forbidden_codepoint_reason(codepoint)) { - return Status::InvalidArgument(prefix, " ", reason); + return Status::InvalidArgument(prefix, "[", format_name(value), "] ", + reason); } position += static_cast(bytes); } diff --git a/src/db/sqlengine/common/util.h b/src/db/sqlengine/common/util.h index 0f2012ed8..18f3644be 100644 --- a/src/db/sqlengine/common/util.h +++ b/src/db/sqlengine/common/util.h @@ -16,7 +16,6 @@ #include #include #include -#include "db/common/constants.h" namespace zvec::sqlengine { diff --git a/src/include/zvec/db/schema.h b/src/include/zvec/db/schema.h index 856d5aa59..0b03a1c66 100644 --- a/src/include/zvec/db/schema.h +++ b/src/include/zvec/db/schema.h @@ -15,7 +15,6 @@ #include #include -#include #include #include #include @@ -408,15 +407,10 @@ class ZVEC_API CollectionSchema { private: void copy_fields(const FieldSchemaPtrList &fields) { - // Constructors cannot return a Status. Reject missing field objects here - // instead of dereferencing them or silently omitting part of the schema. - for (const auto &field : fields) { + for (auto &field : fields) { if (!field) { - throw std::invalid_argument( - "Invalid schema: field schema must not be null"); + continue; } - } - for (auto &field : fields) { auto c = std::make_shared(*field); fields_.push_back(c); fields_map_[field->name()] = c; diff --git a/tests/db/index/common/doc_test.cc b/tests/db/index/common/doc_test.cc index f8bfe6a9c..ffac59767 100644 --- a/tests/db/index/common/doc_test.cc +++ b/tests/db/index/common/doc_test.cc @@ -1738,6 +1738,25 @@ TEST_F(DocDetailedTest, DeserializeChecksNestedCountsAndBooleanRepresentation) { EXPECT_EQ(Doc::deserialize(buffer.data(), buffer.size()), nullptr); } +TEST_F(DocDetailedTest, + LongFieldNamesKeepDistinctIdentitiesAcrossSerialization) { + const std::string prefix(32, 'f'); + const std::vector names{prefix, prefix + "a", + prefix + std::string(31, 'a') + "x", + prefix + std::string(31, 'a') + "y"}; + for (size_t i = 0; i < names.size(); ++i) { + ASSERT_TRUE(test_doc_->set(names[i], static_cast(i))); + } + const auto bytes = test_doc_->serialize(); + const auto restored = Doc::deserialize(bytes.data(), bytes.size()); + ASSERT_NE(restored, nullptr); + for (size_t i = 0; i < names.size(); ++i) { + const auto value = restored->get(names[i]); + ASSERT_TRUE(value.has_value()) << names[i]; + EXPECT_EQ(value.value(), static_cast(i)); + } +} + TEST_F(DocDetailedTest, DeserializePreservesHistoricalTextWithoutRevalidation) { Doc doc; doc.set_pk(std::string("old\0id", 6)); diff --git a/tests/db/index/common/identifier_validation_test.cc b/tests/db/index/common/identifier_validation_test.cc index c3cc36864..99dcc2204 100644 --- a/tests/db/index/common/identifier_validation_test.cc +++ b/tests/db/index/common/identifier_validation_test.cc @@ -41,6 +41,13 @@ void ExpectInvalid(const Status &status, const std::string &message) { EXPECT_EQ(status.message().find("offset"), std::string::npos); } +void ExpectInvalidName(const Utf8NameValidator &validator, + std::string_view value, const std::string &reason) { + ExpectInvalid( + validator.validate(value), + std::string(validator.prefix) + "[" + format_name(value) + "] " + reason); +} + std::string Repeat(std::string_view text, size_t count) { std::string result; result.reserve(text.size() * count); @@ -97,14 +104,12 @@ TEST(IdentifierValidationTest, MeasuresLimitsInUtf8Bytes) { EXPECT_TRUE( validator.validate(std::string(max_bytes - 3, 'a') + u8"中").ok()); - const auto expected = std::string(validator.prefix) + " exceeds " + - std::to_string(max_bytes) + " bytes (got " + - std::to_string(max_bytes + 1) + ")"; - ExpectInvalid(validator.validate(std::string(max_bytes + 1, 'a')), - expected); - ExpectInvalid(validator.validate(emoji + "a"), expected); - ExpectInvalid(validator.validate(std::string(max_bytes - 2, 'a') + u8"中"), - expected); + const auto reason = "exceeds " + std::to_string(max_bytes) + + " bytes (got " + std::to_string(max_bytes + 1) + ")"; + ExpectInvalidName(validator, std::string(max_bytes + 1, 'a'), reason); + ExpectInvalidName(validator, emoji + "a", reason); + ExpectInvalidName(validator, std::string(max_bytes - 2, 'a') + u8"中", + reason); } } @@ -134,10 +139,8 @@ TEST(IdentifierValidationTest, RejectsMalformedUtf8) { for (const auto &validator : kUtf8NameValidators) { SCOPED_TRACE(validator.prefix); for (const auto &value : malformed) { - ExpectInvalid(validator.validate(value), - std::string(validator.prefix) + " is not valid UTF-8"); - ExpectInvalid(validator.validate("prefix" + value), - std::string(validator.prefix) + " is not valid UTF-8"); + ExpectInvalidName(validator, value, "is not valid UTF-8"); + ExpectInvalidName(validator, "prefix" + value, "is not valid UTF-8"); } } } @@ -147,10 +150,10 @@ TEST(IdentifierValidationTest, HonorsStringViewLengthAndEmbeddedNulls) { for (const auto &validator : kUtf8NameValidators) { SCOPED_TRACE(validator.prefix); EXPECT_TRUE(validator.validate(std::string_view(backing.data(), 6)).ok()); - ExpectInvalid(validator.validate(std::string_view(backing.data(), 5)), - std::string(validator.prefix) + " is not valid UTF-8"); - ExpectInvalid(validator.validate(std::string("a\0b", 3)), - std::string(validator.prefix) + " contains a null character"); + ExpectInvalidName(validator, std::string_view(backing.data(), 5), + "is not valid UTF-8"); + ExpectInvalidName(validator, std::string("a\0b", 3), + "contains a null character"); } } @@ -175,8 +178,7 @@ TEST(IdentifierValidationTest, RejectsEveryC0AndC1ControlByCodepoint) { } else if (codepoint == '\t') { reason = "contains a tab"; } - ExpectInvalid(validator.validate("a" + value + "b"), - std::string(validator.prefix) + " " + reason); + ExpectInvalidName(validator, "a" + value + "b", reason); } // These continuation bytes overlap the C1 byte range, but their decoded // codepoints are ordinary letters/symbols and must not be rejected. @@ -186,11 +188,31 @@ TEST(IdentifierValidationTest, RejectsEveryC0AndC1ControlByCodepoint) { TEST(IdentifierValidationTest, DistinguishesUnicodeLineAndParagraphSeparators) { for (const auto &validator : kUtf8NameValidators) { - ExpectInvalid(validator.validate(u8"a\u2028b"), - std::string(validator.prefix) + " contains a line separator"); + ExpectInvalidName(validator, u8"a\u2028b", "contains a line separator"); + ExpectInvalidName(validator, u8"a\u2029b", + "contains a paragraph separator"); + } +} + +TEST(IdentifierValidationTest, Utf8ErrorsIncludeEscapedAndBoundedPreviews) { + for (const auto &validator : kUtf8NameValidators) { + const std::string prefix = validator.prefix; + ExpectInvalid(validator.validate("order\n[123]"), + prefix + "[order\\n\\[123\\]] contains a newline"); + ExpectInvalid(validator.validate("order\xff"), + prefix + "[order\\xFF] is not valid UTF-8"); ExpectInvalid( - validator.validate(u8"a\u2029b"), - std::string(validator.prefix) + " contains a paragraph separator"); + validator.validate(std::string(40, 'a') + "\n"), + prefix + "[" + std::string(32, 'a') + "...] contains a newline"); + const auto status = validator.validate(std::string(10000, '\xff')); + EXPECT_EQ(status.message().find(prefix + "[" + Repeat("\\xFF", 32) + + "...] exceeds "), + 0u); + EXPECT_LT(status.message().size(), 256u); + for (unsigned char byte : status.message()) { + EXPECT_GE(byte, 0x20); + EXPECT_LE(byte, 0x7E); + } } } diff --git a/tests/db/index/common/schema_test.cc b/tests/db/index/common/schema_test.cc index 897f6382f..b72f1bc8d 100644 --- a/tests/db/index/common/schema_test.cc +++ b/tests/db/index/common/schema_test.cc @@ -19,12 +19,21 @@ using namespace zvec; -TEST(CollectionSchemaTest, RejectsNullFieldObjectsWithoutDroppingThem) { +TEST(CollectionSchemaTest, SkipsNullFieldObjectsDuringConstruction) { + CollectionSchema empty("schema", {nullptr}); + EXPECT_TRUE(empty.fields().empty()); + EXPECT_EQ(empty.validate().code(), StatusCode::INVALID_ARGUMENT); + auto field = std::make_shared("valid", DataType::INT32); - EXPECT_THROW(CollectionSchema("schema", {nullptr}), std::invalid_argument); - EXPECT_THROW(CollectionSchema("schema", {field, nullptr}), - std::invalid_argument); + CollectionSchema schema("schema", {nullptr, field, nullptr}); + ASSERT_EQ(schema.fields().size(), 1u); + ASSERT_NE(schema.get_field("valid"), nullptr); + EXPECT_EQ(schema.get_field("valid")->data_type(), DataType::INT32); + EXPECT_TRUE(schema.validate().ok()); +} +TEST(CollectionSchemaTest, RejectsNullFieldObjectsInMutations) { + auto field = std::make_shared("valid", DataType::INT32); CollectionSchema schema("schema", {field}); const CollectionSchema before(schema); EXPECT_EQ(schema.add_field(nullptr).code(), StatusCode::INVALID_ARGUMENT); diff --git a/tests/db/relaxed_validation_test.cc b/tests/db/relaxed_validation_test.cc index ef7978067..bbf642752 100644 --- a/tests/db/relaxed_validation_test.cc +++ b/tests/db/relaxed_validation_test.cc @@ -164,8 +164,13 @@ TEST_F(RelaxedValidationTest, ShortAndMaximumLengthCollectionNamesPersist) { TEST_F(RelaxedValidationTest, MaximumLengthFieldsSupportIndexesAndFilters) { const std::string scalar = "s" + std::string(63, 'a'); + const std::string second_scalar = scalar.substr(0, 63) + "b"; const std::string vector = "v" + std::string(63, 'b'); auto schema = MakeSchema("x", scalar); + ASSERT_TRUE(schema + .add_field(std::make_shared( + second_scalar, DataType::INT32, false)) + .ok()); ASSERT_TRUE(schema .add_field(std::make_shared( vector, DataType::VECTOR_FP32, 4, false, @@ -174,6 +179,7 @@ TEST_F(RelaxedValidationTest, MaximumLengthFieldsSupportIndexesAndFilters) { ASSERT_NO_FATAL_FAILURE(Create(schema)); const std::vector values{1.0f, 2.0f, 3.0f, 4.0f}; Doc doc = MakeDoc(u8"文档:1", 42, scalar); + ASSERT_TRUE(doc.set(second_scalar, 7)); ASSERT_TRUE(doc.set>(vector, values)); std::vector docs{doc}; ASSERT_NO_FATAL_FAILURE(ExpectWrite(collection_->insert(docs), 1)); @@ -185,6 +191,9 @@ TEST_F(RelaxedValidationTest, MaximumLengthFieldsSupportIndexesAndFilters) { status = collection_->create_index( vector, std::make_shared(MetricType::L2)); ASSERT_TRUE(status.ok()) << status.message(); + status = collection_->create_index(second_scalar, + std::make_shared()); + ASSERT_TRUE(status.ok()) << status.message(); ASSERT_NO_FATAL_FAILURE(Reopen()); // Fetch supplies the query vector, covering lookup by a newly allowed ID. @@ -201,13 +210,27 @@ TEST_F(RelaxedValidationTest, MaximumLengthFieldsSupportIndexesAndFilters) { query.target_.set_vector( std::string(reinterpret_cast(stored_vector->data()), stored_vector->size() * sizeof(float))); - query.filter_ = scalar + " = 42"; - query.output_fields_ = std::vector{scalar}; + query.filter_ = scalar + " = 42 AND " + second_scalar + " = 7"; + query.output_fields_ = std::vector{scalar, second_scalar}; auto matches = collection_->query(query); ASSERT_TRUE(matches.has_value()) << matches.error().message(); ASSERT_EQ(matches.value().size(), 1u); EXPECT_EQ(matches.value()[0]->pk(), doc.pk()); EXPECT_EQ(matches.value()[0]->get(scalar), 42); + EXPECT_EQ(matches.value()[0]->get(second_scalar), 7); +} + +TEST_F(RelaxedValidationTest, MaximumLengthFieldsPersistWithBufferedStorage) { + options_.enable_mmap_ = false; + const std::string field(64, 'f'); + ASSERT_NO_FATAL_FAILURE(Create(MakeSchema("buffered", field))); + std::vector docs{MakeDoc("doc", 42, field)}; + ASSERT_NO_FATAL_FAILURE(ExpectWrite(collection_->insert(docs), 1)); + const auto status = collection_->flush(); + ASSERT_TRUE(status.ok()) << status.message(); + ASSERT_NO_FATAL_FAILURE(Reopen()); + EXPECT_TRUE(collection_->schema().value().has_field(field)); + ASSERT_NO_FATAL_FAILURE(ExpectValue("doc", 42, field)); } TEST_F(RelaxedValidationTest, InvalidRenameLeavesSchemaAndDataUnchanged) { @@ -255,6 +278,7 @@ TEST_F(RelaxedValidationTest, InvalidIdRejectsWholeBatchBeforeWriting) { EXPECT_EQ(result.error().message().find("Invalid doc:"), 0u); EXPECT_NE(result.error().message().find("null character"), std::string::npos); + EXPECT_NE(result.error().message().find("id[bad\\0id]"), std::string::npos); EXPECT_EQ(result.error().message().find("offset"), std::string::npos); ASSERT_NO_FATAL_FAILURE(ExpectValue("existing", 1)); auto missing = collection_->fetch({"new:id"}); From c4c34992332ffd2cf806117a07de1d512a14b373 Mon Sep 17 00:00:00 2001 From: Qinren Zhou Date: Wed, 16 Sep 2026 14:54:38 +0800 Subject: [PATCH 04/11] fix format --- src/binding/c/c_api.cc | 9 +- src/db/collection.cc | 4 +- src/db/index/common/doc.cc | 512 ++++++++++-------- src/db/index/common/schema.cc | 14 +- .../sqlengine/planner/vector_recall_node.cc | 2 +- src/include/zvec/db/doc.h | 5 + tests/c/c_api_test.c | 35 +- .../relaxed_validation_recovery_test.cc | 26 +- tests/db/index/common/doc_test.cc | 93 ---- tests/db/index/common/schema_test.cc | 52 +- tests/db/relaxed_validation_test.cc | 67 ++- 11 files changed, 418 insertions(+), 401 deletions(-) diff --git a/src/binding/c/c_api.cc b/src/binding/c/c_api.cc index b562ac7d9..9e3fab7d4 100644 --- a/src/binding/c/c_api.cc +++ b/src/binding/c/c_api.cc @@ -4731,10 +4731,8 @@ zvec_error_code_t zvec_doc_serialize(const zvec_doc_t *doc, uint8_t **data, zvec_error_code_t zvec_doc_deserialize(const uint8_t *data, size_t size, zvec_doc_t **doc) { - if (doc) *doc = nullptr; if (!data || !doc || size == 0) { - SET_LAST_ERROR(ZVEC_ERROR_INVALID_ARGUMENT, - "Invalid doc: data, size and document output must be provided"); + set_last_error("Invalid arguments"); return ZVEC_ERROR_INVALID_ARGUMENT; } @@ -4742,9 +4740,8 @@ zvec_error_code_t zvec_doc_deserialize(const uint8_t *data, size_t size, "Failed to deserialize document", auto deserialized_doc = zvec::Doc::deserialize(data, size); if (!deserialized_doc) { - SET_LAST_ERROR(ZVEC_ERROR_INVALID_ARGUMENT, - "Invalid doc: serialized data is incomplete or invalid"); - return ZVEC_ERROR_INVALID_ARGUMENT; + set_last_error("Failed to deserialize document"); + return ZVEC_ERROR_INTERNAL_ERROR; } // Create a new Doc by copying the deserialized content diff --git a/src/db/collection.cc b/src/db/collection.cc index 3b3ca6483..a96964494 100644 --- a/src/db/collection.cc +++ b/src/db/collection.cc @@ -1478,8 +1478,8 @@ Status CollectionImpl::alter_column(const std::string &column_name, CHECK_DESTROY_RETURN_STATUS(destroyed_, false); CHECK_CLOSED_RETURN_STATUS(closed_, false); - auto field_copy = - new_column_schema ? std::make_shared(*new_column_schema) + auto field_copy = new_column_schema + ? std::make_shared(*new_column_schema) : nullptr; auto s = validate(column_name, field_copy, "", rename, ColumnOp::ALTER); CHECK_RETURN_STATUS(s); diff --git a/src/db/index/common/doc.cc b/src/db/index/common/doc.cc index 99dceca30..856b59efb 100644 --- a/src/db/index/common/doc.cc +++ b/src/db/index/common/doc.cc @@ -12,11 +12,11 @@ // See the License for the specific language governing permissions and // limitations under the License. -#include #include #include #include #include +#include #include #include #include @@ -122,11 +122,28 @@ namespace { template T byte_swap(T value) { - T result; - const auto *source = reinterpret_cast(&value); - auto *destination = reinterpret_cast(&result); - std::reverse_copy(source, source + sizeof(T), destination); - return result; + if constexpr (std::is_same_v) { + uint16_t val; + std::memcpy(&val, static_cast(&value), sizeof(val)); + val = ailego_bswap16(val); + float16_t result; + std::memcpy(static_cast(&result), &val, sizeof(result)); + return result; + } else if constexpr (sizeof(T) == 1) { + return value; + } else if constexpr (sizeof(T) == 2) { + return (value << 8) | ((value >> 8) & 0xFF); + } else if constexpr (sizeof(T) == 4) { + return static_cast(ailego_bswap32(static_cast(value))); + } else if constexpr (sizeof(T) == 8) { + return static_cast(ailego_bswap64(static_cast(value))); + } else { + T result = 0; + for (size_t i = 0; i < sizeof(T); ++i) { + result |= ((value >> (i * 8)) & 0xFF) << ((sizeof(T) - 1 - i) * 8); + } + return result; + } } template @@ -139,179 +156,17 @@ void write_value_to_buffer(std::vector &buffer, const T &value) { buffer.insert(buffer.end(), bytes, bytes + sizeof(T)); } -// Read persisted values only after checking their complete byte range. Length -// fields are checked before allocation, including each nested array/string. -class DocBufferReader { - public: - DocBufferReader(const uint8_t *data, size_t size) - : data_(data), remaining_(size) {} - - size_t remaining() const { - return remaining_; - } - - bool read_bytes(void *destination, size_t size) { - if (size > remaining_) return false; - if (size != 0) { - std::memcpy(destination, data_, size); - data_ += size; - remaining_ -= size; - } - return true; - } - - template - bool read_native(T &value) { - return read_bytes(&value, sizeof(T)); - } - - bool read_string_bytes(std::string &value, size_t size) { - if (size > remaining_) return false; - value.assign(reinterpret_cast(data_), size); - data_ += size; - remaining_ -= size; - return true; - } - - bool read_value(Doc::Value &value) { - uint8_t type; - if (!read_native(type)) return false; - switch (type) { - case TYPE_EMPTY: - value = std::monostate{}; - return true; - case TYPE_BOOL: - return read_as(value); - case TYPE_INT32: - return read_as(value); - case TYPE_UINT32: - return read_as(value); - case TYPE_INT64: - return read_as(value); - case TYPE_UINT64: - return read_as(value); - case TYPE_FLOAT: - return read_as(value); - case TYPE_DOUBLE: - return read_as(value); - case TYPE_STRING: - return read_as(value); - case TYPE_VECTOR_BOOL: - return read_as>(value); - case TYPE_VECTOR_INT8: - return read_as>(value); - case TYPE_VECTOR_INT16: - return read_as>(value); - case TYPE_VECTOR_INT32: - return read_as>(value); - case TYPE_VECTOR_INT64: - return read_as>(value); - case TYPE_VECTOR_UINT32: - return read_as>(value); - case TYPE_VECTOR_UINT64: - return read_as>(value); - case TYPE_VECTOR_FLOAT16: - return read_as>(value); - case TYPE_VECTOR_FLOAT: - return read_as>(value); - case TYPE_VECTOR_DOUBLE: - return read_as>(value); - case TYPE_VECTOR_STRING: - return read_as>(value); - case TYPE_VECTOR_PAIR_INT_FLOAT: - return read_as, std::vector>>( - value); - case TYPE_VECTOR_PAIR_INT_FLOAT16: - return read_as< - std::pair, std::vector>>(value); - default: - return false; - } - } - - private: - template - bool read_little_endian(T &value) { - if (!read_native(value)) return false; - if (IS_BIG_ENDIAN) { - auto *bytes = reinterpret_cast(&value); - std::reverse(bytes, bytes + sizeof(T)); - } - return true; - } - - template - bool read_as(Doc::Value &out) { - T value; - if (!read(value)) return false; - out = std::move(value); - return true; - } - - template - bool read(T &value) { - return read_little_endian(value); - } - - bool read(bool &value) { - static_assert(sizeof(bool) == sizeof(uint8_t)); - uint8_t byte; - if (!read_native(byte) || byte > 1) return false; - value = byte != 0; - return true; - } - - bool read(std::string &value) { - uint32_t size; - return read_little_endian(size) && read_string_bytes(value, size); - } - - template - bool read(std::vector &values) { - uint32_t count; - if (!read_little_endian(count)) return false; - if constexpr (std::is_same_v) { - // Each string contains at least its four-byte length prefix. - if (count > remaining_ / sizeof(uint32_t)) return false; - values.reserve(count); - for (uint32_t i = 0; i < count; ++i) { - std::string value; - if (!read(value)) return false; - values.push_back(std::move(value)); - } - } else if constexpr (std::is_same_v) { - if (count > remaining_ / sizeof(bool)) return false; - values.reserve(count); - for (uint32_t i = 0; i < count; ++i) { - bool value; - if (!read(value)) return false; - values.push_back(value); - } - } else { - // Division avoids overflow before checking the allocation/copy size. - if (count > remaining_ / sizeof(T)) return false; - values.resize(count); - if (!read_bytes(values.data(), static_cast(count) * sizeof(T))) { - return false; - } - if (IS_BIG_ENDIAN) { - for (auto &value : values) { - auto *bytes = reinterpret_cast(&value); - std::reverse(bytes, bytes + sizeof(T)); - } - } - } - return true; - } +template +T read_value_from_buffer(const uint8_t *&data) { + T value; + std::memcpy(&value, data, sizeof(T)); + data += sizeof(T); - template - bool read(std::pair, std::vector> &value) { - return read(value.first) && read(value.second); + if (IS_BIG_ENDIAN) { + value = byte_swap(value); } - - const uint8_t *data_; - size_t remaining_; -}; + return value; +} template std::string vec_to_string(const std::vector &v) { @@ -343,6 +198,11 @@ void Doc::write_to_buffer(std::vector &buffer, const void *src, buffer.insert(buffer.end(), bytes, bytes + size); } +void Doc::read_from_buffer(const uint8_t *&data, void *dest, size_t size) { + std::memcpy(dest, data, size); + data += size; +} + void Doc::serialize_value(std::vector &buffer, const Value &value) { std::visit( [&buffer](const auto &v) { @@ -576,6 +436,238 @@ void Doc::serialize_value(std::vector &buffer, const Value &value) { } +Doc::Value Doc::deserialize_value(const uint8_t *&data) { + uint8_t type; + read_from_buffer(data, &type, sizeof(type)); + + switch (type) { + case TYPE_EMPTY: { + return std::monostate{}; + } + case TYPE_BOOL: { + bool v; + read_from_buffer(data, &v, sizeof(v)); + return v; + } + case TYPE_INT32: { + return read_value_from_buffer(data); + } + case TYPE_INT64: { + return read_value_from_buffer(data); + } + case TYPE_UINT32: { + return read_value_from_buffer(data); + } + case TYPE_UINT64: { + return read_value_from_buffer(data); + } + case TYPE_FLOAT: { + return read_value_from_buffer(data); + } + case TYPE_DOUBLE: { + return read_value_from_buffer(data); + } + case TYPE_STRING: { + uint32_t len = read_value_from_buffer(data); + std::string v(reinterpret_cast(data), len); + data += len; + return v; + } + case TYPE_VECTOR_BOOL: { + uint32_t len = read_value_from_buffer(data); + std::vector v; + v.reserve(len); + for (uint32_t i = 0; i < len; ++i) { + bool b; + read_from_buffer(data, &b, sizeof(b)); + v.push_back(b); + } + return v; + } + case TYPE_VECTOR_INT8: { + uint32_t len = read_value_from_buffer(data); + std::vector v(len); + read_from_buffer(data, v.data(), len * sizeof(int8_t)); + return v; + } + case TYPE_VECTOR_INT16: { + uint32_t len = read_value_from_buffer(data); + std::vector v(len); + if (IS_BIG_ENDIAN) { + for (uint32_t i = 0; i < len; ++i) { + v[i] = byte_swap(read_value_from_buffer(data)); + } + } else { + read_from_buffer(data, v.data(), len * sizeof(int16_t)); + } + return v; + } + case TYPE_VECTOR_INT32: { + uint32_t len = read_value_from_buffer(data); + std::vector v(len); + if (IS_BIG_ENDIAN) { + for (uint32_t i = 0; i < len; ++i) { + v[i] = byte_swap(read_value_from_buffer(data)); + } + } else { + read_from_buffer(data, v.data(), len * sizeof(int32_t)); + } + return v; + } + case TYPE_VECTOR_INT64: { + uint32_t len = read_value_from_buffer(data); + std::vector v(len); + if (IS_BIG_ENDIAN) { + for (uint32_t i = 0; i < len; ++i) { + v[i] = byte_swap(read_value_from_buffer(data)); + } + } else { + read_from_buffer(data, v.data(), len * sizeof(int64_t)); + } + return v; + } + case TYPE_VECTOR_UINT32: { + uint32_t len = read_value_from_buffer(data); + std::vector v(len); + if (IS_BIG_ENDIAN) { + for (uint32_t i = 0; i < len; ++i) { + v[i] = byte_swap(read_value_from_buffer(data)); + } + } else { + read_from_buffer(data, v.data(), len * sizeof(uint32_t)); + } + return v; + } + case TYPE_VECTOR_UINT64: { + uint32_t len = read_value_from_buffer(data); + std::vector v(len); + if (IS_BIG_ENDIAN) { + for (uint32_t i = 0; i < len; ++i) { + v[i] = byte_swap(read_value_from_buffer(data)); + } + } else { + read_from_buffer(data, v.data(), len * sizeof(uint64_t)); + } + return v; + } + case TYPE_VECTOR_FLOAT: { + uint32_t len = read_value_from_buffer(data); + std::vector v(len); + if (IS_BIG_ENDIAN) { + for (uint32_t i = 0; i < len; ++i) { + v[i] = byte_swap(read_value_from_buffer(data)); + } + } else { + read_from_buffer(data, v.data(), len * sizeof(float)); + } + return v; + } + case TYPE_VECTOR_DOUBLE: { + uint32_t len = read_value_from_buffer(data); + std::vector v(len); + if (IS_BIG_ENDIAN) { + for (uint32_t i = 0; i < len; ++i) { + v[i] = byte_swap(read_value_from_buffer(data)); + } + } else { + read_from_buffer(data, v.data(), len * sizeof(double)); + } + return v; + } + case TYPE_VECTOR_FLOAT16: { + uint32_t len = read_value_from_buffer(data); + std::vector v(len); + if (IS_BIG_ENDIAN) { + for (uint32_t i = 0; i < len; ++i) { + v[i] = byte_swap(read_value_from_buffer(data)); + } + } else { + read_from_buffer(data, v.data(), len * sizeof(float16_t)); + } + return v; + } + case TYPE_VECTOR_STRING: { + uint32_t len = read_value_from_buffer(data); + std::vector v; + v.reserve(len); + for (uint32_t i = 0; i < len; ++i) { + uint32_t str_len = read_value_from_buffer(data); + std::string s(reinterpret_cast(data), str_len); + data += str_len; + v.push_back(s); + } + return v; + } + case TYPE_VECTOR_PAIR_INT_FLOAT: { + uint32_t len = read_value_from_buffer(data); + std::pair, std::vector> v; + v.first.reserve(len); + if (IS_BIG_ENDIAN) { + for (uint32_t i = 0; i < len; ++i) { + v.first.push_back( + byte_swap(read_value_from_buffer(data))); + } + } else { + for (uint32_t i = 0; i < len; ++i) { + uint32_t first; + read_from_buffer(data, &first, sizeof(first)); + v.first.push_back(first); + } + } + len = read_value_from_buffer(data); + v.second.reserve(len); + if (IS_BIG_ENDIAN) { + for (uint32_t i = 0; i < len; ++i) { + v.second.push_back( + byte_swap(read_value_from_buffer(data))); + } + } else { + for (uint32_t i = 0; i < len; ++i) { + float second; + read_from_buffer(data, &second, sizeof(second)); + v.second.push_back(second); + } + } + return v; + } + case TYPE_VECTOR_PAIR_INT_FLOAT16: { + uint32_t len = read_value_from_buffer(data); + std::pair, std::vector> v; + v.first.reserve(len); + if (IS_BIG_ENDIAN) { + for (uint32_t i = 0; i < len; ++i) { + v.first.push_back( + byte_swap(read_value_from_buffer(data))); + } + } else { + for (uint32_t i = 0; i < len; ++i) { + uint32_t first; + read_from_buffer(data, &first, sizeof(first)); + v.first.push_back(first); + } + } + len = read_value_from_buffer(data); + v.second.reserve(len); + if (IS_BIG_ENDIAN) { + for (uint32_t i = 0; i < len; ++i) { + v.second.push_back( + byte_swap(read_value_from_buffer(data))); + } + } else { + for (uint32_t i = 0; i < len; ++i) { + float16_t second; + read_from_buffer(data, &second, sizeof(second)); + v.second.push_back(second); + } + } + return v; + } + + default: + throw std::runtime_error("Unknown value type: " + std::to_string(type)); + } +} + std::vector Doc::serialize() const { std::vector buffer; uint32_t pk_len = static_cast(pk_.size()); @@ -600,40 +692,37 @@ std::vector Doc::serialize() const { return buffer; } -Doc::Ptr Doc::deserialize(const uint8_t *data, size_t size) { - if (!data) return nullptr; - DocBufferReader reader(data, size); - auto doc = std::make_shared(); - uint32_t pk_length; - uint32_t operation; - uint32_t field_count; - // The document header and field-name lengths retain their existing native - // representation; value payloads use the existing little-endian encoding. - if (!reader.read_native(pk_length) || - !reader.read_string_bytes(doc->pk_, pk_length) || - !reader.read_native(doc->score_) || !reader.read_native(doc->doc_id_) || - !reader.read_native(operation) || - operation > static_cast(Operator::DELETE) || - !reader.read_native(field_count)) { - return nullptr; - } - doc->op_ = static_cast(operation); - // Even an empty name and a null value require a length prefix and type byte. - if (field_count > reader.remaining() / (sizeof(uint32_t) + sizeof(uint8_t))) { - return nullptr; - } +Doc::Ptr Doc::deserialize(const uint8_t *data, size_t /*size*/) { + const uint8_t *ptr = data; + Doc::Ptr doc = std::make_shared(); + + uint32_t pk_len = read_value_from_buffer(ptr); + std::string pk(reinterpret_cast(ptr), pk_len); + ptr += pk_len; + doc->set_pk(pk); + + float score = read_value_from_buffer(ptr); + doc->set_score(score); + + uint64_t doc_id = read_value_from_buffer(ptr); + doc->set_doc_id(doc_id); + + Operator op; + read_from_buffer(ptr, &op, sizeof(op)); + doc->set_operator(op); + + uint32_t field_count = read_value_from_buffer(ptr); + for (uint32_t i = 0; i < field_count; ++i) { - uint32_t name_length; - std::string name; - Value value; - if (!reader.read_native(name_length) || - !reader.read_string_bytes(name, name_length) || - !reader.read_value(value) || - !doc->fields_.emplace(std::move(name), std::move(value)).second) { - return nullptr; - } + uint32_t name_len = read_value_from_buffer(ptr); + std::string field_name(reinterpret_cast(ptr), name_len); + ptr += name_len; + + Doc::Value value = deserialize_value(ptr); + doc->fields_[field_name] = value; } - return reader.remaining() == 0 ? doc : nullptr; + + return doc; } Status Doc::validate_and_sanitize(const CollectionSchema::Ptr &schema, @@ -642,9 +731,8 @@ Status Doc::validate_and_sanitize(const CollectionSchema::Ptr &schema, return Status::InternalError("schema is null during doc validation"); } - auto id_status = validate_document_id(pk_); - if (!id_status.ok()) { - return id_status; + if (auto s = validate_document_id(pk_); !s.ok()) { + return s; } // check doc fields match schema diff --git a/src/db/index/common/schema.cc b/src/db/index/common/schema.cc index 5f3e91b78..9e59b62ec 100644 --- a/src/db/index/common/schema.cc +++ b/src/db/index/common/schema.cc @@ -81,8 +81,9 @@ static Status validate_fts_index_params(const FieldSchema &field) { } Status FieldSchema::validate() const { - auto name_status = validate_field_name(name_); - CHECK_RETURN_STATUS(name_status); + if (auto s = validate_field_name(name_); !s.ok()) { + return s; + } if (data_type_ == DataType::UNDEFINED) { return Status::InvalidArgument("Invalid schema: field[", name_, @@ -405,8 +406,9 @@ std::string FieldSchema::to_string_formatted(int indent_level) const { } Status CollectionSchema::validate() const { - auto name_status = validate_collection_name(name_); - CHECK_RETURN_STATUS(name_status); + if (auto s = validate_collection_name(name_); !s.ok()) { + return s; + } std::unordered_set names; for (const auto &field : fields_) { if (!field) { @@ -490,7 +492,7 @@ Status CollectionSchema::add_field(FieldSchema::Ptr column_schema) { } // Check if field already exists if (has_field(column_schema->name())) { - return Status::AlreadyExists("field[", format_name(column_schema->name()), + return Status::AlreadyExists("field[", column_schema->name(), "] already exists in schema"); } @@ -524,7 +526,7 @@ Status CollectionSchema::alter_field( // If renaming to an existing field name (and it's not the same field) if (new_column_name != column_name && has_field(new_column_name)) { - return Status::AlreadyExists("field[", format_name(new_column_name), + return Status::AlreadyExists("field[", new_column_name, "] already exists in schema"); } diff --git a/src/db/sqlengine/planner/vector_recall_node.cc b/src/db/sqlengine/planner/vector_recall_node.cc index 58e4e8d43..b1900b809 100644 --- a/src/db/sqlengine/planner/vector_recall_node.cc +++ b/src/db/sqlengine/planner/vector_recall_node.cc @@ -67,7 +67,7 @@ VectorRecallNode::VectorRecallNode(Segment::Ptr segment, arrow::AsyncGenerator> VectorRecallNode::gen() { auto state_ptr = std::make_shared(shared_from_this()); return [state_ptr = std::move(state_ptr)]() mutable - -> arrow::Future> { + -> arrow::Future> { auto &state = *state_ptr; if (!state.iter_) { auto vector_ret = state.self_->prepare(); diff --git a/src/include/zvec/db/doc.h b/src/include/zvec/db/doc.h index 2366708dd..785d9c1d8 100644 --- a/src/include/zvec/db/doc.h +++ b/src/include/zvec/db/doc.h @@ -310,9 +310,14 @@ class ZVEC_API Doc { private: static void serialize_value(std::vector &buffer, const Value &value); + static Value deserialize_value(const uint8_t *&data, uint8_t type); + static Value deserialize_value(const uint8_t *&data); + static void write_to_buffer(std::vector &buffer, const void *src, size_t size); + static void read_from_buffer(const uint8_t *&data, void *dest, size_t size); + struct ValueEqual; private: diff --git a/tests/c/c_api_test.c b/tests/c/c_api_test.c index 6b0666b35..a4402703f 100644 --- a/tests/c/c_api_test.c +++ b/tests/c/c_api_test.c @@ -1189,7 +1189,7 @@ void test_validation_last_error(void) { TEST_ASSERT(zvec_collection_schema_validate(schema, NULL) == ZVEC_ERROR_INVALID_ARGUMENT); check_last_error(ZVEC_ERROR_INVALID_ARGUMENT, - "Invalid schema: collection name is not valid UTF-8"); + "Invalid schema: collection name[\\xFF] is not valid UTF-8"); // Replacing a previous error must update both text and code. TEST_ASSERT(zvec_field_schema_validate(field, NULL) == ZVEC_ERROR_INVALID_ARGUMENT); @@ -1221,7 +1221,7 @@ void test_validation_last_error(void) { ZVEC_ERROR_INVALID_ARGUMENT); TEST_ASSERT(collection == NULL); check_last_error(ZVEC_ERROR_INVALID_ARGUMENT, - "Invalid schema: collection name is not valid UTF-8"); + "Invalid schema: collection name[\\xFF] is not valid UTF-8"); zvec_field_schema_destroy(field); zvec_collection_schema_destroy(schema); @@ -1268,7 +1268,7 @@ void test_batch_validation_errors(void) { zvec_doc_set_pk(invalid_doc, "\xff"); const zvec_doc_t *invalid_inputs[] = {NULL, invalid_doc}; const char *reasons[] = {"document must not be null", - "id is not valid UTF-8"}; + "id[\\xFF] is not valid UTF-8"}; for (size_t i = 0; i < 2; ++i) { const zvec_doc_t *docs[] = {valid_doc, invalid_inputs[i]}; for (size_t op = 0; op < 3; ++op) { @@ -1278,7 +1278,9 @@ void test_batch_validation_errors(void) { TEST_ASSERT(success_count == 0); TEST_ASSERT(error_count == 2); check_last_error(ZVEC_ERROR_INVALID_ARGUMENT, reasons[i]); - check_last_error(ZVEC_ERROR_INVALID_ARGUMENT, "document at index 1"); + if (i == 0) { + check_last_error(ZVEC_ERROR_INVALID_ARGUMENT, "document at index 1"); + } zvec_write_result_t *results = (zvec_write_result_t *)(uintptr_t)1; size_t result_count = 123; @@ -3645,31 +3647,6 @@ void test_doc_serialization(void) { TEST_ASSERT(err == ZVEC_OK); TEST_ASSERT(deserialized_int32 == -2147483648); - const size_t truncated_sizes[] = {1, data_size / 2, data_size - 1}; - for (size_t i = 0; i < sizeof(truncated_sizes) / sizeof(truncated_sizes[0]); - ++i) { - zvec_doc_t *invalid_doc = (zvec_doc_t *)(uintptr_t)1; - TEST_ASSERT(zvec_doc_deserialize(serialized_data, truncated_sizes[i], - &invalid_doc) == - ZVEC_ERROR_INVALID_ARGUMENT); - TEST_ASSERT(invalid_doc == NULL); - check_last_error(ZVEC_ERROR_INVALID_ARGUMENT, - "Invalid doc: serialized data is incomplete or invalid"); - } - zvec_doc_t *invalid_doc = (zvec_doc_t *)(uintptr_t)1; - TEST_ASSERT(zvec_doc_deserialize(NULL, data_size, &invalid_doc) == - ZVEC_ERROR_INVALID_ARGUMENT); - TEST_ASSERT(invalid_doc == NULL); - invalid_doc = (zvec_doc_t *)(uintptr_t)1; - TEST_ASSERT(zvec_doc_deserialize(serialized_data, 0, &invalid_doc) == - ZVEC_ERROR_INVALID_ARGUMENT); - TEST_ASSERT(invalid_doc == NULL); - TEST_ASSERT(zvec_doc_deserialize(serialized_data, data_size, NULL) == - ZVEC_ERROR_INVALID_ARGUMENT); - check_last_error( - ZVEC_ERROR_INVALID_ARGUMENT, - "Invalid doc: data, size and document output must be provided"); - zvec_free_uint8_array(serialized_data); free(string_field.value.string_value.data); zvec_doc_destroy(deserialized_doc); diff --git a/tests/db/crash_recovery/relaxed_validation_recovery_test.cc b/tests/db/crash_recovery/relaxed_validation_recovery_test.cc index c33cf820c..ddd9ebf54 100644 --- a/tests/db/crash_recovery/relaxed_validation_recovery_test.cc +++ b/tests/db/crash_recovery/relaxed_validation_recovery_test.cc @@ -64,8 +64,7 @@ std::map ReadManifests(const std::string &path) { // Only called in ASSERT_EXIT children. Intentionally skip collection cleanup. void WriteStringDocsAndExit(const std::string &path, const std::vector &ids, - const std::vector &values, - bool malformed_record = false) { + const std::vector &values) { CollectionSchema schema("wal_recovery"); if (!schema .add_field( @@ -88,13 +87,6 @@ void WriteStringDocsAndExit(const std::string &path, for (const auto &status : result.value()) { if (!status.ok()) std::_Exit(4); } - if (malformed_record) { - auto wal = WalFile::Create(FindWal(path)); - if (wal->open(WalOptions{}) != 0 || - wal->append("invalid encoded document") != 0) { - std::_Exit(5); - } - } std::_Exit(0); } @@ -478,22 +470,6 @@ TEST_F(RelaxedValidationDeathTest, CorruptWalFailsOpenWithoutReplacingFiles) { } } -TEST_F(RelaxedValidationDeathTest, InvalidDocumentPayloadFailsRecovery) { - ::testing::FLAGS_gtest_death_test_style = "threadsafe"; - ASSERT_EXIT(WriteStringDocsAndExit(path_, {"prefix"}, {"before"}, true), - ::testing::ExitedWithCode(0), ""); - const auto wal_path = FindWal(path_); - ASSERT_FALSE(wal_path.empty()); - const auto bytes = ReadFileBytes(wal_path); - const auto manifests = ReadManifests(path_); - auto opened = Collection::Open(path_, CollectionOptions{}); - ASSERT_FALSE(opened.has_value()); - EXPECT_NE(opened.error().message().find("Corrupt WAL document"), - std::string::npos); - EXPECT_EQ(ReadFileBytes(wal_path), bytes); - EXPECT_EQ(ReadManifests(path_), manifests); -} - TEST_F(RelaxedValidationDeathTest, CorruptTailDoesNotApplyUpsertPrefix) { ::testing::FLAGS_gtest_death_test_style = "threadsafe"; ASSERT_EXIT(WriteUpsertsAndExit(path_, true, true), diff --git a/tests/db/index/common/doc_test.cc b/tests/db/index/common/doc_test.cc index ffac59767..95342df33 100644 --- a/tests/db/index/common/doc_test.cc +++ b/tests/db/index/common/doc_test.cc @@ -14,7 +14,6 @@ #include "zvec/db/doc.h" #include -#include #include #include #include @@ -1026,14 +1025,6 @@ TEST_F(DocDetailedTest, SerializeValueCoverage) { auto buffer = doc.serialize(); EXPECT_FALSE(buffer.empty()); - // Exercise every truncated header, scalar, vector and nested string using - // independently allocated buffers, so ASan also catches accidental overreads. - for (size_t size = 0; size < buffer.size(); ++size) { - SCOPED_TRACE(size); - std::vector truncated(buffer.begin(), buffer.begin() + size); - EXPECT_EQ(Doc::deserialize(truncated.data(), truncated.size()), nullptr); - } - auto deserialized_doc = Doc::deserialize(buffer.data(), buffer.size()); EXPECT_NE(deserialized_doc, nullptr); @@ -1654,90 +1645,6 @@ TEST_F(DocDetailedTest, FieldExistenceChecks) { } -TEST_F(DocDetailedTest, DeserializeRejectsMalformedLengthsAndTags) { - Doc doc; - doc.set_pk("id"); - doc.set("field", "payload"); - const auto valid = doc.serialize(); - const size_t operation = - sizeof(uint32_t) + 2 + sizeof(float) + sizeof(uint64_t); - const size_t fields_count = operation + sizeof(uint32_t); - const size_t field = fields_count + sizeof(uint32_t); - const size_t tag = field + sizeof(uint32_t) + 5; - const auto maximum = std::numeric_limits::max(); - for (size_t offset : {size_t{0}, fields_count, field, tag + 1}) { - SCOPED_TRACE(offset); - auto invalid = valid; - std::memcpy(invalid.data() + offset, &maximum, sizeof(maximum)); - EXPECT_EQ(Doc::deserialize(invalid.data(), invalid.size()), nullptr); - } - auto invalid = valid; - std::memcpy(invalid.data() + operation, &maximum, sizeof(maximum)); - EXPECT_EQ(Doc::deserialize(invalid.data(), invalid.size()), nullptr); - invalid = valid; - invalid[tag] = 255; - EXPECT_EQ(Doc::deserialize(invalid.data(), invalid.size()), nullptr); - invalid = valid; - invalid.push_back(0); - EXPECT_EQ(Doc::deserialize(invalid.data(), invalid.size()), nullptr); - invalid = valid; - const uint32_t two = 2; - std::memcpy(invalid.data() + fields_count, &two, sizeof(two)); - invalid.insert(invalid.end(), valid.begin() + field, valid.end()); - EXPECT_EQ(Doc::deserialize(invalid.data(), invalid.size()), nullptr); - EXPECT_EQ(Doc::deserialize(nullptr, valid.size()), nullptr); - EXPECT_EQ(Doc::deserialize(valid.data(), 1), nullptr); -} - -TEST_F(DocDetailedTest, DeserializeChecksNestedCountsAndBooleanRepresentation) { - const size_t tag = sizeof(uint32_t) + 2 + sizeof(float) + sizeof(uint64_t) + - sizeof(uint32_t) + sizeof(uint32_t) + sizeof(uint32_t) + 1; - const std::vector values{ - std::vector{1, 2}, std::vector{"hello", "world"}, - std::vector{true, false}, - std::pair, std::vector>{{1, 2}, {1.f, 2.f}}, - std::pair, std::vector>{ - {1, 2}, {zvec::float16_t(1.f), zvec::float16_t(2.f)}}}; - for (const auto &value : values) { - Doc doc; - doc.set_pk("id"); - std::visit( - [&](const auto &item) { - if constexpr (std::is_same_v, - std::monostate>) { - doc.set_null("v"); - } else { - doc.set("v", item); - } - }, - value); - auto buffer = doc.serialize(); - ASSERT_NE(Doc::deserialize(buffer.data(), buffer.size()), nullptr); - // A UINT32_MAX element count cannot fit into any of these payloads. - std::fill_n(buffer.begin() + tag + 1, sizeof(uint32_t), uint8_t{0xff}); - EXPECT_EQ(Doc::deserialize(buffer.data(), buffer.size()), nullptr); - } - Doc sparse; - sparse.set_pk("id"); - sparse.set("v", std::pair, std::vector>{ - {1, 2}, {1.f, 2.f}}); - auto buffer = sparse.serialize(); - std::fill_n( - buffer.begin() + tag + 1 + sizeof(uint32_t) + 2 * sizeof(uint32_t), - sizeof(uint32_t), uint8_t{0xff}); - EXPECT_EQ(Doc::deserialize(buffer.data(), buffer.size()), nullptr); - Doc boolean; - boolean.set_pk("id"); - boolean.set("v", true); - buffer = boolean.serialize(); - buffer[tag + 1] = 2; - EXPECT_EQ(Doc::deserialize(buffer.data(), buffer.size()), nullptr); - boolean.set("v", std::vector{true}); - buffer = boolean.serialize(); - buffer[tag + 1 + sizeof(uint32_t)] = 2; - EXPECT_EQ(Doc::deserialize(buffer.data(), buffer.size()), nullptr); -} - TEST_F(DocDetailedTest, LongFieldNamesKeepDistinctIdentitiesAcrossSerialization) { const std::string prefix(32, 'f'); diff --git a/tests/db/index/common/schema_test.cc b/tests/db/index/common/schema_test.cc index b72f1bc8d..fdba5698b 100644 --- a/tests/db/index/common/schema_test.cc +++ b/tests/db/index/common/schema_test.cc @@ -44,24 +44,40 @@ TEST(CollectionSchemaTest, RejectsNullFieldObjectsInMutations) { } TEST(CollectionSchemaTest, ValidatesDuplicateNamesFromConstructorsAndCopies) { - CollectionSchema schema( - "schema", {std::make_shared("duplicate", DataType::INT32), - std::make_shared("duplicate", DataType::INT64)}); - // Retain the invalid input until explicit validation, rather than hiding one - // field and allowing a collection whose schema changes after reopening. - ASSERT_EQ(schema.fields().size(), 2u); - auto status = schema.validate(); - EXPECT_EQ(status.code(), StatusCode::INVALID_ARGUMENT); - EXPECT_EQ(status.message(), - "Invalid schema: duplicate field name [duplicate]; field names " - "must be unique"); - CollectionSchema copied(schema); - CollectionSchema assigned; - assigned = schema; - EXPECT_EQ(copied.validate(), status); - EXPECT_EQ(assigned.validate(), status); - EXPECT_EQ(copied.fields().size(), 2u); - EXPECT_EQ(assigned.fields().size(), 2u); + auto scalar = std::make_shared("duplicate", DataType::INT32); + auto other_scalar = + std::make_shared("duplicate", DataType::INT64); + auto vector = std::make_shared("duplicate", + DataType::VECTOR_FP32, 4, false); + auto other_vector = std::make_shared( + "duplicate", DataType::VECTOR_FP32, 8, false); + const std::vector cases{ + {scalar, std::make_shared(*scalar)}, + {scalar, other_scalar}, + {scalar, scalar}, + {vector, other_vector}, + {scalar, vector}, + {vector, scalar}, + {nullptr, scalar, nullptr, vector}}; + for (size_t i = 0; i < cases.size(); ++i) { + SCOPED_TRACE(i); + CollectionSchema schema("schema", cases[i]); + ASSERT_EQ(schema.fields().size(), 2u); + auto status = schema.validate(); + EXPECT_EQ(status.code(), StatusCode::INVALID_ARGUMENT); + EXPECT_EQ(status.message(), + "Invalid schema: duplicate field name [duplicate]; field names " + "must be unique"); + CollectionSchema copied(schema); + CollectionSchema assigned( + "old", {std::make_shared("existing", DataType::INT32)}); + assigned = schema; + EXPECT_EQ(copied.validate(), status); + EXPECT_EQ(assigned.validate(), status); + EXPECT_EQ(copied.fields().size(), 2u); + EXPECT_EQ(assigned.fields().size(), 2u); + EXPECT_FALSE(assigned.has_field("existing")); + } } TEST(CollectionSchemaTest, CopyOwnsIndependentFieldObjects) { diff --git a/tests/db/relaxed_validation_test.cc b/tests/db/relaxed_validation_test.cc index bbf642752..1515aa5dd 100644 --- a/tests/db/relaxed_validation_test.cc +++ b/tests/db/relaxed_validation_test.cc @@ -323,15 +323,64 @@ TEST_F(RelaxedValidationTest, ReservedNamesAndDuplicatesFailBeforeCreation) { EXPECT_NE(result.error().message().find("is reserved"), std::string::npos); EXPECT_FALSE(ailego::FileHelper::IsExist(path_.c_str())); } - CollectionSchema duplicate( - "x", {std::make_shared("value", DataType::INT32), - std::make_shared("value", DataType::INT64)}); - auto result = Collection::CreateAndOpen(path_, duplicate, options_); - ASSERT_FALSE(result.has_value()); - EXPECT_EQ(result.error().code(), StatusCode::INVALID_ARGUMENT); - EXPECT_NE(result.error().message().find("duplicate field name"), - std::string::npos); - EXPECT_FALSE(ailego::FileHelper::IsExist(path_.c_str())); + auto scalar = std::make_shared("value", DataType::INT32); + auto other_scalar = std::make_shared("value", DataType::INT64); + auto vector = + std::make_shared("value", DataType::VECTOR_FP32, 4, false); + auto other_vector = + std::make_shared("value", DataType::VECTOR_FP32, 8, false); + const std::vector cases{ + {scalar, std::make_shared(*scalar)}, + {scalar, other_scalar}, + {scalar, scalar}, + {vector, other_vector}, + {scalar, vector}, + {vector, scalar}}; + for (size_t i = 0; i < cases.size(); ++i) { + SCOPED_TRACE(i); + CollectionSchema duplicate("x", cases[i]); + auto result = Collection::CreateAndOpen(path_, duplicate, options_); + ASSERT_FALSE(result.has_value()); + EXPECT_EQ(result.error().code(), StatusCode::INVALID_ARGUMENT); + EXPECT_NE(result.error().message().find("duplicate field name [value]"), + std::string::npos); + EXPECT_FALSE(ailego::FileHelper::IsExist(path_.c_str())); + } +} + +TEST_F(RelaxedValidationTest, DuplicateDdlTargetsLeaveSchemaAndDataUnchanged) { + auto schema = MakeSchema(); + ASSERT_TRUE(schema + .add_field(std::make_shared( + "other", DataType::INT32, true)) + .ok()); + ASSERT_TRUE(schema + .add_field(std::make_shared( + "embedding", DataType::VECTOR_FP32, 4, true)) + .ok()); + ASSERT_NO_FATAL_FAILURE(Create(schema)); + std::vector docs{MakeDoc("id", 42)}; + ASSERT_TRUE(docs[0].set>("embedding", {1, 2, 3, 4})); + ASSERT_NO_FATAL_FAILURE(ExpectWrite(collection_->insert(docs), 1)); + const auto before = collection_->schema().value(); + for (const std::string name : {"other", "embedding"}) { + SCOPED_TRACE(name); + auto field = std::make_shared(name, DataType::INT32, true); + auto status = collection_->add_column(field, ""); + EXPECT_EQ(status.code(), StatusCode::INVALID_ARGUMENT); + EXPECT_NE(status.message().find("already exists"), std::string::npos); + status = collection_->alter_column("value", name); + EXPECT_EQ(status.code(), StatusCode::INVALID_ARGUMENT); + EXPECT_NE(status.message().find("already exists"), std::string::npos); + status = collection_->alter_column("value", "", field); + EXPECT_EQ(status.code(), StatusCode::ALREADY_EXISTS); + EXPECT_NE(status.message().find("already exists"), std::string::npos); + EXPECT_EQ(collection_->schema().value(), before); + ASSERT_NO_FATAL_FAILURE(ExpectValue("id", 42)); + } + ASSERT_NO_FATAL_FAILURE(Reopen()); + EXPECT_EQ(collection_->schema().value(), before); + ASSERT_NO_FATAL_FAILURE(ExpectValue("id", 42)); } TEST_F(RelaxedValidationTest, ReservedDdlTargetsLeaveTheCollectionUnchanged) { From 90db5d92bad86cc2a18f641a82df9e9d6853c16a Mon Sep 17 00:00:00 2001 From: Qinren Zhou Date: Wed, 16 Sep 2026 15:10:26 +0800 Subject: [PATCH 05/11] refactor --- src/db/index/common/doc.cc | 64 +-- .../index/common/manifest/manifest_codec.cc | 24 +- src/db/index/common/manifest_codec.h | 3 +- src/db/index/segment/segment.cc | 299 +++------- src/db/index/storage/wal/local_wal_file.cc | 240 ++++---- src/db/index/storage/wal/local_wal_file.h | 15 +- src/db/index/storage/wal/wal_file.h | 9 +- .../relaxed_validation_recovery_test.cc | 515 ------------------ tests/db/index/common/doc_test.cc | 20 +- .../common/manifest_codec_golden_test.cc | 32 -- tests/db/index/storage/wal_file_test.cc | 231 +++----- 11 files changed, 315 insertions(+), 1137 deletions(-) diff --git a/src/db/index/common/doc.cc b/src/db/index/common/doc.cc index 856b59efb..48742ba2f 100644 --- a/src/db/index/common/doc.cc +++ b/src/db/index/common/doc.cc @@ -739,7 +739,7 @@ Status Doc::validate_and_sanitize(const CollectionSchema::Ptr &schema, for (auto &[name, value] : fields_) { if (!schema->has_field(name)) { return Status::InvalidArgument( - "Invalid doc: doc[", format_name(pk_), "]: field[", format_name(name), + "Invalid doc[", format_name(pk_), "]: field[", format_name(name), "] does not exist in the collection schema"); } } @@ -752,16 +752,16 @@ Status Doc::validate_and_sanitize(const CollectionSchema::Ptr &schema, if (field_schema->nullable() || is_update) { continue; } - return Status::InvalidArgument("Invalid doc: doc[", format_name(pk_), - "]: field[", format_name(field_name), + return Status::InvalidArgument("Invalid doc[", format_name(pk_), + "]: field[", field_name, "] is required but not provided"); } else { if (std::holds_alternative(field_pair->second)) { if (field_schema->nullable()) { continue; } - return Status::InvalidArgument("Invalid doc: doc[", format_name(pk_), - "]: field[", format_name(field_name), + return Status::InvalidArgument("Invalid doc[", format_name(pk_), + "]: field[", field_name, "] is required but its value is null"); } } @@ -892,24 +892,21 @@ Status Doc::validate_and_sanitize(const CollectionSchema::Ptr &schema, field_value); if (sparse_values.size() != sparse_indices.size()) { return Status::InvalidArgument( - "Invalid doc: doc[", format_name(pk_), - "]: sparse vector field[", format_name(field_name), - "] has mismatched indices and values sizes"); + "Invalid doc[", format_name(pk_), "]: sparse vector field[", + field_name, "] has mismatched indices and values sizes"); } if (sparse_indices.size() > kSparseMaxDimSize) { return Status::InvalidArgument( - "Invalid doc: doc[", format_name(pk_), - "]: sparse vector field[", format_name(field_name), - "] exceeds the maximum number of sparse indices (", + "Invalid doc[", format_name(pk_), "]: sparse vector field[", + field_name, "] exceeds the maximum number of sparse indices (", kSparseMaxDimSize, ")"); } auto status = need_sanitize_sparse(sparse_indices.data(), sparse_indices.size()); if (status == SparseIndicesStatus::kHasDuplicate) { return Status::InvalidArgument( - "Invalid doc: doc[", format_name(pk_), - "]: sparse vector field[", format_name(field_name), - "] contains duplicate indices"); + "Invalid doc[", format_name(pk_), "]: sparse vector field[", + field_name, "] contains duplicate indices"); } if (status == SparseIndicesStatus::kNeedSort) { if (sort_and_find_duplicates( @@ -917,9 +914,8 @@ Status Doc::validate_and_sanitize(const CollectionSchema::Ptr &schema, reinterpret_cast(sparse_values.data()), sparse_indices.size(), sizeof(float16_t))) { return Status::InvalidArgument( - "Invalid doc: doc[", format_name(pk_), - "]: sparse vector field[", format_name(field_name), - "] contains duplicate indices"); + "Invalid doc[", format_name(pk_), "]: sparse vector field[", + field_name, "] contains duplicate indices"); } } } @@ -934,24 +930,21 @@ Status Doc::validate_and_sanitize(const CollectionSchema::Ptr &schema, field_value); if (sparse_values.size() != sparse_indices.size()) { return Status::InvalidArgument( - "Invalid doc: doc[", format_name(pk_), - "]: sparse vector field[", format_name(field_name), - "] has mismatched indices and values sizes"); + "Invalid doc[", format_name(pk_), "]: sparse vector field[", + field_name, "] has mismatched indices and values sizes"); } if (sparse_indices.size() > kSparseMaxDimSize) { return Status::InvalidArgument( - "Invalid doc: doc[", format_name(pk_), - "]: sparse vector field[", format_name(field_name), - "] exceeds the maximum number of sparse indices (", + "Invalid doc[", format_name(pk_), "]: sparse vector field[", + field_name, "] exceeds the maximum number of sparse indices (", kSparseMaxDimSize, ")"); } auto status = need_sanitize_sparse(sparse_indices.data(), sparse_indices.size()); if (status == SparseIndicesStatus::kHasDuplicate) { return Status::InvalidArgument( - "Invalid doc: doc[", format_name(pk_), - "]: sparse vector field[", format_name(field_name), - "] contains duplicate indices"); + "Invalid doc[", format_name(pk_), "]: sparse vector field[", + field_name, "] contains duplicate indices"); } if (status == SparseIndicesStatus::kNeedSort) { if (sort_and_find_duplicates( @@ -959,34 +952,33 @@ Status Doc::validate_and_sanitize(const CollectionSchema::Ptr &schema, reinterpret_cast(sparse_values.data()), sparse_indices.size(), sizeof(float))) { return Status::InvalidArgument( - "Invalid doc: doc[", format_name(pk_), - "]: sparse vector field[", format_name(field_name), - "] contains duplicate indices"); + "Invalid doc[", format_name(pk_), "]: sparse vector field[", + field_name, "] contains duplicate indices"); } } } break; } default: - return Status::InvalidArgument("Invalid doc: doc[", format_name(pk_), - "]: field[", format_name(field_name), + return Status::InvalidArgument("Invalid doc[", format_name(pk_), + "]: field[", field_name, "] has unsupported data type"); break; } if (!type_match) { return Status::InvalidArgument( - "Invalid doc: doc[", format_name(pk_), "]: field[", - format_name(field_name), "] type mismatch, expected ", + "Invalid doc[", format_name(pk_), "]: field[", field_name, + "] type mismatch, expected ", DataTypeCodeBook::AsString(expected_type), " but got ", get_value_type_name(field_value, field_schema->is_vector_field())); } if (field_schema->is_dense_vector()) { if (value_dimension != field_schema->dimension()) { return Status::InvalidArgument( - "Invalid doc: doc[", format_name(pk_), "]: field[", - format_name(field_name), "] dimension mismatch, expected ", - field_schema->dimension(), " but got ", value_dimension); + "Invalid doc[", format_name(pk_), "]: field[", field_name, + "] dimension mismatch, expected ", field_schema->dimension(), + " but got ", value_dimension); } } } diff --git a/src/db/index/common/manifest/manifest_codec.cc b/src/db/index/common/manifest/manifest_codec.cc index 1b22ef5ba..06af8e43b 100644 --- a/src/db/index/common/manifest/manifest_codec.cc +++ b/src/db/index/common/manifest/manifest_codec.cc @@ -803,7 +803,7 @@ void ManifestCodec::EncodeCollectionSchema(const CollectionSchema &schema, schema.max_doc_count_per_segment()); } -Result ManifestCodec::DecodeCollectionSchema( +CollectionSchema::Ptr ManifestCodec::DecodeCollectionSchema( std::string_view buf) { auto schema = std::make_shared(); // The protobuf-based converter read max_doc_count_per_segment straight from @@ -816,14 +816,9 @@ Result ManifestCodec::DecodeCollectionSchema( case f_collection::kName: schema->set_name(r.string_value()); break; - case f_collection::kFields: { - auto status = schema->add_field(DecodeFieldSchema(r.bytes())); - if (!status.ok()) { - return tl::make_unexpected(Status::InternalError( - "Malformed manifest schema: ", status.message())); - } + case f_collection::kFields: + schema->add_field(DecodeFieldSchema(r.bytes())); break; - } case f_collection::kMaxDocCountPerSegment: schema->set_max_doc_count_per_segment(r.varint()); break; @@ -831,10 +826,6 @@ Result ManifestCodec::DecodeCollectionSchema( break; } } - if (!r.ok()) { - return tl::make_unexpected( - Status::InternalError("Malformed manifest schema")); - } return schema; } @@ -961,14 +952,9 @@ Status ManifestCodec::Decode(std::string_view buf, ManifestData *data) { case f_manifest::kVersion: data->version = r.uint32_value(); break; - case f_manifest::kSchema: { - auto schema = DecodeCollectionSchema(r.bytes()); - if (!schema.has_value()) { - return schema.error(); - } - data->schema = std::move(schema).value(); + case f_manifest::kSchema: + data->schema = DecodeCollectionSchema(r.bytes()); break; - } case f_manifest::kEnableMmap: data->enable_mmap = r.bool_value(); break; diff --git a/src/db/index/common/manifest_codec.h b/src/db/index/common/manifest_codec.h index c6a22b86c..4dc7c3154 100644 --- a/src/db/index/common/manifest_codec.h +++ b/src/db/index/common/manifest_codec.h @@ -67,8 +67,7 @@ struct ManifestCodec { static void EncodeCollectionSchema(const CollectionSchema &schema, std::string *out); - static Result DecodeCollectionSchema( - std::string_view buf); + static CollectionSchema::Ptr DecodeCollectionSchema(std::string_view buf); static void EncodeBlockMeta(const BlockMeta &meta, std::string *out); static BlockMeta::Ptr DecodeBlockMeta(std::string_view buf); diff --git a/src/db/index/segment/segment.cc b/src/db/index/segment/segment.cc index 26742c8d7..efc27f758 100644 --- a/src/db/index/segment/segment.cc +++ b/src/db/index/segment/segment.cc @@ -317,12 +317,10 @@ class SegmentImpl : public Segment, Status insert_vector_indexer(Doc &doc); Status internal_insert(Doc &doc); Status internal_update(Doc &doc); + Status internal_upsert(Doc &doc); Status internal_delete(const Doc &doc); Status recover(); - Result> find_legacy_upsert_predecessors( - const std::unordered_set &keys, - uint64_t first_replay_id) const; Status open_wal_file(); Status append_wal(const Doc &doc); Status update_version(uint32_t delete_snapshot_path_suffix); @@ -952,6 +950,15 @@ Status SegmentImpl::internal_update(Doc &doc) { return internal_insert(doc); } +Status SegmentImpl::internal_upsert(Doc &doc) { + uint64_t g_doc_id; + bool exist = id_map_->has(doc.pk_ref(), &g_doc_id); + if (exist) { + delete_store_->mark_deleted(g_doc_id); + } + return internal_insert(doc); +} + Status SegmentImpl::internal_delete(const Doc &doc) { delete_store_->mark_deleted(doc.doc_id()); id_map_->remove(doc.pk_ref()); @@ -996,28 +1003,13 @@ Status SegmentImpl::Update(Doc &doc) { Status SegmentImpl::Upsert(Doc &doc) { std::lock_guard lock(seg_mtx_); - // Persist the predecessor explicitly. RocksDB may flush the new ID mapping - // before the deletion snapshot is committed, so recovery cannot safely - // infer the superseded document from the current ID map. - const auto original_doc_id = doc.doc_id(); - uint64_t previous_id; - const bool exists = id_map_->has(doc.pk_ref(), &previous_id); - if (exists) { - doc.set_doc_id(previous_id); - doc.set_operator(Operator::UPDATE); - } else { - doc.set_operator(Operator::INSERT); - } - - auto status = append_wal(doc); - // Preserve the public operation on the caller's document; only the WAL - // uses the already-supported INSERT/UPDATE representation. doc.set_operator(Operator::UPSERT); - if (!status.ok()) { - doc.set_doc_id(original_doc_id); - return status; - } - return exists ? internal_update(doc) : internal_insert(doc); + + // append WAL + auto s = append_wal(doc); + CHECK_RETURN_STATUS(s); + + return internal_upsert(doc); } Status SegmentImpl::Delete(const std::string &pk) { @@ -4274,82 +4266,10 @@ Status SegmentImpl::init_memory_components() { return Status::OK(); } -Result> SegmentImpl::find_legacy_upsert_predecessors( - const std::unordered_set &keys, - uint64_t first_replay_id) const { - std::vector predecessors; - const auto version = version_manager_->get_current_version(); - auto segments = version.persisted_segment_metas(); - if (auto writing = version.writing_segment_meta()) { - segments.push_back(std::move(writing)); - } - for (const auto &segment : segments) { - for (const auto &block : segment->persisted_blocks()) { - if (block.type() != BlockType::SCALAR || - !block.contain_column(GLOBAL_DOC_ID) || - !block.contain_column(USER_ID)) { - continue; - } - const auto forward_path = FileHelper::MakeForwardBlockPath( - path_, segment->id(), block.id(), !options_.enable_mmap_); - BaseForwardStore::Ptr store; - // BufferPoolForwardStore only supports Parquet; IPC always uses the - // mapped store. Release each store and its buffers before the next block. - if (options_.enable_mmap_ || - InferFileFormat(forward_path) == FileFormat::IPC) { - store = std::make_shared(forward_path); - } else { - store = std::make_shared(forward_path); - } - auto status = store->Open(); - if (!status.ok()) { - return tl::make_unexpected(Status::InternalError( - "Failed to open committed rows for legacy WAL recovery: path[", - forward_path, "], reason[", status.message(), "]")); - } - auto reader = store->scan({GLOBAL_DOC_ID, USER_ID}); - if (!reader) { - return tl::make_unexpected(Status::InternalError( - "Failed to scan committed rows for legacy WAL recovery: ", - forward_path)); - } - while (true) { - std::shared_ptr batch; - const auto read_status = reader->ReadNext(&batch); - if (!read_status.ok()) { - return tl::make_unexpected(Status::InternalError( - "Failed to read committed rows for legacy WAL recovery: path[", - forward_path, "], reason[", read_status.message(), "]")); - } - if (!batch) break; - const auto ids = std::dynamic_pointer_cast( - batch->GetColumnByName(GLOBAL_DOC_ID)); - const auto pks = std::dynamic_pointer_cast( - batch->GetColumnByName(USER_ID)); - if (!ids || !pks || ids->length() != pks->length() || - ids->null_count() != 0 || pks->null_count() != 0) { - return tl::make_unexpected(Status::InternalError( - "Invalid committed identity columns during legacy WAL recovery: ", - forward_path)); - } - for (int64_t row = 0; row < ids->length(); ++row) { - const auto doc_id = ids->Value(row); - if (doc_id < first_replay_id && !delete_store_->is_deleted(doc_id) && - keys.find(pks->GetString(row)) != keys.end()) { - predecessors.push_back(doc_id); - } - } - } - } - } - return predecessors; -} - Status SegmentImpl::recover() { // recover mem block meta auto &mem_block = segment_meta_->writing_forward_block().value(); - const auto first_replay_id = mem_block.min_doc_id(); - doc_id_allocator_.store(first_replay_id); + doc_id_allocator_.store(mem_block.min_doc_id()); std::string wal_file_path = FileHelper::MakeWalPath(path_, segment_meta_->id(), mem_block.id_); @@ -4366,142 +4286,100 @@ Status SegmentImpl::recover() { 0) { LOG_ERROR("WAL recovery failed: unable to open WAL file [%s]", wal_file_path.c_str()); - return Status::InternalError("Failed to open WAL for recovery: ", - wal_file_path); + return Status::OK(); } std::array(Operator::DELETE) + 1> recovered_doc_count{}; uint64_t total_recovered_doc_count{0}; - std::unordered_set legacy_upsert_keys; + + int ret = recover_wal_file->prepare_for_read(); + if (ret != 0) { + LOG_ERROR( + "WAL recovery failed: unable to prepare file for reading, path[%s], " + "segment[%d], ret[%d]", + wal_file_path.c_str(), id(), ret); + return Status::InternalError( + "Failed to prepare WAL file for reading: path[", wal_file_path, + "], segment[", id(), "], ret[", ret, "]"); + } LOG_INFO("WAL recovery started: path[%s], segment[%d]", wal_file_path.c_str(), id()); std::lock_guard lock(seg_mtx_); - // Validate the complete stream before changing the ID map or indexes. A - // failed open can close and flush those stores, so discovering corruption - // after applying a prefix would otherwise persist a partial recovery. - // Read one WAL record at a time. Legacy recovery additionally holds unique - // UPSERT keys, matched predecessor IDs, and the existing forward-store - // buffers for one committed block. New WAL records need no committed scan. - for (int pass = 0; pass < 2; ++pass) { - const bool replay = pass == 1; - total_recovered_doc_count = 0; - int ret = recover_wal_file->prepare_for_read(); - if (ret != 0) { + while (true) { + std::string buf = recover_wal_file->next(); + if (buf.empty()) { + break; + } + total_recovered_doc_count++; + auto doc = Doc::deserialize(reinterpret_cast(buf.data()), + buf.size()); + if (doc == nullptr) { LOG_ERROR( - "WAL recovery failed: unable to prepare file for reading, path[%s], " - "segment[%d], ret[%d]", - wal_file_path.c_str(), id(), ret); - return Status::InternalError( - "Failed to prepare WAL file for reading: path[", wal_file_path, - "], segment[", id(), "], ret[", ret, "]"); - } - - if (replay && !legacy_upsert_keys.empty()) { - // Old UPSERT records did not include their predecessor ID. Reconstruct - // it from committed rows even if an interrupted replay already replaced - // its ID-map entry. Finish the entire scan before changing tombstones. - auto predecessors = - find_legacy_upsert_predecessors(legacy_upsert_keys, first_replay_id); - if (!predecessors.has_value()) return predecessors.error(); - for (const auto doc_id : predecessors.value()) { - delete_store_->mark_deleted(doc_id); - } + "WAL record recovery failed: path[%s], segment[%d], record[%zu], " + "reason[deserialization failed]", + wal_file_path.c_str(), id(), (size_t)total_recovered_doc_count); + continue; } - while (true) { - auto record = recover_wal_file->next(); - if (!record.has_value()) { - return Status::InternalError( - "Failed to read WAL during recovery: path[", wal_file_path, - "], segment[", id(), "], reason[", record.error().message(), "]"); - } - if (!record.value().has_value()) { + Status status; + switch (doc->get_operator()) { + case Operator::INSERT: { + internal_insert(*doc); break; } - if (options_.read_only_) { - return Status::FailedPrecondition( - "WAL recovery is required; open the collection in read-write mode " - "once to recover before opening it read-only"); - } - const auto &buf = record.value().value(); - total_recovered_doc_count++; - auto doc = Doc::deserialize(reinterpret_cast(buf.data()), - buf.size()); - if (doc == nullptr) { - LOG_ERROR( - "WAL record recovery failed: path[%s], segment[%d], record[%zu], " - "reason[deserialization failed]", - wal_file_path.c_str(), id(), (size_t)total_recovered_doc_count); - return Status::InternalError( - "Corrupt WAL document: path[", wal_file_path, "], segment[", id(), - "], record[", total_recovered_doc_count, "]"); + case Operator::UPDATE: { + internal_update(*doc); + break; } - - if (!replay) { - if (doc->get_operator() == Operator::UPSERT) { - legacy_upsert_keys.insert(doc->pk_ref()); - } - continue; + case Operator::UPSERT: { + internal_upsert(*doc); + break; } - - Status status; - switch (doc->get_operator()) { - case Operator::INSERT: { - status = internal_insert(*doc); - break; - } - case Operator::UPDATE: { - status = internal_update(*doc); - break; - } - case Operator::UPSERT: { - // A previous interrupted replay may already have persisted this - // record's ID (or a later ID for the same key) in RocksDB. Only an - // older document is superseded; marking this/later replay ID deleted - // would hide a successfully recovered document on retry. - uint64_t previous_id; - if (id_map_->has(doc->pk_ref(), &previous_id) && - previous_id < doc_id_allocator_.load()) { - delete_store_->mark_deleted(previous_id); - } - status = internal_insert(*doc); - break; - } - case Operator::DELETE: { - status = internal_delete(*doc); - break; - } - default: - LOG_ERROR( - "WAL record recovery failed: path[%s], segment[%d], record[%zu], " - "operator[%d], reason[unknown operator]", - wal_file_path.c_str(), id(), (size_t)total_recovered_doc_count, - static_cast(doc->get_operator())); - return Status::InternalError("Unknown WAL document operator: path[", - wal_file_path, "], record[", - total_recovered_doc_count, "]"); + case Operator::DELETE: { + internal_delete(*doc); + break; } - - if (!status.ok()) { + default: LOG_ERROR( "WAL record recovery failed: path[%s], segment[%d], record[%zu], " - "operator[%d], reason[%s]", + "operator[%d], reason[unknown operator]", wal_file_path.c_str(), id(), (size_t)total_recovered_doc_count, - static_cast(doc->get_operator()), status.message().c_str()); - return Status(status.code(), - ailego::StringHelper::Concat( - "Failed to apply WAL record: path[", wal_file_path, - "], record[", total_recovered_doc_count, "], reason[", - status.message(), "]")); - } + static_cast(doc->get_operator())); + break; + } - recovered_doc_count[static_cast(doc->get_operator())]++; + if (!status.ok()) { + LOG_ERROR( + "WAL record recovery failed: path[%s], segment[%d], record[%zu], " + "operator[%d], reason[%s]", + wal_file_path.c_str(), id(), (size_t)total_recovered_doc_count, + static_cast(doc->get_operator()), status.message().c_str()); + continue; } + + recovered_doc_count[static_cast(doc->get_operator())]++; } + const auto added_docs = recovered_doc_count[0] + // INSERT + recovered_doc_count[1] + // UPSERT + recovered_doc_count[2]; // UPDATE + mem_block.max_doc_id_ += added_docs; + + ret = recover_wal_file->close(); + if (ret != 0) { + LOG_ERROR( + "WAL recovery failed: unable to close file, path[%s], " + "segment[%d], ret[%d]", + wal_file_path.c_str(), id(), ret); + return Status::InternalError("Failed to close recovered WAL file: path[", + wal_file_path, "], segment[", id(), "], ret[", + ret, "]"); + } + recover_wal_file.reset(); + LOG_INFO( "WAL recovery completed: path[%s], segment[%d], total[%zu], " "insert[%zu], upsert[%zu], update[%zu], delete[%zu]", @@ -4516,10 +4394,7 @@ Status SegmentImpl::recover() { // optimize() flush the writing segment before sealing it; without an open // member WAL, flush() treats the recovered memory components as empty and // returns without persisting them. - // Retain the reader's valid-tail position so a later append can discard an - // incomplete crash record. Opening read-only does not truncate the WAL. - wal_file_ = std::move(recover_wal_file); - return Status::OK(); + return open_wal_file(); } Status SegmentImpl::open_wal_file() { diff --git a/src/db/index/storage/wal/local_wal_file.cc b/src/db/index/storage/wal/local_wal_file.cc index e9d0493c5..6a50e5c34 100644 --- a/src/db/index/storage/wal/local_wal_file.cc +++ b/src/db/index/storage/wal/local_wal_file.cc @@ -13,9 +13,6 @@ // limitations under the License. #include "local_wal_file.h" -#include -#include -#include #ifndef _MSC_VER #include #endif @@ -25,73 +22,49 @@ #include "db/common/file_helper.h" #include "db/common/typedef.h" +#define MAX_RECORD_SIZE 4194304 // 4Mb + namespace zvec { int LocalWalFile::append(std::string &&data) { - if (data.empty() || data.size() > std::numeric_limits::max()) { - WLOG_ERROR("Wal record length is not representable: %zu", data.size()); - return -1; - } - WalRecord record; - record.length_ = static_cast(data.size()); - record.crc_ = ailego::Crc32c::Hash(data.data(), data.size(), 0); - record.content_ = std::move(data); + record.length_ = data.size(); + record.crc_ = ailego::Crc32c::Hash( + reinterpret_cast(data.data()), record.length_, 0); + record.content_ = std::forward(data); - std::lock_guard lock(file_mutex_); - if (!opened_ || failed_) { - return -1; - } - if (incomplete_tail_offset_) { - if (!file_.truncate(*incomplete_tail_offset_)) { - WLOG_ERROR("Wal incomplete tail truncation failed"); - failed_ = true; - return -1; - } - incomplete_tail_offset_.reset(); - } - if (!file_.seek(0, ailego::File::Origin::End)) { - return -1; - } if (write_record(record) < 0) { + WLOG_ERROR("Wal write record error. record.length_[%zu]", + (size_t)record.length_); return -1; } - // Keep the flush counter and flush in the same critical section as writes. + // if max_docs_wal_flush_ is 0, no need flush if (max_docs_wal_flush_ != 0 && docs_count_ >= max_docs_wal_flush_) { if (!file_.flush()) { WLOG_ERROR("Wal flush error. docs_count_[%zu] max_docs_wal_flush_[%zu]", (size_t)docs_count_, (size_t)max_docs_wal_flush_); - failed_ = true; - return -1; } docs_count_ = 0; } return 0; } -Result> LocalWalFile::next() { - std::lock_guard lock(file_mutex_); - if (!opened_ || failed_) { - return tl::make_unexpected( - Status::InternalError("WAL is not open for reading or has failed")); - } +std::string LocalWalFile::next() { WalRecord record; - auto result = read_record(record); - if (!result.has_value()) { - failed_ = true; - return tl::make_unexpected(result.error()); - } - if (!result.value()) { - return std::nullopt; - } - const uint32_t crc = - ailego::Crc32c::Hash(record.content_.data(), record.content_.size(), 0); - if (crc != record.crc_) { - failed_ = true; - return tl::make_unexpected( - Status::InternalError("WAL record CRC mismatch")); + if (read_record(record) > 0) { + uint32_t tmp_crc = ailego::Crc32c::Hash( + reinterpret_cast(record.content_.data()), record.length_, + 0); + if (tmp_crc == record.crc_) { + return std::move(record.content_); + } else { + WLOG_ERROR( + "Wal next error. record.length_[%zu] crc_[%zu] != tmp_crc[%zu]", + (size_t)record.length_, (size_t)record.crc_, (size_t)tmp_crc); + } } - return std::optional(std::move(record.content_)); + // end of file or read error + return std::string(); } int LocalWalFile::open(const WalOptions &wal_option) { @@ -109,7 +82,7 @@ int LocalWalFile::open(const WalOptions &wal_option) { } // write wal header - size_t write_size = file_.write((const void *)&header_, sizeof(header_)); + int write_size = file_.write((const void *)&header_, sizeof(header_)); if (write_size != sizeof(header_)) { WLOG_ERROR("Wal write header error. create_new[%d]", wal_option.create_new); @@ -129,16 +102,11 @@ int LocalWalFile::open(const WalOptions &wal_option) { } // open default for write - if (!file_.seek(0, ailego::File::Origin::End)) { - return -1; - } + file_.seek(0, ailego::File::Origin::End); } max_docs_wal_flush_ = wal_option.max_docs_wal_flush; opened_ = true; - failed_ = false; - incomplete_tail_offset_.reset(); - docs_count_ = 0; WLOG_INFO("Wal open success. create_new[%d]", wal_option.create_new); return 0; @@ -174,98 +142,106 @@ int LocalWalFile::flush() { int LocalWalFile::prepare_for_read() { CHECK_STATUS(opened_, true); - incomplete_tail_offset_.reset(); - if (failed_ || !file_.seek(0, ailego::File::Origin::Begin)) { + if (!file_.seek(0, ailego::File::Origin::Begin)) { return -1; } - size_t read_size = file_.read((void *)&header_, sizeof(header_)); + int read_size = file_.read((void *)&header_, sizeof(header_)); if (read_size != sizeof(header_)) { WLOG_ERROR("Wal read header error."); - failed_ = true; return -1; } if (header_.wal_version != 0UL) { WLOG_ERROR("Wal version not support error."); - failed_ = true; return -1; } return 0; } -// Caller holds file_mutex_. A failed write must not strand future successful -// appends behind its incomplete record. +//! Return 1 if success or -1 if write error int LocalWalFile::write_record(WalRecord &record) { - const auto start = file_.offset(); - if (start < static_cast(sizeof(header_))) { - failed_ = true; - return -1; - } - if (file_.write(&record.length_, LENGTH_SIZE) != LENGTH_SIZE || - file_.write(&record.crc_, CRC_SIZE) != CRC_SIZE || - file_.write(record.content_.data(), record.content_.size()) != - record.content_.size()) { - WLOG_ERROR("Wal write record failed. record.length_[%zu]", - record.content_.size()); - if (!file_.truncate(static_cast(start)) || - !file_.seek(start, ailego::File::Origin::Begin)) { - failed_ = true; + CHECK_STATUS(opened_, true); + + int write_size = 0; + int ret = -1; + + std::lock_guard lock(file_mutex_); + do { + write_size = file_.write((const void *)&record.length_, LENGTH_SIZE); + if (write_size != LENGTH_SIZE) { + WLOG_ERROR("Wal write error. record.length_ error write_size[%d]", + write_size); + break; } - return -1; - } - ++docs_count_; - return 1; + + write_size = file_.write((const void *)&record.crc_, CRC_SIZE); + if (write_size != CRC_SIZE) { + WLOG_ERROR("Wal write error. record.crc_ error write_size[%d]", + write_size); + break; + } + + write_size = + file_.write((const void *)record.content_.data(), record.length_); + if (write_size != (int)record.length_) { + WLOG_ERROR("Wal write error. record.content_ error write_size[%d]", + write_size); + break; + } + ret = 1; // write one record success + docs_count_++; + } while (false); + + return ret; } -Result LocalWalFile::read_record(WalRecord &record) { - if (incomplete_tail_offset_) { - return false; - } - // File::read reports bytes read for both EOF and I/O failures. Check the - // physical extent first: a short read within that extent is an I/O error, - // whereas a final frame that does not fit is a tolerated interrupted write. - const auto start = file_.offset(); - const size_t file_size = file_.size(); - if (!file_.is_valid() || start < static_cast(sizeof(header_)) || - file_size < sizeof(header_) || static_cast(start) > file_size) { - return tl::make_unexpected( - Status::InternalError("Failed to determine WAL read position or size")); - } - const size_t remaining = file_size - static_cast(start); - if (remaining == 0) { - return false; - } - if (remaining < LENGTH_SIZE + CRC_SIZE) { - incomplete_tail_offset_ = static_cast(start); - return false; - } - if (file_.read(&record.length_, LENGTH_SIZE) != LENGTH_SIZE || - file_.read(&record.crc_, CRC_SIZE) != CRC_SIZE) { - return tl::make_unexpected( - Status::InternalError("Failed to read WAL record header")); - } - if (record.length_ == 0) { - return tl::make_unexpected( - Status::InternalError("WAL record has zero length")); - } - if (record.length_ > remaining - LENGTH_SIZE - CRC_SIZE) { - incomplete_tail_offset_ = static_cast(start); - return false; - } - try { +//! Return 1 if success or 0 if eof or -1 if read error +int LocalWalFile::read_record(WalRecord &record) { + CHECK_STATUS(opened_, true); + + int read_size = 0; + std::string err_msg; + int ret = -1; + + do { + read_size = + file_.read(reinterpret_cast(&record.length_), LENGTH_SIZE); + if (read_size == 0) { + ret = 0; + WLOG_INFO("Wal read finished. end of file"); + break; + } + + if (read_size != LENGTH_SIZE) { + WLOG_ERROR("Wal read error. record.length_ error read_size[%d]", + read_size); + break; + } + + read_size = file_.read(reinterpret_cast(&record.crc_), CRC_SIZE); + if (read_size != CRC_SIZE) { + WLOG_ERROR("Wal read error. record.crc_ error read_size[%d]", read_size); + break; + } + + // resize may crash if record.length_ very large + if (record.length_ <= 0 || record.length_ > MAX_RECORD_SIZE) { + WLOG_ERROR("Wal read error. record.length_ value error read_size[%d]", + read_size); + break; + } + record.content_.resize(record.length_); - } catch (const std::bad_alloc &) { - return tl::make_unexpected(Status(StatusCode::RESOURCE_EXHAUSTED, - "Unable to allocate WAL record buffer")); - } catch (const std::length_error &) { - return tl::make_unexpected(Status(StatusCode::RESOURCE_EXHAUSTED, - "WAL record exceeds string capacity")); - } - if (file_.read(record.content_.data(), record.content_.size()) != - record.content_.size()) { - return tl::make_unexpected( - Status::InternalError("Failed to read WAL record payload")); - } - return true; + read_size = file_.read((void *)const_cast(record.content_.data()), + record.length_); + if (read_size != (int)record.length_) { + WLOG_ERROR("Wal read error. record.content_ error read_size[%d]", + read_size); + break; + } + ret = 1; // read one record success + } while (false); + + return ret; } -} // namespace zvec +}; // namespace zvec \ No newline at end of file diff --git a/src/db/index/storage/wal/local_wal_file.h b/src/db/index/storage/wal/local_wal_file.h index c2a8dcadf..d4f392fa8 100644 --- a/src/db/index/storage/wal/local_wal_file.h +++ b/src/db/index/storage/wal/local_wal_file.h @@ -14,7 +14,12 @@ #pragma once #include +#include +#include +#include #include +#include +#include #include #include "wal_file.h" @@ -25,7 +30,7 @@ namespace zvec { */ struct WalHeader { uint64_t wal_version{0U}; - uint64_t reserved_[7]{}; + uint64_t reserved_[7]; }; static_assert(sizeof(WalHeader) % 64 == 0, @@ -56,7 +61,7 @@ class LocalWalFile : public WalFile { public: int append(std::string &&data) override; int prepare_for_read() override; - Result> next() override; + std::string next() override; public: int open(const WalOptions &wal_option) override; @@ -73,7 +78,7 @@ class LocalWalFile : public WalFile { private: int write_record(WalRecord &record); - Result read_record(WalRecord &record); + int read_record(WalRecord &record); private: ailego::File file_; @@ -88,10 +93,6 @@ class LocalWalFile : public WalFile { WalHeader header_; bool opened_{false}; - bool failed_{false}; - // Preserve the complete prefix and remove a torn final record before the - // next append. Merely reading a WAL must not modify it. - std::optional incomplete_tail_offset_; }; diff --git a/src/db/index/storage/wal/wal_file.h b/src/db/index/storage/wal/wal_file.h index 34c2570b3..6e17f7384 100644 --- a/src/db/index/storage/wal/wal_file.h +++ b/src/db/index/storage/wal/wal_file.h @@ -13,11 +13,8 @@ // limitations under the License. #pragma once -#include #include -#include #include -#include namespace zvec { @@ -49,9 +46,7 @@ class WalFile { public: virtual int append(std::string &&data) = 0; virtual int prepare_for_read() = 0; - // A successful empty optional means EOF or an incomplete final crash record. - // Read failures and complete but corrupt records return an error. - virtual Result> next() = 0; + virtual std::string next() = 0; public: //! Open and initialize WalFile @@ -69,4 +64,4 @@ class WalFile { virtual bool has_record() = 0; }; -}; // namespace zvec +}; // namespace zvec \ No newline at end of file diff --git a/tests/db/crash_recovery/relaxed_validation_recovery_test.cc b/tests/db/crash_recovery/relaxed_validation_recovery_test.cc index ddd9ebf54..b29d5f07e 100644 --- a/tests/db/crash_recovery/relaxed_validation_recovery_test.cc +++ b/tests/db/crash_recovery/relaxed_validation_recovery_test.cc @@ -14,193 +14,17 @@ #include #include -#include -#include -#include -#include -#include -#include -#include #include -#include #include #include #include #include #include #include -#include "db/index/common/id_map.h" -#include "db/index/common/version_manager.h" -#include "db/index/storage/wal/wal_file.h" namespace zvec { namespace { -std::string ReadFileBytes(const std::string &path) { - std::ifstream file(path, std::ios::binary); - return std::string(std::istreambuf_iterator(file), {}); -} - -std::string FindWal(const std::string &path) { - for (const auto &entry : - std::filesystem::recursive_directory_iterator(path)) { - if (entry.path().extension() == ".wal") return entry.path().string(); - } - return {}; -} - -std::map ReadManifests(const std::string &path) { - std::map result; - for (const auto &entry : - std::filesystem::recursive_directory_iterator(path)) { - if (entry.path().filename().string().rfind("manifest", 0) == 0) { - result.emplace(entry.path().string(), - ReadFileBytes(entry.path().string())); - } - } - return result; -} - -// Only called in ASSERT_EXIT children. Intentionally skip collection cleanup. -void WriteStringDocsAndExit(const std::string &path, - const std::vector &ids, - const std::vector &values) { - CollectionSchema schema("wal_recovery"); - if (!schema - .add_field( - std::make_shared("text", DataType::STRING, false)) - .ok()) { - std::_Exit(1); - } - auto created = Collection::CreateAndOpen(path, schema, CollectionOptions{}); - if (!created.has_value()) std::_Exit(2); - auto collection = std::move(created).value(); - std::vector docs; - for (size_t i = 0; i < ids.size(); ++i) { - Doc doc; - doc.set_pk(ids[i]); - doc.set("text", values[i]); - docs.push_back(std::move(doc)); - } - auto result = collection->insert(docs); - if (!result.has_value()) std::_Exit(3); - for (const auto &status : result.value()) { - if (!status.ok()) std::_Exit(4); - } - std::_Exit(0); -} - -std::string FindIdMap(const std::string &path) { - for (const auto &entry : std::filesystem::directory_iterator(path)) { - if (entry.is_directory() && - entry.path().filename().string().rfind("idmap", 0) == 0) { - return entry.path().string(); - } - } - return {}; -} - -void WriteUpsertsAndExit(const std::string &path, bool persisted_base, - bool append_suffix, bool separate_segment = false) { - CollectionSchema schema("wal_recovery"); - if (!schema - .add_field( - std::make_shared("text", DataType::STRING, false)) - .ok()) { - std::_Exit(1); - } - auto created = Collection::CreateAndOpen(path, schema, CollectionOptions{}); - if (!created.has_value()) std::_Exit(2); - auto collection = std::move(created).value(); - Doc doc; - doc.set_pk("target"); - if (persisted_base) { - doc.set("text", "original"); - std::vector docs{doc}; - auto inserted = collection->insert(docs); - if (!inserted.has_value() || !inserted.value().front().ok() || - !collection->flush().ok()) - std::_Exit(3); - if (separate_segment && !collection->optimize().ok()) std::_Exit(6); - } - for (const auto &value : {"first", "second"}) { - doc.set("text", value); - std::vector docs{doc}; - auto updated = collection->upsert(docs); - if (!updated.has_value() || !updated.value().front().ok()) std::_Exit(4); - } - if (append_suffix) { - doc.set_pk("broken"); - std::vector docs{doc}; - auto inserted = collection->insert(docs); - if (!inserted.has_value() || !inserted.value().front().ok()) std::_Exit(5); - } - std::_Exit(0); -} - -void ReadWalDocuments(const std::string &path, std::vector *docs) { - const auto wal_path = FindWal(path); - ASSERT_FALSE(wal_path.empty()); - auto wal = WalFile::Create(wal_path); - ASSERT_EQ(wal->open(WalOptions{}), 0); - ASSERT_EQ(wal->prepare_for_read(), 0); - while (true) { - auto record = wal->next(); - ASSERT_TRUE(record.has_value()) << record.error().message(); - if (!record.value().has_value()) break; - const auto &bytes = record.value().value(); - auto doc = Doc::deserialize(reinterpret_cast(bytes.data()), - bytes.size()); - ASSERT_NE(doc, nullptr); - docs->push_back(std::move(doc)); - } - ASSERT_EQ(wal->close(), 0); -} - -// Recreate historical UPSERT records through the serializer and WAL writer so -// the framing and checksums remain valid. New public writes use INSERT/UPDATE. -void RewriteWalAsLegacyUpserts(const std::string &path) { - std::vector docs; - ASSERT_NO_FATAL_FAILURE(ReadWalDocuments(path, &docs)); - ASSERT_FALSE(docs.empty()); - auto wal = WalFile::Create(FindWal(path)); - ASSERT_EQ(wal->remove(), 0); - WalOptions options; - options.create_new = true; - ASSERT_EQ(wal->open(options), 0); - for (auto &doc : docs) { - if (doc->pk_ref() == "target") { - doc->set_operator(Operator::UPSERT); - // Legacy UPSERT did not record its predecessor's ID. - doc->set_doc_id(0); - } - auto bytes = doc->serialize(); - ASSERT_EQ(wal->append(std::string(bytes.begin(), bytes.end())), 0); - } - ASSERT_EQ(wal->flush(), 0); - ASSERT_EQ(wal->close(), 0); -} - -void ExpectOnlyTarget(const Collection::Ptr &collection, - const std::string &value) { - auto fetched = collection->fetch({"target"}); - ASSERT_TRUE(fetched.has_value()) << fetched.error().message(); - ASSERT_NE(fetched.value().at("target"), nullptr); - EXPECT_EQ(fetched.value().at("target")->get("text"), value); - - SearchQuery query; - query.topk_ = 10; - query.filter_ = "text != ''"; - auto matches = collection->query(query); - ASSERT_TRUE(matches.has_value()) << matches.error().message(); - ASSERT_EQ(matches.value().size(), 1u); - EXPECT_EQ(matches.value().front()->pk_ref(), "target"); - EXPECT_EQ(matches.value().front()->get("text"), value); - auto stats = collection->stats(); - ASSERT_TRUE(stats.has_value()) << stats.error().message(); - EXPECT_EQ(stats.value().doc_count, 1u); -} - class RelaxedValidationDeathTest : public ::testing::Test { protected: void SetUp() override { @@ -218,80 +42,6 @@ class RelaxedValidationDeathTest : public ::testing::Test { ailego::FileHelper::RemovePath(path_.c_str()); } - void CheckLegacyCommittedPredecessor(bool separate_segment) { - ASSERT_EXIT(WriteUpsertsAndExit(path_, true, false, separate_segment), - ::testing::ExitedWithCode(0), ""); - if (!::testing::internal::InDeathTestChild()) { - auto recovered_version = VersionManager::Recovery(path_); - ASSERT_TRUE(recovered_version.has_value()); - const auto version = recovered_version.value()->get_current_version(); - EXPECT_EQ(version.persisted_segment_metas().empty(), !separate_segment); - if (!separate_segment) { - EXPECT_FALSE( - version.writing_segment_meta()->persisted_blocks().empty()); - } - ASSERT_EQ( - version.writing_segment_meta()->writing_forward_block()->min_doc_id(), - 1u); - - // Pin the new writer's format before constructing a historical WAL. - std::vector docs; - ASSERT_NO_FATAL_FAILURE(ReadWalDocuments(path_, &docs)); - ASSERT_EQ(docs.size(), 2u); - EXPECT_EQ(docs[0]->get_operator(), Operator::UPDATE); - EXPECT_EQ(docs[0]->doc_id(), 0u); - EXPECT_EQ(docs[1]->get_operator(), Operator::UPDATE); - EXPECT_EQ(docs[1]->doc_id(), 1u); - ASSERT_NO_FATAL_FAILURE(RewriteWalAsLegacyUpserts(path_)); - - auto map = - IDMap::CreateAndOpen("wal_recovery", FindIdMap(path_), false, false); - ASSERT_NE(map, nullptr); - // The map can reach disk ahead of the manifest's deletion snapshot. - // Its latest replay ID no longer identifies committed predecessor ID 0. - ASSERT_TRUE(map->upsert("target", 2).ok()); - ASSERT_TRUE(map->flush().ok()); - } - - // Neither child flushes its recovered deletion bitmap. Both retries must - // rediscover ID 0, including when it belongs to another persisted segment. - for (int attempt = 0; attempt < 2; ++attempt) { - ASSERT_EXIT( - { - auto opened = Collection::Open(path_, CollectionOptions{}); - if (!opened.has_value()) { - std::cerr << opened.error() << std::endl; - std::_Exit(1); - } - ExpectOnlyTarget(opened.value(), "second"); - std::_Exit(::testing::Test::HasFailure() ? 2 : 0); - }, - ::testing::ExitedWithCode(0), ""); - } - - { - auto opened = Collection::Open(path_, CollectionOptions{}); - ASSERT_TRUE(opened.has_value()) << opened.error().message(); - ASSERT_NO_FATAL_FAILURE(ExpectOnlyTarget(opened.value(), "second")); - ASSERT_TRUE(opened.value()->flush().ok()); - } - CollectionOptions options; - options.read_only_ = true; - auto reopened = Collection::Open(path_, options); - ASSERT_TRUE(reopened.has_value()) << reopened.error().message(); - ASSERT_NO_FATAL_FAILURE(ExpectOnlyTarget(reopened.value(), "second")); - auto iterator = reopened.value()->create_iterator(); - ASSERT_TRUE(iterator.has_value()) << iterator.error().message(); - auto first = iterator.value()->next(); - ASSERT_TRUE(first.has_value()); - ASSERT_NE(first.value(), nullptr); - EXPECT_EQ(first.value()->pk_ref(), "target"); - EXPECT_EQ(first.value()->get("text"), "second"); - auto end = iterator.value()->next(); - ASSERT_TRUE(end.has_value()); - EXPECT_EQ(end.value(), nullptr); - } - const std::string path_{"relaxed_validation_recovery_db"}; }; @@ -353,270 +103,5 @@ TEST_F(RelaxedValidationDeathTest, Utf8AndLongIdsRecoverFromUnflushedWal) { EXPECT_EQ(reopened.value()->stats().value().doc_count, ids.size()); } -TEST_F(RelaxedValidationDeathTest, LargeStringAndTrailingDocRecoverFromWal) { - const std::vector ids{"prefix", std::string(1024, 'x'), - "suffix"}; - // The same value fits below 4MiB with the old 64-byte ID limit. A 1024-byte - // ID takes its serialized WAL record over that former reader-only limit. - const std::vector values{"before", std::string(4193700, 'v'), - "after"}; - ::testing::FLAGS_gtest_death_test_style = "threadsafe"; - ASSERT_EXIT(WriteStringDocsAndExit(path_, ids, values), - ::testing::ExitedWithCode(0), ""); - for (int reopen = 0; reopen < 2; ++reopen) { - auto opened = Collection::Open(path_, CollectionOptions{}); - ASSERT_TRUE(opened.has_value()) << opened.error().message(); - auto collection = std::move(opened).value(); - auto fetched = collection->fetch(ids); - ASSERT_TRUE(fetched.has_value()) << fetched.error().message(); - for (size_t i = 0; i < ids.size(); ++i) { - auto found = fetched.value().find(ids[i]); - ASSERT_NE(found, fetched.value().end()); - ASSERT_NE(found->second, nullptr); - EXPECT_EQ(found->second->pk_ref(), ids[i]); - EXPECT_EQ(found->second->get("text"), values[i]); - } - EXPECT_EQ(collection->stats().value().doc_count, ids.size()); - ASSERT_TRUE(collection->flush().ok()); - } -} - -TEST_F(RelaxedValidationDeathTest, IncompleteTailAllowsLaterCrashRecovery) { - ::testing::FLAGS_gtest_death_test_style = "threadsafe"; - ASSERT_EXIT( - WriteStringDocsAndExit(path_, {"prefix", "torn"}, {"before", "tail"}), - ::testing::ExitedWithCode(0), ""); - if (!::testing::internal::InDeathTestChild()) { - const auto wal_path = FindWal(path_); - ASSERT_FALSE(wal_path.empty()); - std::filesystem::resize_file(wal_path, - std::filesystem::file_size(wal_path) - 1); - const auto tail_bytes = ReadFileBytes(wal_path); - const auto manifests = ReadManifests(path_); - ASSERT_FALSE(manifests.empty()); - { - CollectionOptions options; - options.read_only_ = true; - auto opened = Collection::Open(path_, options); - // Recovery needs writable index stores. A read-only attempt must fail - // explicitly and leave the WAL and committed manifest untouched. - ASSERT_FALSE(opened.has_value()); - EXPECT_EQ(opened.error().code(), StatusCode::FAILED_PRECONDITION); - EXPECT_NE(opened.error().message().find("read-write mode once"), - std::string::npos); - } - EXPECT_EQ(ReadFileBytes(wal_path), tail_bytes); - EXPECT_EQ(ReadManifests(path_), manifests); - } - - ASSERT_EXIT( - { - auto opened = Collection::Open(path_, CollectionOptions{}); - if (!opened.has_value()) { - std::cerr << opened.error() << std::endl; - std::_Exit(1); - } - Doc doc; - doc.set_pk("suffix"); - doc.set("text", "after"); - std::vector docs{doc}; - auto inserted = opened.value()->insert(docs); - if (!inserted.has_value() || !inserted.value().front().ok()) - std::_Exit(2); - std::_Exit(0); - }, - ::testing::ExitedWithCode(0), ""); - auto opened = Collection::Open(path_, CollectionOptions{}); - ASSERT_TRUE(opened.has_value()) << opened.error().message(); - auto fetched = opened.value()->fetch({"prefix", "torn", "suffix"}); - ASSERT_TRUE(fetched.has_value()); - ASSERT_NE(fetched.value().at("prefix"), nullptr); - ASSERT_NE(fetched.value().at("suffix"), nullptr); - EXPECT_TRUE(fetched.value().find("torn") == fetched.value().end() || - fetched.value().at("torn") == nullptr); - EXPECT_EQ(opened.value()->stats().value().doc_count, 2); -} - -TEST_F(RelaxedValidationDeathTest, CorruptWalFailsOpenWithoutReplacingFiles) { - ::testing::FLAGS_gtest_death_test_style = "threadsafe"; - ASSERT_EXIT( - WriteStringDocsAndExit(path_, {"prefix", "broken"}, {"before", "after"}), - ::testing::ExitedWithCode(0), ""); - const auto wal_path = FindWal(path_); - ASSERT_FALSE(wal_path.empty()); - { - std::fstream file(wal_path, - std::ios::in | std::ios::out | std::ios::binary); - ASSERT_TRUE(file.is_open()); - uint32_t first_length; - file.seekg(64); - file.read(reinterpret_cast(&first_length), sizeof(first_length)); - ASSERT_TRUE(file.good()); - // Damage the second record CRC, keeping its complete framing intact. - file.seekp(64 + 8 + first_length + 4); - const uint32_t bad_crc = 0; - file.write(reinterpret_cast(&bad_crc), sizeof(bad_crc)); - ASSERT_TRUE(file.good()); - } - const auto bytes = ReadFileBytes(wal_path); - const auto manifests = ReadManifests(path_); - ASSERT_FALSE(manifests.empty()); - for (int attempt = 0; attempt < 2; ++attempt) { - auto opened = Collection::Open(path_, CollectionOptions{}); - ASSERT_FALSE(opened.has_value()); - EXPECT_NE(opened.error().message().find("CRC mismatch"), std::string::npos); - EXPECT_EQ(ReadFileBytes(wal_path), bytes); - EXPECT_EQ(ReadManifests(path_), manifests); - } -} - -TEST_F(RelaxedValidationDeathTest, CorruptTailDoesNotApplyUpsertPrefix) { - ::testing::FLAGS_gtest_death_test_style = "threadsafe"; - ASSERT_EXIT(WriteUpsertsAndExit(path_, true, true), - ::testing::ExitedWithCode(0), ""); - ASSERT_NO_FATAL_FAILURE(RewriteWalAsLegacyUpserts(path_)); - const auto wal_path = FindWal(path_); - ASSERT_FALSE(wal_path.empty()); - auto bytes = ReadFileBytes(wal_path); - size_t prefix_end = 64; - for (int i = 0; i < 2; ++i) { - uint32_t length; - ASSERT_GE(bytes.size() - prefix_end, 8u); - std::memcpy(&length, bytes.data() + prefix_end, sizeof(length)); - prefix_end += 8 + length; - ASSERT_LE(prefix_end, bytes.size()); - } - ASSERT_GE(bytes.size() - prefix_end, 8u); - bytes[prefix_end + 4] ^= 1; // Corrupt only the trailing record's CRC. - { - std::ofstream file(wal_path, std::ios::binary | std::ios::trunc); - file.write(bytes.data(), bytes.size()); - ASSERT_TRUE(file.good()); - } - const auto idmap_path = FindIdMap(path_); - ASSERT_FALSE(idmap_path.empty()); - { - auto map = IDMap::CreateAndOpen("wal_recovery", idmap_path, false, false); - ASSERT_NE(map, nullptr); - // The original row is committed as ID 0. Make the pre-recovery mapping - // deterministic even if RocksDB flushed uncommitted writes before exit. - ASSERT_TRUE(map->upsert("target", 0).ok()); - map->remove("broken"); - ASSERT_TRUE(map->flush().ok()); - } - const auto manifests = ReadManifests(path_); - for (int attempt = 0; attempt < 2; ++attempt) { - auto opened = Collection::Open(path_, CollectionOptions{}); - ASSERT_FALSE(opened.has_value()); - EXPECT_NE(opened.error().message().find("CRC mismatch"), std::string::npos); - EXPECT_EQ(ReadFileBytes(wal_path), bytes); - EXPECT_EQ(ReadManifests(path_), manifests); - auto map = IDMap::CreateAndOpen("wal_recovery", idmap_path, false, true); - ASSERT_NE(map, nullptr); - uint64_t original_id; - ASSERT_TRUE(map->has("target", &original_id)); - EXPECT_EQ(original_id, 0u); - EXPECT_FALSE(map->has("broken")); - } - - // Remove the damaged last record and retry the intact UPSERT prefix. - std::filesystem::resize_file(wal_path, prefix_end); - { - auto opened = Collection::Open(path_, CollectionOptions{}); - ASSERT_TRUE(opened.has_value()) << opened.error().message(); - auto fetched = opened.value()->fetch({"target"}); - ASSERT_TRUE(fetched.has_value()); - ASSERT_NE(fetched.value().at("target"), nullptr); - EXPECT_EQ(fetched.value().at("target")->get("text"), "second"); - EXPECT_EQ(opened.value()->stats().value().doc_count, 1); - Doc update; - update.set_pk("target"); - update.set("text", "third"); - std::vector updates{update}; - auto updated = opened.value()->upsert(updates); - ASSERT_TRUE(updated.has_value()); - ASSERT_TRUE(updated.value().front().ok()); - ASSERT_TRUE(opened.value()->flush().ok()); - } - auto reopened = Collection::Open(path_, CollectionOptions{}); - ASSERT_TRUE(reopened.has_value()); - auto fetched = reopened.value()->fetch({"target"}); - ASSERT_TRUE(fetched.has_value()); - ASSERT_NE(fetched.value().at("target"), nullptr); - EXPECT_EQ(fetched.value().at("target")->get("text"), "third"); - EXPECT_EQ(reopened.value()->stats().value().doc_count, 1); -} - -TEST_F(RelaxedValidationDeathTest, UpsertReplayIgnoresSameAndLaterReplayIds) { - ::testing::FLAGS_gtest_death_test_style = "threadsafe"; - ASSERT_EXIT(WriteUpsertsAndExit(path_, false, false), - ::testing::ExitedWithCode(0), ""); - if (!::testing::internal::InDeathTestChild()) { - ASSERT_NO_FATAL_FAILURE(RewriteWalAsLegacyUpserts(path_)); - const auto idmap_path = FindIdMap(path_); - ASSERT_FALSE(idmap_path.empty()); - auto map = IDMap::CreateAndOpen("wal_recovery", idmap_path, false, false); - ASSERT_NE(map, nullptr); - // Simulate a partial previous recovery that persisted the second UPSERT's - // ID. It is later than the first replay ID and equal to the second. - ASSERT_TRUE(map->upsert("target", 1).ok()); - ASSERT_TRUE(map->flush().ok()); - } - ASSERT_EXIT( - { - auto opened = Collection::Open(path_, CollectionOptions{}); - if (!opened.has_value()) { - std::cerr << opened.error() << std::endl; - std::_Exit(1); - } - auto fetched = opened.value()->fetch({"target"}); - if (!fetched.has_value() || !fetched.value().at("target") || - fetched.value().at("target")->get("text") != - "second" || - opened.value()->stats().value().doc_count != 1) - std::_Exit(2); - Doc update; - update.set_pk("target"); - update.set("text", "third"); - std::vector updates{update}; - auto updated = opened.value()->upsert(updates); - if (!updated.has_value() || !updated.value().front().ok()) - std::_Exit(3); - std::_Exit(0); - }, - ::testing::ExitedWithCode(0), ""); - auto reopened = Collection::Open(path_, CollectionOptions{}); - ASSERT_TRUE(reopened.has_value()) << reopened.error().message(); - auto fetched = reopened.value()->fetch({"target"}); - ASSERT_TRUE(fetched.has_value()); - ASSERT_NE(fetched.value().at("target"), nullptr); - EXPECT_EQ(fetched.value().at("target")->get("text"), "third"); - EXPECT_EQ(reopened.value()->stats().value().doc_count, 1); -} - -TEST_F(RelaxedValidationDeathTest, - LegacyUpsertsRecoverCommittedPredecessorInWritingSegment) { - ASSERT_NO_FATAL_FAILURE(CheckLegacyCommittedPredecessor(false)); -} - -TEST_F(RelaxedValidationDeathTest, - LegacyUpsertsRecoverCommittedPredecessorInPersistedSegment) { - ASSERT_NO_FATAL_FAILURE(CheckLegacyCommittedPredecessor(true)); -} - -TEST_F(RelaxedValidationDeathTest, UpsertWalRecordsInsertOrUpdatePredecessor) { - ASSERT_EXIT(WriteUpsertsAndExit(path_, false, false), - ::testing::ExitedWithCode(0), ""); - std::vector docs; - ASSERT_NO_FATAL_FAILURE(ReadWalDocuments(path_, &docs)); - ASSERT_EQ(docs.size(), 2u); - EXPECT_EQ(docs[0]->get_operator(), Operator::INSERT); - EXPECT_EQ(docs[0]->pk_ref(), "target"); - EXPECT_EQ(docs[0]->get("text"), "first"); - EXPECT_EQ(docs[1]->get_operator(), Operator::UPDATE); - EXPECT_EQ(docs[1]->doc_id(), 0u); - EXPECT_EQ(docs[1]->get("text"), "second"); -} - } // namespace } // namespace zvec diff --git a/tests/db/index/common/doc_test.cc b/tests/db/index/common/doc_test.cc index 95342df33..32e875651 100644 --- a/tests/db/index/common/doc_test.cc +++ b/tests/db/index/common/doc_test.cc @@ -14,6 +14,7 @@ #include "zvec/db/doc.h" #include +#include #include #include #include @@ -1337,12 +1338,21 @@ TEST(SearchQuery, ValidateAndSanitize) { v.size() * sizeof(float)); }; auto decode_idx = [](const std::string &buf) { - const auto *p = reinterpret_cast(buf.data()); - return std::vector(p, p + buf.size() / sizeof(uint32_t)); + EXPECT_EQ(buf.size() % sizeof(uint32_t), 0u); + std::vector indices(buf.size() / sizeof(uint32_t)); + if (!indices.empty()) { + std::memcpy(indices.data(), buf.data(), + indices.size() * sizeof(uint32_t)); + } + return indices; }; auto decode_val = [](const std::string &buf) { - const auto *p = reinterpret_cast(buf.data()); - return std::vector(p, p + buf.size() / sizeof(float)); + EXPECT_EQ(buf.size() % sizeof(float), 0u); + std::vector values(buf.size() / sizeof(float)); + if (!values.empty()) { + std::memcpy(values.data(), buf.data(), values.size() * sizeof(float)); + } + return values; }; FieldSchema schema = FieldSchema("field_name", DataType::SPARSE_VECTOR_FP32); @@ -1688,7 +1698,7 @@ TEST_F(DocDetailedTest, ValidationErrorsEscapeAndBoundDocumentAndFieldNames) { doc.set(name, int32_t{42}); const auto status = doc.validate_and_sanitize(schema); EXPECT_EQ(status.code(), StatusCode::INVALID_ARGUMENT); - EXPECT_EQ(status.message().find("Invalid doc:"), 0u); + EXPECT_EQ(status.message().find("Invalid doc["), 0u); EXPECT_NE(status.message().find("does not exist"), std::string::npos); EXPECT_LT(status.message().size(), 256u); for (unsigned char byte : status.message()) { diff --git a/tests/db/index/common/manifest_codec_golden_test.cc b/tests/db/index/common/manifest_codec_golden_test.cc index 7846022e5..1bde56b1e 100644 --- a/tests/db/index/common/manifest_codec_golden_test.cc +++ b/tests/db/index/common/manifest_codec_golden_test.cc @@ -1122,38 +1122,6 @@ TEST(ManifestCodecGolden, LegacyFieldNamesSurviveWithoutRevalidation) { EXPECT_EQ(reencoded, encoded_manifest); } -TEST(ManifestCodecGolden, DuplicateFieldsFailInsteadOfBeingSilentlyDropped) { - // This malformed schema could previously be created through the C++ list - // constructor. Restoring only its first field silently changes its meaning. - CollectionSchema schema( - "legacy", {std::make_shared("duplicate", DataType::INT32), - std::make_shared("duplicate", DataType::INT64)}); - std::string schema_bytes; - ManifestCodec::EncodeCollectionSchema(schema, &schema_bytes); - auto decoded_schema = ManifestCodec::DecodeCollectionSchema(schema_bytes); - ASSERT_FALSE(decoded_schema.has_value()); - EXPECT_EQ(decoded_schema.error().code(), StatusCode::INTERNAL_ERROR); - EXPECT_NE(decoded_schema.error().message().find("duplicate"), - std::string::npos); - - std::string manifest_bytes; - pbwire::Writer(&manifest_bytes).PutMessage(2, schema_bytes); - ManifestData restored; - auto status = ManifestCodec::Decode(manifest_bytes, &restored); - EXPECT_EQ(status.code(), StatusCode::INTERNAL_ERROR); - EXPECT_EQ(restored.schema, nullptr); -} - -TEST(ManifestCodecGolden, MalformedNestedSchemaFailsExplicitly) { - const std::string malformed_schema("\x0a\x05x", 3); - std::string manifest_bytes; - pbwire::Writer(&manifest_bytes).PutMessage(2, malformed_schema); - ManifestData restored; - auto status = ManifestCodec::Decode(manifest_bytes, &restored); - EXPECT_EQ(status.code(), StatusCode::INTERNAL_ERROR); - EXPECT_EQ(restored.schema, nullptr); -} - TEST(ManifestCodecGolden, StructuralSchemaHelpersPreserveLegacyNames) { CollectionSchema schema("legacy_collection"); auto status = schema.add_field( diff --git a/tests/db/index/storage/wal_file_test.cc b/tests/db/index/storage/wal_file_test.cc index a29e80e3d..50cd122da 100644 --- a/tests/db/index/storage/wal_file_test.cc +++ b/tests/db/index/storage/wal_file_test.cc @@ -12,14 +12,20 @@ // See the License for the specific language governing permissions and // limitations under the License. +#ifdef _MSC_VER +#define _ALLOW_KEYWORD_MACROS +#endif +#define private public +#define protected public #include "db/index/storage/wal/wal_file.h" +#undef private +#undef protected + #include #include #include #include -#include #include -#include #include #include #include @@ -41,16 +47,6 @@ class WalFileTest : public testing::Test { } void TearDown() override {} - - // Legacy success-path tests use empty string as their loop sentinel, but - // still assert that the new API did not report an error. - std::string ReadRecord(const WalFilePtr &wal_file) { - auto result = wal_file->next(); - EXPECT_TRUE(result.has_value()) - << (result.has_value() ? "" : result.error().message()); - if (!result.has_value() || !result.value().has_value()) return {}; - return std::move(result.value().value()); - } }; TEST_F(WalFileTest, TestGeneral) { @@ -130,14 +126,14 @@ TEST_F(WalFileTest, TestGeneral) { uint32_t idx = 0; ret = wal_file->prepare_for_read(); ASSERT_EQ(ret, 0); - std::string record = ReadRecord(wal_file); + std::string record = wal_file->next(); while (!record.empty()) { if (idx < 100) { ASSERT_EQ(record, "hello"); } else { ASSERT_EQ(record, std::string("hello") + std::to_string(idx)); } - record = ReadRecord(wal_file); + record = wal_file->next(); idx++; } ASSERT_EQ(idx, 400); @@ -209,9 +205,9 @@ TEST_F(WalFileTest, TestMultiThread) { uint32_t idx = 0; ret = wal_file->prepare_for_read(); ASSERT_EQ(ret, 0); - std::string record = ReadRecord(wal_file); + std::string record = wal_file->next(); while (!record.empty()) { - record = ReadRecord(wal_file); + record = wal_file->next(); idx++; } ASSERT_EQ(idx, 30000); @@ -247,9 +243,9 @@ TEST_F(WalFileTest, TestBoundaryCondition) { ret = wal_file->open(wal_option); ASSERT_EQ(ret, 0); uint32_t idx = 0; - std::string record = ReadRecord(wal_file); + std::string record = wal_file->next(); while (!record.empty()) { - record = ReadRecord(wal_file); + record = wal_file->next(); idx++; } ASSERT_EQ(idx, 0); @@ -275,13 +271,13 @@ TEST_F(WalFileTest, TestBoundaryCondition) { idx = 0; ret = wal_file->prepare_for_read(); ASSERT_EQ(ret, 0); - record = ReadRecord(wal_file); + record = wal_file->next(); while (!record.empty()) { ASSERT_EQ(record.size(), 4); for (size_t i = 0; i < 4; i++) { ASSERT_EQ(record[i], i); } - record = ReadRecord(wal_file); + record = wal_file->next(); idx++; } ASSERT_EQ(idx, 1); @@ -316,13 +312,13 @@ TEST_F(WalFileTest, TestBoundaryCondition) { idx = 0; ret = wal_file->prepare_for_read(); ASSERT_EQ(ret, 0); - record = ReadRecord(wal_file); + record = wal_file->next(); while (!record.empty()) { ASSERT_EQ(record.size(), BIG_DATA_SIZE); for (size_t i = 0; i < BIG_DATA_SIZE; i++) { ASSERT_EQ((uint8_t)record[i], i % 256); } - record = ReadRecord(wal_file); + record = wal_file->next(); idx++; } ASSERT_EQ(idx, 1); @@ -353,10 +349,10 @@ TEST_F(WalFileTest, TestBoundaryCondition) { idx = 0; ret = wal_file->prepare_for_read(); ASSERT_EQ(ret, 0); - record = ReadRecord(wal_file); + record = wal_file->next(); while (!record.empty()) { ASSERT_EQ(record, std::string("hello") + std::to_string(idx)); - record = ReadRecord(wal_file); + record = wal_file->next(); idx++; } ASSERT_EQ(idx, 99); @@ -421,15 +417,16 @@ TEST_F(WalFileTest, TestFirstErrorCase) { ret = wal_file->open(wal_option); ASSERT_EQ(ret, 0); - ASSERT_EQ(wal_file->prepare_for_read(), 0); - for (size_t i = 0; i < 0; ++i) { - EXPECT_EQ(ReadRecord(wal_file), "hello"); + uint32_t idx = 0; + ret = wal_file->prepare_for_read(); + ASSERT_EQ(ret, 0); + std::string record = wal_file->next(); + while (!record.empty()) { + ASSERT_EQ(record, "hello"); + record = wal_file->next(); + idx++; } - auto result = wal_file->next(); - ASSERT_FALSE(result.has_value()); - EXPECT_EQ(result.error().code(), StatusCode::INTERNAL_ERROR); - EXPECT_NE(result.error().message().find("CRC mismatch"), std::string::npos); - EXPECT_NE(wal_file->append("after corruption"), 0); + ASSERT_EQ(idx, 0); // close ret = wal_file->close(); ASSERT_EQ(ret, 0); @@ -480,15 +477,16 @@ TEST_F(WalFileTest, TestMiddleErrorCase) { ret = wal_file->open(wal_option); ASSERT_EQ(ret, 0); - ASSERT_EQ(wal_file->prepare_for_read(), 0); - for (size_t i = 0; i < 5; ++i) { - EXPECT_EQ(ReadRecord(wal_file), "hello"); + uint32_t idx = 0; + ret = wal_file->prepare_for_read(); + ASSERT_EQ(ret, 0); + std::string record = wal_file->next(); + while (!record.empty()) { + ASSERT_EQ(record, "hello"); + record = wal_file->next(); + idx++; } - auto result = wal_file->next(); - ASSERT_FALSE(result.has_value()); - EXPECT_EQ(result.error().code(), StatusCode::INTERNAL_ERROR); - EXPECT_NE(result.error().message().find("CRC mismatch"), std::string::npos); - EXPECT_NE(wal_file->append("after corruption"), 0); + ASSERT_EQ(idx, 5); // close ret = wal_file->close(); ASSERT_EQ(ret, 0); @@ -540,10 +538,10 @@ TEST_F(WalFileTest, TestLastErrorCase) { uint32_t idx = 0; ret = wal_file->prepare_for_read(); ASSERT_EQ(ret, 0); - std::string record = ReadRecord(wal_file); + std::string record = wal_file->next(); while (!record.empty()) { ASSERT_EQ(record, "hello"); - record = ReadRecord(wal_file); + record = wal_file->next(); idx++; } ASSERT_EQ(idx, 9); @@ -596,15 +594,16 @@ TEST_F(WalFileTest, TestLengthSmallErrorCase) { ret = wal_file->open(wal_option); ASSERT_EQ(ret, 0); - ASSERT_EQ(wal_file->prepare_for_read(), 0); - for (size_t i = 0; i < 0; ++i) { - EXPECT_EQ(ReadRecord(wal_file), "hello"); + uint32_t idx = 0; + ret = wal_file->prepare_for_read(); + ASSERT_EQ(ret, 0); + std::string record = wal_file->next(); + while (!record.empty()) { + ASSERT_EQ(record, "hello"); + record = wal_file->next(); + idx++; } - auto result = wal_file->next(); - ASSERT_FALSE(result.has_value()); - EXPECT_EQ(result.error().code(), StatusCode::INTERNAL_ERROR); - EXPECT_NE(result.error().message().find("CRC mismatch"), std::string::npos); - EXPECT_NE(wal_file->append("after corruption"), 0); + ASSERT_EQ(idx, 0); // close ret = wal_file->close(); ASSERT_EQ(ret, 0); @@ -644,7 +643,7 @@ TEST_F(WalFileTest, TestLengthBigErrorCase) { dir_path, "data.wal.", std::to_string(segment_id)); int wal_fd = open(wal_path.c_str(), O_RDWR, 0644); ASSERT_GT(wal_fd, 0); - uint32_t err_length = std::numeric_limits::max(); + uint32_t err_length = 200; // exceed file size 130 lseek(wal_fd, 64, SEEK_SET); write(wal_fd, (const void *)&err_length, 4); @@ -658,10 +657,10 @@ TEST_F(WalFileTest, TestLengthBigErrorCase) { uint32_t idx = 0; ret = wal_file->prepare_for_read(); ASSERT_EQ(ret, 0); - std::string record = ReadRecord(wal_file); + std::string record = wal_file->next(); while (!record.empty()) { ASSERT_EQ(record, "hello"); - record = ReadRecord(wal_file); + record = wal_file->next(); idx++; } ASSERT_EQ(idx, 0); @@ -715,15 +714,16 @@ TEST_F(WalFileTest, TestCRCErrorCase) { ret = wal_file->open(wal_option); ASSERT_EQ(ret, 0); - ASSERT_EQ(wal_file->prepare_for_read(), 0); - for (size_t i = 0; i < 1; ++i) { - EXPECT_EQ(ReadRecord(wal_file), "hello"); + uint32_t idx = 0; + ret = wal_file->prepare_for_read(); + ASSERT_EQ(ret, 0); + std::string record = wal_file->next(); + while (!record.empty()) { + ASSERT_EQ(record, "hello"); + record = wal_file->next(); + idx++; } - auto result = wal_file->next(); - ASSERT_FALSE(result.has_value()); - EXPECT_EQ(result.error().code(), StatusCode::INTERNAL_ERROR); - EXPECT_NE(result.error().message().find("CRC mismatch"), std::string::npos); - EXPECT_NE(wal_file->append("after corruption"), 0); + ASSERT_EQ(idx, 1); // close ret = wal_file->close(); ASSERT_EQ(ret, 0); @@ -732,115 +732,6 @@ TEST_F(WalFileTest, TestCRCErrorCase) { ASSERT_EQ(ret, 0); } -TEST_F(WalFileTest, RecordLargerThanFourMiBPreservesFollowingRecord) { - const std::string path = "./data.wal.large"; - auto wal = WalFile::Create(path); - WalOptions options; - options.create_new = true; - ASSERT_EQ(wal->open(options), 0); - const std::string large(4 * 1024 * 1024 + 1024, 'x'); - ASSERT_EQ(wal->append("prefix"), 0); - ASSERT_EQ(wal->append(std::string(large)), 0); - ASSERT_EQ(wal->append("suffix"), 0); - ASSERT_EQ(wal->close(), 0); - options.create_new = false; - ASSERT_EQ(wal->open(options), 0); - ASSERT_EQ(wal->prepare_for_read(), 0); - EXPECT_EQ(ReadRecord(wal), "prefix"); - EXPECT_EQ(ReadRecord(wal), large); - EXPECT_EQ(ReadRecord(wal), "suffix"); - auto end = wal->next(); - ASSERT_TRUE(end.has_value()); - EXPECT_FALSE(end.value().has_value()); -} - -TEST_F(WalFileTest, IncompleteTailIsRemovedOnlyBeforeAppend) { - const std::string path = "./data.wal.tail"; - constexpr size_t prefix_end = 64 + 8 + 6; - // Exercise every partial header and partial payload boundary. - for (size_t tail_size = 1; tail_size < 8 + 4; ++tail_size) { - SCOPED_TRACE(tail_size); - auto wal = WalFile::Create(path); - WalOptions options; - options.create_new = true; - ASSERT_EQ(wal->open(options), 0); - ASSERT_EQ(wal->append("prefix"), 0); - ASSERT_EQ(wal->append("torn"), 0); - ASSERT_EQ(wal->close(), 0); - ailego::File file; - ASSERT_TRUE(file.open(path, false)); - ASSERT_TRUE(file.truncate(prefix_end + tail_size)); - file.close(); - - options.create_new = false; - ASSERT_EQ(wal->open(options), 0); - ASSERT_EQ(wal->prepare_for_read(), 0); - EXPECT_EQ(ReadRecord(wal), "prefix"); - auto end = wal->next(); - ASSERT_TRUE(end.has_value()); - EXPECT_FALSE(end.value().has_value()); - ASSERT_TRUE(file.open(path, true)); - EXPECT_EQ(file.size(), prefix_end + tail_size); - file.close(); - - ASSERT_EQ(wal->append("suffix"), 0); - ASSERT_EQ(wal->close(), 0); - ASSERT_EQ(wal->open(options), 0); - ASSERT_EQ(wal->prepare_for_read(), 0); - EXPECT_EQ(ReadRecord(wal), "prefix"); - EXPECT_EQ(ReadRecord(wal), "suffix"); - end = wal->next(); - ASSERT_TRUE(end.has_value()); - EXPECT_FALSE(end.value().has_value()); - ASSERT_EQ(wal->remove(), 0); - } -} - -TEST_F(WalFileTest, ClosedReaderReturnsError) { - auto wal = WalFile::Create("./data.wal.closed"); - auto result = wal->next(); - ASSERT_FALSE(result.has_value()); - EXPECT_EQ(result.error().code(), StatusCode::INTERNAL_ERROR); -} - -TEST_F(WalFileTest, ZeroLengthRecordIsCorruption) { - const std::string path = "./data.wal.zero"; - auto wal = WalFile::Create(path); - WalOptions options; - options.create_new = true; - ASSERT_EQ(wal->open(options), 0); - EXPECT_NE(wal->append(""), 0); - ASSERT_EQ(wal->append("payload"), 0); - ASSERT_EQ(wal->close(), 0); - ailego::File file; - ASSERT_TRUE(file.open(path, false)); - const uint32_t length = 0; - ASSERT_EQ(file.write(64, &length, sizeof(length)), sizeof(length)); - file.close(); - options.create_new = false; - ASSERT_EQ(wal->open(options), 0); - ASSERT_EQ(wal->prepare_for_read(), 0); - auto result = wal->next(); - ASSERT_FALSE(result.has_value()); - EXPECT_NE(result.error().message().find("zero length"), std::string::npos); -} - -TEST_F(WalFileTest, TruncatedFileHeaderReturnsError) { - const std::string path = "./data.wal.header"; - auto wal = WalFile::Create(path); - WalOptions options; - options.create_new = true; - ASSERT_EQ(wal->open(options), 0); - ASSERT_EQ(wal->close(), 0); - ailego::File file; - ASSERT_TRUE(file.open(path, false)); - ASSERT_TRUE(file.truncate(63)); - file.close(); - options.create_new = false; - ASSERT_EQ(wal->open(options), 0); - EXPECT_NE(wal->prepare_for_read(), 0); -} - #if defined(__GNUC__) || defined(__GNUG__) #pragma GCC diagnostic pop #endif \ No newline at end of file From 94db918313d74eb19f91893fc190f2369255bb5f Mon Sep 17 00:00:00 2001 From: Qinren Zhou Date: Wed, 16 Sep 2026 15:36:18 +0800 Subject: [PATCH 06/11] refactor: ut --- tests/db/collection_test.cc | 413 +++++++++++++++ .../relaxed_validation_recovery_test.cc | 107 ---- .../db/crash_recovery/write_recovery_test.cc | 62 +++ .../common/manifest_codec_golden_test.cc | 62 --- tests/db/relaxed_validation_test.cc | 485 ------------------ 5 files changed, 475 insertions(+), 654 deletions(-) delete mode 100644 tests/db/crash_recovery/relaxed_validation_recovery_test.cc delete mode 100644 tests/db/relaxed_validation_test.cc diff --git a/tests/db/collection_test.cc b/tests/db/collection_test.cc index e811be6fc..aa663f6d5 100644 --- a/tests/db/collection_test.cc +++ b/tests/db/collection_test.cc @@ -45,6 +45,7 @@ #include #include #include "db/collection_query_internal.h" +#include "db/common/constants.h" #include "db/common/file_helper.h" #include "db/index/common/type_helper.h" #include "db/index/common/version_manager.h" @@ -74,9 +75,64 @@ class CollectionTest : public ::testing::Test { } void TearDown() override { + collection_.reset(); FileHelper::RemoveDirectory(col_path); ailego::FileHelper::RemoveDirectory("demo"); } + + CollectionSchema MakeSchema(const std::string &name = "x", + const std::string &field = "value") { + CollectionSchema schema(name); + EXPECT_TRUE(schema + .add_field(std::make_shared( + field, DataType::INT32, false)) + .ok()); + return schema; + } + + void Create(const CollectionSchema &schema) { + auto result = Collection::CreateAndOpen(col_path, schema, options_); + ASSERT_TRUE(result.has_value()) << result.error().message(); + collection_ = std::move(result).value(); + } + + void Reopen() { + collection_.reset(); + auto result = Collection::Open(col_path, options_); + ASSERT_TRUE(result.has_value()) << result.error().message(); + collection_ = std::move(result).value(); + } + + Doc MakeDoc(const std::string &id, int32_t value, + const std::string &field = "value") { + Doc doc; + doc.set_pk(id); + EXPECT_TRUE(doc.set(field, value)); + return doc; + } + + void ExpectWrite(const Result &result, size_t count) { + ASSERT_TRUE(result.has_value()) << result.error().message(); + ASSERT_EQ(result.value().size(), count); + for (const auto &status : result.value()) { + ASSERT_TRUE(status.ok()) << status.message(); + } + } + + void ExpectValue(const std::string &id, int32_t expected, + const std::string &field = "value") { + auto result = collection_->fetch({id}); + ASSERT_TRUE(result.has_value()) << result.error().message(); + ASSERT_EQ(result.value().size(), 1u); + auto found = result.value().find(id); + ASSERT_NE(found, result.value().end()); + ASSERT_NE(found->second, nullptr); + EXPECT_EQ(found->second->pk(), id); + EXPECT_EQ(found->second->get(field), expected); + } + + CollectionOptions options_; + Collection::Ptr collection_; }; class DirectoryWriteBlockerForTest { @@ -329,6 +385,46 @@ TEST_F(CollectionTest, Feature_OpenReadOnly_WithReadOnlyLockFile) { fs::perm_options::replace, ec); } +TEST_F(CollectionTest, Feature_CreateAndOpen_NameBoundaries) { + for (const auto &name : std::vector{ + "x", "xy", u8"集", std::string(253, 'n') + u8"集"}) { + SCOPED_TRACE(name); + ASSERT_NO_FATAL_FAILURE(Create(MakeSchema(name))); + std::vector docs{MakeDoc("id", 1)}; + ASSERT_NO_FATAL_FAILURE(ExpectWrite(collection_->insert(docs), 1)); + auto status = collection_->flush(); + ASSERT_TRUE(status.ok()) << status.message(); + ASSERT_NO_FATAL_FAILURE(Reopen()); + auto schema = collection_->schema(); + ASSERT_TRUE(schema.has_value()) << schema.error().message(); + EXPECT_EQ(schema.value().name(), name); + ASSERT_NO_FATAL_FAILURE(ExpectValue("id", 1)); + status = collection_->destroy(); + ASSERT_TRUE(status.ok()) << status.message(); + collection_.reset(); + } +} + +TEST_F(CollectionTest, Feature_CreateAndOpen_InvalidSchema) { + auto result = Collection::CreateAndOpen( + col_path, MakeSchema("x", "_zvec_uid_"), options_); + ASSERT_FALSE(result.has_value()); + EXPECT_EQ(result.error().code(), StatusCode::INVALID_ARGUMENT); + EXPECT_NE(result.error().message().find("is reserved"), std::string::npos); + EXPECT_FALSE(ailego::FileHelper::IsExist(col_path.c_str())); + + CollectionSchema duplicate( + "x", {std::make_shared("value", DataType::INT32), + std::make_shared("value", DataType::VECTOR_FP32, 4, + false)}); + result = Collection::CreateAndOpen(col_path, duplicate, options_); + ASSERT_FALSE(result.has_value()); + EXPECT_EQ(result.error().code(), StatusCode::INVALID_ARGUMENT); + EXPECT_NE(result.error().message().find("duplicate field name [value]"), + std::string::npos); + EXPECT_FALSE(ailego::FileHelper::IsExist(col_path.c_str())); +} + TEST_F(CollectionTest, Feature_CreateAndOpen_Empty) { int doc_count = 0; int loop_count = 100; @@ -505,6 +601,36 @@ TEST_F(CollectionTest, Feature_CreateAndOpen_MultiThread) { ASSERT_FALSE(has_error.load()); } +TEST_F(CollectionTest, Feature_Write_InvalidIdRejectsWholeBatch) { + ASSERT_NO_FATAL_FAILURE(Create(MakeSchema())); + std::vector initial{MakeDoc("existing", 1)}; + ASSERT_NO_FATAL_FAILURE(ExpectWrite(collection_->insert(initial), 1)); + for (int operation = 0; operation < 3; ++operation) { + SCOPED_TRACE(operation); + const std::string first_id = operation == 0 ? "new:id" : "existing"; + std::vector batch{MakeDoc(first_id, 99), + MakeDoc(std::string("bad\0id", 6), 100)}; + auto result = operation == 0 ? collection_->insert(batch) + : operation == 1 ? collection_->update(batch) + : collection_->upsert(batch); + ASSERT_FALSE(result.has_value()); + EXPECT_EQ(result.error().code(), StatusCode::INVALID_ARGUMENT); + EXPECT_EQ(result.error().message().find("Invalid doc:"), 0u); + EXPECT_NE(result.error().message().find("null character"), + std::string::npos); + EXPECT_NE(result.error().message().find("id[bad\\0id]"), std::string::npos); + EXPECT_EQ(result.error().message().find("offset"), std::string::npos); + ASSERT_NO_FATAL_FAILURE(ExpectValue("existing", 1)); + auto missing = collection_->fetch({"new:id"}); + ASSERT_TRUE(missing.has_value()) << missing.error().message(); + ASSERT_EQ(missing.value().size(), 1u); + EXPECT_EQ(missing.value().at("new:id"), nullptr); + } + ASSERT_NO_FATAL_FAILURE(Reopen()); + ASSERT_NO_FATAL_FAILURE(ExpectValue("existing", 1)); + EXPECT_EQ(collection_->stats().value().doc_count, 1u); +} + TEST_F(CollectionTest, Feature_Write_Batch_Validate) { FileHelper::RemoveDirectory(col_path); @@ -538,6 +664,50 @@ TEST_F(CollectionTest, Feature_Write_Batch_Validate) { ASSERT_FALSE(upsert_exceed_status.ok()); } +TEST_F(CollectionTest, Feature_Write_Utf8IdsAndReopen) { + const std::string name = u8"测试 集合/v1"; + ASSERT_NO_FATAL_FAILURE(Create(MakeSchema(name))); + const std::vector ids = { + "user:123", "https://example.com/document/42", + u8"订单-😀", std::string(1021, 'i') + u8"中", + "doc", " doc", + "doc ", " ", + u8"café", u8"cafe\u0301", + "DOC"}; + std::vector docs; + for (size_t i = 0; i < ids.size(); ++i) { + docs.push_back(MakeDoc(ids[i], static_cast(i))); + } + ASSERT_NO_FATAL_FAILURE(ExpectWrite(collection_->insert(docs), docs.size())); + + std::vector updates{MakeDoc(ids[0], 100)}; + ASSERT_NO_FATAL_FAILURE(ExpectWrite(collection_->update(updates), 1)); + std::vector upserts{MakeDoc(ids[3], 103), MakeDoc(u8"新增:文档", 200)}; + ASSERT_NO_FATAL_FAILURE(ExpectWrite(collection_->upsert(upserts), 2)); + ASSERT_NO_FATAL_FAILURE(ExpectWrite(collection_->delete_({ids[1]}), 1)); + + auto flush_status = collection_->flush(); + ASSERT_TRUE(flush_status.ok()) << flush_status.message(); + ASSERT_NO_FATAL_FAILURE(Reopen()); + auto schema = collection_->schema(); + ASSERT_TRUE(schema.has_value()) << schema.error().message(); + EXPECT_EQ(schema.value().name(), name); + for (size_t i = 0; i < ids.size(); ++i) { + if (i == 1) { + continue; + } + int32_t expected = static_cast(i); + if (i == 0) expected = 100; + if (i == 3) expected = 103; + ASSERT_NO_FATAL_FAILURE(ExpectValue(ids[i], expected)); + } + ASSERT_NO_FATAL_FAILURE(ExpectValue(u8"新增:文档", 200)); + auto deleted = collection_->fetch({ids[1]}); + ASSERT_TRUE(deleted.has_value()) << deleted.error().message(); + ASSERT_EQ(deleted.value().size(), 1u); + EXPECT_EQ(deleted.value().at(ids[1]), nullptr); +} + TEST_F(CollectionTest, Feature_Insert_General) { auto func = [&](bool enable_mmap, bool schema_nullable, bool doc_nullable, int doc_count = 1000) { @@ -1590,6 +1760,26 @@ TEST_F(CollectionTest, Feature_Update_Empty) { } } +TEST_F(CollectionTest, Feature_FetchAndDelete_InvalidIdsRemainMissing) { + ASSERT_NO_FATAL_FAILURE(Create(MakeSchema())); + // These are invalid for a new document, but lookup must retain its existing + // missing-key behavior rather than introducing input validation errors. + const std::vector absent_ids{"", std::string("bad\0id", 6), + std::string(1025, 'x'), "\xff"}; + auto fetched = collection_->fetch(absent_ids); + ASSERT_TRUE(fetched.has_value()) << fetched.error().message(); + ASSERT_EQ(fetched.value().size(), absent_ids.size()); + for (const auto &id : absent_ids) { + EXPECT_EQ(fetched.value().at(id), nullptr); + } + auto deleted = collection_->delete_(absent_ids); + ASSERT_TRUE(deleted.has_value()) << deleted.error().message(); + ASSERT_EQ(deleted.value().size(), absent_ids.size()); + for (const auto &status : deleted.value()) { + EXPECT_EQ(status.code(), StatusCode::NOT_FOUND); + } +} + TEST_F(CollectionTest, Feature_Delete_General) { auto func = [&](bool enable_mmap, int doc_count) { auto schema = TestHelper::CreateNormalSchema(); @@ -4407,6 +4597,72 @@ TEST_F(CollectionTest, Feature_Query_Validate) { } } +TEST_F(CollectionTest, Feature_Query_MaximumLengthFieldNames) { + auto check = [&](bool enable_mmap) { + SCOPED_TRACE(enable_mmap); + options_.enable_mmap_ = enable_mmap; + const std::string scalar = "s" + std::string(63, 'a'); + const std::string second_scalar = scalar.substr(0, 63) + "b"; + const std::string vector = "v" + std::string(63, 'b'); + auto schema = MakeSchema("x", scalar); + ASSERT_TRUE(schema + .add_field(std::make_shared( + second_scalar, DataType::INT32, false)) + .ok()); + ASSERT_TRUE(schema + .add_field(std::make_shared( + vector, DataType::VECTOR_FP32, 4, false, + std::make_shared(MetricType::L2))) + .ok()); + ASSERT_NO_FATAL_FAILURE(Create(schema)); + const std::vector values{1.0f, 2.0f, 3.0f, 4.0f}; + Doc doc = MakeDoc(u8"文档:1", 42, scalar); + ASSERT_TRUE(doc.set(second_scalar, 7)); + ASSERT_TRUE(doc.set>(vector, values)); + std::vector docs{doc}; + ASSERT_NO_FATAL_FAILURE(ExpectWrite(collection_->insert(docs), 1)); + auto status = collection_->flush(); + ASSERT_TRUE(status.ok()) << status.message(); + status = collection_->create_index(scalar, + std::make_shared()); + ASSERT_TRUE(status.ok()) << status.message(); + status = collection_->create_index( + vector, std::make_shared(MetricType::L2)); + ASSERT_TRUE(status.ok()) << status.message(); + status = collection_->create_index(second_scalar, + std::make_shared()); + ASSERT_TRUE(status.ok()) << status.message(); + ASSERT_NO_FATAL_FAILURE(Reopen()); + + // Fetch supplies the query vector, covering lookup by a newly allowed ID. + auto fetched = collection_->fetch({doc.pk()}); + ASSERT_TRUE(fetched.has_value()) << fetched.error().message(); + ASSERT_EQ(fetched.value().size(), 1u); + const auto stored_vector = + fetched.value().at(doc.pk())->get>(vector); + ASSERT_TRUE(stored_vector.has_value()); + EXPECT_EQ(stored_vector.value(), values); + SearchQuery query; + query.topk_ = 1; + query.target_.field_name_ = vector; + query.target_.set_vector( + std::string(reinterpret_cast(stored_vector->data()), + stored_vector->size() * sizeof(float))); + query.filter_ = scalar + " = 42 AND " + second_scalar + " = 7"; + query.output_fields_ = std::vector{scalar, second_scalar}; + auto matches = collection_->query(query); + ASSERT_TRUE(matches.has_value()) << matches.error().message(); + ASSERT_EQ(matches.value().size(), 1u); + EXPECT_EQ(matches.value()[0]->pk(), doc.pk()); + EXPECT_EQ(matches.value()[0]->get(scalar), 42); + EXPECT_EQ(matches.value()[0]->get(second_scalar), 7); + collection_.reset(); + FileHelper::RemoveDirectory(col_path); + }; + ASSERT_NO_FATAL_FAILURE(check(true)); + ASSERT_NO_FATAL_FAILURE(check(false)); +} + TEST_F(CollectionTest, Feature_Query_General) { auto func = [&](bool enable_mmap, std::string field_name) { FileHelper::RemoveDirectory(col_path); @@ -5215,6 +5471,135 @@ TEST_F(CollectionTest, Feature_MultiQuery_CallbackReranker) { TEST_F(CollectionTest, Feature_GroupByQuery) {} +TEST_F(CollectionTest, Feature_ColumnDDL_DuplicateNamesKeepData) { + auto schema = MakeSchema(); + ASSERT_TRUE(schema + .add_field(std::make_shared( + "other", DataType::INT32, true)) + .ok()); + ASSERT_TRUE(schema + .add_field(std::make_shared( + "embedding", DataType::VECTOR_FP32, 4, true)) + .ok()); + ASSERT_NO_FATAL_FAILURE(Create(schema)); + std::vector docs{MakeDoc("id", 42)}; + ASSERT_TRUE(docs[0].set>("embedding", {1, 2, 3, 4})); + ASSERT_NO_FATAL_FAILURE(ExpectWrite(collection_->insert(docs), 1)); + const auto before = collection_->schema().value(); + for (const std::string name : {"other", "embedding"}) { + SCOPED_TRACE(name); + auto field = std::make_shared(name, DataType::INT32, true); + auto status = collection_->add_column(field, ""); + EXPECT_EQ(status.code(), StatusCode::INVALID_ARGUMENT); + EXPECT_NE(status.message().find("already exists"), std::string::npos); + status = collection_->alter_column("value", name); + EXPECT_EQ(status.code(), StatusCode::INVALID_ARGUMENT); + EXPECT_NE(status.message().find("already exists"), std::string::npos); + status = collection_->alter_column("value", "", field); + EXPECT_EQ(status.code(), StatusCode::ALREADY_EXISTS); + EXPECT_NE(status.message().find("already exists"), std::string::npos); + EXPECT_EQ(collection_->schema().value(), before); + ASSERT_NO_FATAL_FAILURE(ExpectValue("id", 42)); + } + ASSERT_NO_FATAL_FAILURE(Reopen()); + EXPECT_EQ(collection_->schema().value(), before); + ASSERT_NO_FATAL_FAILURE(ExpectValue("id", 42)); +} + +TEST_F(CollectionTest, Feature_ColumnDDL_ReservedNamesKeepData) { + ASSERT_NO_FATAL_FAILURE(Create(MakeSchema())); + std::vector docs{MakeDoc("id", 42)}; + ASSERT_NO_FATAL_FAILURE(ExpectWrite(collection_->insert(docs), 1)); + const auto before = collection_->schema().value(); + const std::string name = "_zvec_uid_"; + auto status = collection_->add_column( + std::make_shared(name, DataType::INT32, true), ""); + EXPECT_EQ(status.code(), StatusCode::INVALID_ARGUMENT); + EXPECT_NE(status.message().find("is reserved"), std::string::npos); + status = collection_->alter_column("value", name); + EXPECT_EQ(status.code(), StatusCode::INVALID_ARGUMENT); + EXPECT_NE(status.message().find("is reserved"), std::string::npos); + status = collection_->alter_column( + "value", "", std::make_shared(name, DataType::INT32)); + EXPECT_EQ(status.code(), StatusCode::INVALID_ARGUMENT); + EXPECT_NE(status.message().find("is reserved"), std::string::npos); + EXPECT_EQ(collection_->schema().value(), before); + ASSERT_NO_FATAL_FAILURE(ExpectValue("id", 42)); + ASSERT_NO_FATAL_FAILURE(Reopen()); + EXPECT_EQ(collection_->schema().value(), before); + ASSERT_NO_FATAL_FAILURE(ExpectValue("id", 42)); +} + +TEST_F(CollectionTest, Feature_ColumnDDL_FieldCountLimits) { + CollectionSchema schema("x"); + for (uint32_t i = 0; i < kMaxScalarFieldSize; ++i) { + ASSERT_TRUE(schema + .add_field(std::make_shared( + "f" + std::to_string(i), DataType::INT32, true)) + .ok()); + } + ASSERT_NO_FATAL_FAILURE(Create(schema)); + auto status = collection_->add_column( + std::make_shared("excess", DataType::INT32, true), ""); + EXPECT_EQ(status.code(), StatusCode::INVALID_ARGUMENT); + EXPECT_NE(status.message().find("1024 scalar fields"), std::string::npos); + EXPECT_EQ(collection_->schema().value(), schema); + EXPECT_FALSE(collection_->schema().value().has_field("excess")); + status = collection_->destroy(); + ASSERT_TRUE(status.ok()) << status.message(); + collection_.reset(); + + ASSERT_NO_FATAL_FAILURE(Create(MakeSchema())); + std::vector docs{MakeDoc("id", 42)}; + ASSERT_NO_FATAL_FAILURE(ExpectWrite(collection_->insert(docs), 1)); + status = collection_->drop_column("value"); + EXPECT_EQ(status.code(), StatusCode::INVALID_ARGUMENT); + EXPECT_NE(status.message().find("last field"), std::string::npos); + ASSERT_NO_FATAL_FAILURE(ExpectValue("id", 42)); + ASSERT_NO_FATAL_FAILURE(Reopen()); + EXPECT_TRUE(collection_->schema().value().has_field("value")); + ASSERT_NO_FATAL_FAILURE(ExpectValue("id", 42)); +} + +TEST_F(CollectionTest, Feature_ColumnDDL_CopiesCallerSchema) { + ASSERT_NO_FATAL_FAILURE(Create(MakeSchema())); + std::vector initial{MakeDoc("original", 1)}; + ASSERT_NO_FATAL_FAILURE(ExpectWrite(collection_->insert(initial), 1)); + auto added = std::make_shared("extra", DataType::INT32, true); + auto status = collection_->add_column(added, ""); + ASSERT_TRUE(status.ok()) << status.message(); + added->set_name("bad name"); + added->set_data_type(DataType::STRING); + added->set_nullable(false); + + auto current = collection_->schema().value(); + ASSERT_TRUE(current.has_field("extra")); + EXPECT_EQ(current.get_field("extra")->data_type(), DataType::INT32); + EXPECT_TRUE(current.get_field("extra")->nullable()); + EXPECT_FALSE(current.has_field("bad name")); + Doc doc = MakeDoc("new", 2); + ASSERT_TRUE(doc.set("extra", 7)); + std::vector docs{doc}; + ASSERT_NO_FATAL_FAILURE(ExpectWrite(collection_->insert(docs), 1)); + + auto altered = std::make_shared("extra", DataType::INT64, true); + status = collection_->alter_column("extra", "", altered); + ASSERT_TRUE(status.ok()) << status.message(); + altered->set_name("another bad name"); + altered->set_data_type(DataType::STRING); + current = collection_->schema().value(); + ASSERT_TRUE(current.has_field("extra")); + EXPECT_EQ(current.get_field("extra")->data_type(), DataType::INT64); + EXPECT_FALSE(current.has_field("another bad name")); + ASSERT_NO_FATAL_FAILURE(Reopen()); + ASSERT_NO_FATAL_FAILURE(ExpectValue("original", 1)); + ASSERT_NO_FATAL_FAILURE(ExpectValue("new", 2)); + auto fetched = collection_->fetch({"new"}); + ASSERT_TRUE(fetched.has_value()) << fetched.error().message(); + ASSERT_NE(fetched.value().at("new"), nullptr); + EXPECT_EQ(fetched.value().at("new")->get("extra"), 7); +} + TEST_F(CollectionTest, Feature_AddColumn_General) { auto func = [&](bool enable_mmap) { FileHelper::RemoveDirectory(col_path); @@ -5438,6 +5823,34 @@ TEST_F(CollectionTest, Feature_DropColumn_General) { ASSERT_TRUE(!new_schema.has_field("int32")); } +TEST_F(CollectionTest, Feature_AlterColumn_InvalidNameKeepsData) { + ASSERT_NO_FATAL_FAILURE(Create(MakeSchema())); + std::vector docs{MakeDoc("id", 42)}; + ASSERT_NO_FATAL_FAILURE(ExpectWrite(collection_->insert(docs), 1)); + const auto before = collection_->schema().value(); + for (const auto &name : std::vector{ + "user name", "../value", u8"字段", std::string(65, 'f')}) { + SCOPED_TRACE(name); + auto status = collection_->alter_column("value", name); + ASSERT_EQ(status.code(), StatusCode::INVALID_ARGUMENT); + EXPECT_EQ(status.message().find("Invalid schema:"), 0u); + EXPECT_EQ(status.message().find("offset"), std::string::npos); + EXPECT_EQ(collection_->schema().value(), before); + ASSERT_NO_FATAL_FAILURE(ExpectValue("id", 42)); + } + ASSERT_NO_FATAL_FAILURE(Reopen()); + EXPECT_EQ(collection_->schema().value(), before); + ASSERT_NO_FATAL_FAILURE(ExpectValue("id", 42)); + + const std::string renamed(64, 'r'); + auto status = collection_->alter_column("value", renamed); + ASSERT_TRUE(status.ok()) << status.message(); + ASSERT_NO_FATAL_FAILURE(Reopen()); + EXPECT_FALSE(collection_->schema().value().has_field("value")); + EXPECT_TRUE(collection_->schema().value().has_field(renamed)); + ASSERT_NO_FATAL_FAILURE(ExpectValue("id", 42, renamed)); +} + TEST_F(CollectionTest, Feature_AlterColumn_General) { // create collection int doc_count = 1000; diff --git a/tests/db/crash_recovery/relaxed_validation_recovery_test.cc b/tests/db/crash_recovery/relaxed_validation_recovery_test.cc deleted file mode 100644 index b29d5f07e..000000000 --- a/tests/db/crash_recovery/relaxed_validation_recovery_test.cc +++ /dev/null @@ -1,107 +0,0 @@ -// Copyright 2025-present the zvec project -// -// Licensed under the Apache License, Version 2.0 (the "License"); -// you may not use this file except in compliance with the License. -// You may obtain a copy of the License at -// -// http://www.apache.org/licenses/LICENSE-2.0 -// -// Unless required by applicable law or agreed to in writing, software -// distributed under the License is distributed on an "AS IS" BASIS, -// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -// See the License for the specific language governing permissions and -// limitations under the License. - -#include -#include -#include -#include -#include -#include -#include -#include -#include - -namespace zvec { -namespace { - -class RelaxedValidationDeathTest : public ::testing::Test { - protected: - void SetUp() override { - // Re-exec children start with the default style; select threadsafe before - // InDeathTestChild() interprets their internal death-test flag. - ::testing::FLAGS_gtest_death_test_style = "threadsafe"; - // Threadsafe death tests re-exec the fixture. A second crash child must - // reopen the first child's database instead of deleting it in SetUp. - if (!::testing::internal::InDeathTestChild()) { - ailego::FileHelper::RemovePath(path_.c_str()); - } - } - - void TearDown() override { - ailego::FileHelper::RemovePath(path_.c_str()); - } - - const std::string path_{"relaxed_validation_recovery_db"}; -}; - -TEST_F(RelaxedValidationDeathTest, Utf8AndLongIdsRecoverFromUnflushedWal) { - const std::vector ids{u8"订单:😀", - std::string(1021, 'x') + u8"中", " doc ", - u8"café", u8"cafe\u0301"}; - // Re-exec the child before starting collection threads. Exit without stack - // unwinding so Collection destruction cannot flush the writing segment. - ::testing::FLAGS_gtest_death_test_style = "threadsafe"; - ASSERT_EXIT( - { - CollectionSchema schema(u8"恢复 集合"); - if (!schema - .add_field(std::make_shared( - "value", DataType::INT32, false)) - .ok()) { - std::_Exit(1); - } - auto created = - Collection::CreateAndOpen(path_, schema, CollectionOptions{}); - if (!created.has_value()) std::_Exit(2); - auto collection = std::move(created).value(); - std::vector docs; - for (size_t i = 0; i < ids.size(); ++i) { - Doc doc; - doc.set_pk(ids[i]); - doc.set("value", static_cast(i)); - docs.push_back(std::move(doc)); - } - auto inserted = collection->insert(docs); - if (!inserted.has_value()) std::_Exit(3); - for (const auto &status : inserted.value()) { - if (!status.ok()) std::_Exit(4); - } - std::_Exit(0); - }, - ::testing::ExitedWithCode(0), ""); - - auto opened = Collection::Open(path_, CollectionOptions{}); - ASSERT_TRUE(opened.has_value()) << opened.error().message(); - auto collection = std::move(opened).value(); - EXPECT_EQ(collection->schema().value().name(), u8"恢复 集合"); - auto fetched = collection->fetch(ids); - ASSERT_TRUE(fetched.has_value()) << fetched.error().message(); - ASSERT_EQ(fetched.value().size(), ids.size()); - for (size_t i = 0; i < ids.size(); ++i) { - const auto found = fetched.value().find(ids[i]); - ASSERT_NE(found, fetched.value().end()); - ASSERT_NE(found->second, nullptr); - EXPECT_EQ(found->second->pk_ref(), ids[i]); - EXPECT_EQ(found->second->get("value"), static_cast(i)); - } - auto status = collection->flush(); - ASSERT_TRUE(status.ok()) << status.message(); - collection.reset(); - auto reopened = Collection::Open(path_, CollectionOptions{}); - ASSERT_TRUE(reopened.has_value()) << reopened.error().message(); - EXPECT_EQ(reopened.value()->stats().value().doc_count, ids.size()); -} - -} // namespace -} // namespace zvec diff --git a/tests/db/crash_recovery/write_recovery_test.cc b/tests/db/crash_recovery/write_recovery_test.cc index 467ed623c..614fb71ea 100644 --- a/tests/db/crash_recovery/write_recovery_test.cc +++ b/tests/db/crash_recovery/write_recovery_test.cc @@ -14,8 +14,12 @@ #include +#include +#include #include +#include #include +#include #include #include #include @@ -161,6 +165,64 @@ class CrashRecoveryTest : public ::testing::Test { }; +TEST_F(CrashRecoveryTest, Utf8AndLongIdsRecoverFromUnflushedWal) { + const std::vector ids{u8"订单:😀", + std::string(1021, 'x') + u8"中", " doc ", + u8"café", u8"cafe\u0301"}; + // Re-exec the child before starting collection threads. Exit without stack + // unwinding so Collection destruction cannot flush the writing segment. + ::testing::FLAGS_gtest_death_test_style = "threadsafe"; + ASSERT_EXIT( + { + CollectionSchema schema(u8"恢复 集合"); + if (!schema + .add_field(std::make_shared( + "value", DataType::INT32, false)) + .ok()) { + std::_Exit(1); + } + auto created = + Collection::CreateAndOpen(dir_path_, schema, CollectionOptions{}); + if (!created.has_value()) std::_Exit(2); + auto collection = std::move(created).value(); + std::vector docs; + for (size_t i = 0; i < ids.size(); ++i) { + Doc doc; + doc.set_pk(ids[i]); + doc.set("value", static_cast(i)); + docs.push_back(std::move(doc)); + } + auto inserted = collection->insert(docs); + if (!inserted.has_value()) std::_Exit(3); + for (const auto &status : inserted.value()) { + if (!status.ok()) std::_Exit(4); + } + std::_Exit(0); + }, + ::testing::ExitedWithCode(0), ""); + + auto opened = Collection::Open(dir_path_, CollectionOptions{}); + ASSERT_TRUE(opened.has_value()) << opened.error().message(); + auto collection = std::move(opened).value(); + EXPECT_EQ(collection->schema().value().name(), u8"恢复 集合"); + auto fetched = collection->fetch(ids); + ASSERT_TRUE(fetched.has_value()) << fetched.error().message(); + ASSERT_EQ(fetched.value().size(), ids.size()); + for (size_t i = 0; i < ids.size(); ++i) { + const auto found = fetched.value().find(ids[i]); + ASSERT_NE(found, fetched.value().end()); + ASSERT_NE(found->second, nullptr); + EXPECT_EQ(found->second->pk_ref(), ids[i]); + EXPECT_EQ(found->second->get("value"), static_cast(i)); + } + auto status = collection->flush(); + ASSERT_TRUE(status.ok()) << status.message(); + collection.reset(); + auto reopened = Collection::Open(dir_path_, CollectionOptions{}); + ASSERT_TRUE(reopened.has_value()) << reopened.error().message(); + EXPECT_EQ(reopened.value()->stats().value().doc_count, ids.size()); +} + TEST_F(CrashRecoveryTest, BasicInsertAndReopen) { { auto schema = CreateTestSchema(collection_name_); diff --git a/tests/db/index/common/manifest_codec_golden_test.cc b/tests/db/index/common/manifest_codec_golden_test.cc index 1bde56b1e..329f6199e 100644 --- a/tests/db/index/common/manifest_codec_golden_test.cc +++ b/tests/db/index/common/manifest_codec_golden_test.cc @@ -1081,68 +1081,6 @@ TEST(ManifestCodecGolden, IndexParamsLastBranchWins) { EXPECT_EQ(decoded->type(), IndexType::HNSW); } -TEST(ManifestCodecGolden, LegacyFieldNamesSurviveWithoutRevalidation) { - // Older alter_column paths could persist names outside the schema rules. - // Build the wire data without schema helpers so this also catches a decoder - // silently dropping fields after a new validation check is added. - const std::vector names{"user name", std::string(65, 'f'), - u8"历史字段", "_zvec_uid_"}; - std::string encoded_schema; - pbwire::Writer schema_writer(&encoded_schema); - schema_writer.PutString(1, "legacy_collection"); - for (const auto &name : names) { - std::string encoded_field; - pbwire::Writer field_writer(&encoded_field); - field_writer.PutString(1, name); - field_writer.PutVarint(2, 2); // Persisted STRING data type. - schema_writer.PutMessage(2, encoded_field); - } - schema_writer.PutVarint(3, 10000); - std::string encoded_manifest; - pbwire::Writer(&encoded_manifest).PutMessage(2, encoded_schema); - - ManifestData restored; - auto status = ManifestCodec::Decode(encoded_manifest, &restored); - ASSERT_TRUE(status.ok()) << status.message(); - ASSERT_NE(restored.schema, nullptr); - ASSERT_EQ(restored.schema->fields().size(), names.size()); - for (size_t i = 0; i < names.size(); ++i) { - SCOPED_TRACE(i); - const auto *field = restored.schema->get_field(names[i]); - ASSERT_NE(field, nullptr); - EXPECT_EQ(field->name(), names[i]); - EXPECT_EQ(field->data_type(), DataType::STRING); - EXPECT_EQ(restored.schema->fields()[i]->name(), names[i]); - } - EXPECT_EQ(restored.schema->validate().code(), StatusCode::INVALID_ARGUMENT); - - std::string reencoded; - status = ManifestCodec::Encode(restored, &reencoded); - ASSERT_TRUE(status.ok()) << status.message(); - EXPECT_EQ(reencoded, encoded_manifest); -} - -TEST(ManifestCodecGolden, StructuralSchemaHelpersPreserveLegacyNames) { - CollectionSchema schema("legacy_collection"); - auto status = schema.add_field( - std::make_shared("user name", DataType::STRING)); - ASSERT_TRUE(status.ok()) << status.message(); - status = schema.alter_field("user name", std::make_shared( - u8"历史字段", DataType::STRING)); - ASSERT_TRUE(status.ok()) << status.message(); - EXPECT_FALSE(schema.has_field("user name")); - ASSERT_TRUE(schema.has_field(u8"历史字段")); - - // Structural mutation is also used during recovery. Explicit validation is - // kept separate, and legacy names can still be replaced with legal names. - EXPECT_EQ(schema.validate().code(), StatusCode::INVALID_ARGUMENT); - status = schema.alter_field( - u8"历史字段", std::make_shared("renamed", DataType::STRING)); - ASSERT_TRUE(status.ok()) << status.message(); - EXPECT_TRUE(schema.validate().ok()); - EXPECT_TRUE(schema.has_field("renamed")); -} - TEST(ManifestCodecGolden, UnknownFieldsAreIgnored) { // Forward compatibility: a manifest written by a newer zvec may carry fields // this build does not know about. They must be skipped silently. diff --git a/tests/db/relaxed_validation_test.cc b/tests/db/relaxed_validation_test.cc deleted file mode 100644 index 1515aa5dd..000000000 --- a/tests/db/relaxed_validation_test.cc +++ /dev/null @@ -1,485 +0,0 @@ -// Copyright 2025-present the zvec project -// -// Licensed under the Apache License, Version 2.0 (the "License"); -// you may not use this file except in compliance with the License. -// You may obtain a copy of the License at -// -// http://www.apache.org/licenses/LICENSE-2.0 -// -// Unless required by applicable law or agreed to in writing, software -// distributed under the License is distributed on an "AS IS" BASIS, -// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -// See the License for the specific language governing permissions and -// limitations under the License. - -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include "db/common/constants.h" - -namespace zvec { -namespace { - -class RelaxedValidationTest : public ::testing::Test { - protected: - void SetUp() override { - ailego::MemoryLimitPool::get_instance().init(2 * 1024ll * 1024ll * 1024ll); - ailego::FileHelper::RemovePath(path_.c_str()); - } - - void TearDown() override { - collection_.reset(); - ailego::FileHelper::RemovePath(path_.c_str()); - } - - CollectionSchema MakeSchema(const std::string &name = "x", - const std::string &field = "value") { - CollectionSchema schema(name); - EXPECT_TRUE(schema - .add_field(std::make_shared( - field, DataType::INT32, false)) - .ok()); - return schema; - } - - void Create(const CollectionSchema &schema) { - auto result = Collection::CreateAndOpen(path_, schema, options_); - ASSERT_TRUE(result.has_value()) << result.error().message(); - collection_ = std::move(result).value(); - } - - void Reopen() { - collection_.reset(); - auto result = Collection::Open(path_, options_); - ASSERT_TRUE(result.has_value()) << result.error().message(); - collection_ = std::move(result).value(); - } - - Doc MakeDoc(const std::string &id, int32_t value, - const std::string &field = "value") { - Doc doc; - doc.set_pk(id); - EXPECT_TRUE(doc.set(field, value)); - return doc; - } - - void ExpectWrite(const Result &result, size_t count) { - ASSERT_TRUE(result.has_value()) << result.error().message(); - ASSERT_EQ(result.value().size(), count); - for (const auto &status : result.value()) { - ASSERT_TRUE(status.ok()) << status.message(); - } - } - - void ExpectValue(const std::string &id, int32_t expected, - const std::string &field = "value") { - auto result = collection_->fetch({id}); - ASSERT_TRUE(result.has_value()) << result.error().message(); - ASSERT_EQ(result.value().size(), 1u); - auto found = result.value().find(id); - ASSERT_NE(found, result.value().end()); - ASSERT_NE(found->second, nullptr); - EXPECT_EQ(found->second->pk(), id); - EXPECT_EQ(found->second->get(field), expected); - } - - const std::string path_{"relaxed_validation_test_db"}; - CollectionOptions options_; - Collection::Ptr collection_; -}; - -TEST_F(RelaxedValidationTest, Utf8IdsKeepTheirIdentityAcrossCrudAndReopen) { - const std::string name = u8"测试 集合/v1"; - ASSERT_NO_FATAL_FAILURE(Create(MakeSchema(name))); - const std::vector ids = { - "user:123", "https://example.com/document/42", - u8"订单-😀", std::string(1021, 'i') + u8"中", - "doc", " doc", - "doc ", " ", - u8"café", u8"cafe\u0301", - "DOC"}; - std::vector docs; - for (size_t i = 0; i < ids.size(); ++i) { - docs.push_back(MakeDoc(ids[i], static_cast(i))); - } - ASSERT_NO_FATAL_FAILURE(ExpectWrite(collection_->insert(docs), docs.size())); - - std::vector updates{MakeDoc(ids[0], 100)}; - ASSERT_NO_FATAL_FAILURE(ExpectWrite(collection_->update(updates), 1)); - std::vector upserts{MakeDoc(ids[3], 103), MakeDoc(u8"新增:文档", 200)}; - ASSERT_NO_FATAL_FAILURE(ExpectWrite(collection_->upsert(upserts), 2)); - ASSERT_NO_FATAL_FAILURE(ExpectWrite(collection_->delete_({ids[1]}), 1)); - - auto flush_status = collection_->flush(); - ASSERT_TRUE(flush_status.ok()) << flush_status.message(); - ASSERT_NO_FATAL_FAILURE(Reopen()); - auto schema = collection_->schema(); - ASSERT_TRUE(schema.has_value()) << schema.error().message(); - EXPECT_EQ(schema.value().name(), name); - for (size_t i = 0; i < ids.size(); ++i) { - if (i == 1) { - continue; - } - int32_t expected = static_cast(i); - if (i == 0) expected = 100; - if (i == 3) expected = 103; - ASSERT_NO_FATAL_FAILURE(ExpectValue(ids[i], expected)); - } - ASSERT_NO_FATAL_FAILURE(ExpectValue(u8"新增:文档", 200)); - auto deleted = collection_->fetch({ids[1]}); - ASSERT_TRUE(deleted.has_value()) << deleted.error().message(); - ASSERT_EQ(deleted.value().size(), 1u); - EXPECT_EQ(deleted.value().at(ids[1]), nullptr); -} - -TEST_F(RelaxedValidationTest, ShortAndMaximumLengthCollectionNamesPersist) { - for (const auto &name : std::vector{ - "x", "xy", u8"集", std::string(253, 'n') + u8"集"}) { - SCOPED_TRACE(name); - ASSERT_NO_FATAL_FAILURE(Create(MakeSchema(name))); - std::vector docs{MakeDoc("id", 1)}; - ASSERT_NO_FATAL_FAILURE(ExpectWrite(collection_->insert(docs), 1)); - auto status = collection_->flush(); - ASSERT_TRUE(status.ok()) << status.message(); - ASSERT_NO_FATAL_FAILURE(Reopen()); - auto schema = collection_->schema(); - ASSERT_TRUE(schema.has_value()) << schema.error().message(); - EXPECT_EQ(schema.value().name(), name); - ASSERT_NO_FATAL_FAILURE(ExpectValue("id", 1)); - status = collection_->destroy(); - ASSERT_TRUE(status.ok()) << status.message(); - collection_.reset(); - } -} - -TEST_F(RelaxedValidationTest, MaximumLengthFieldsSupportIndexesAndFilters) { - const std::string scalar = "s" + std::string(63, 'a'); - const std::string second_scalar = scalar.substr(0, 63) + "b"; - const std::string vector = "v" + std::string(63, 'b'); - auto schema = MakeSchema("x", scalar); - ASSERT_TRUE(schema - .add_field(std::make_shared( - second_scalar, DataType::INT32, false)) - .ok()); - ASSERT_TRUE(schema - .add_field(std::make_shared( - vector, DataType::VECTOR_FP32, 4, false, - std::make_shared(MetricType::L2))) - .ok()); - ASSERT_NO_FATAL_FAILURE(Create(schema)); - const std::vector values{1.0f, 2.0f, 3.0f, 4.0f}; - Doc doc = MakeDoc(u8"文档:1", 42, scalar); - ASSERT_TRUE(doc.set(second_scalar, 7)); - ASSERT_TRUE(doc.set>(vector, values)); - std::vector docs{doc}; - ASSERT_NO_FATAL_FAILURE(ExpectWrite(collection_->insert(docs), 1)); - auto status = collection_->flush(); - ASSERT_TRUE(status.ok()) << status.message(); - status = - collection_->create_index(scalar, std::make_shared()); - ASSERT_TRUE(status.ok()) << status.message(); - status = collection_->create_index( - vector, std::make_shared(MetricType::L2)); - ASSERT_TRUE(status.ok()) << status.message(); - status = collection_->create_index(second_scalar, - std::make_shared()); - ASSERT_TRUE(status.ok()) << status.message(); - ASSERT_NO_FATAL_FAILURE(Reopen()); - - // Fetch supplies the query vector, covering lookup by a newly allowed ID. - auto fetched = collection_->fetch({doc.pk()}); - ASSERT_TRUE(fetched.has_value()) << fetched.error().message(); - ASSERT_EQ(fetched.value().size(), 1u); - const auto stored_vector = - fetched.value().at(doc.pk())->get>(vector); - ASSERT_TRUE(stored_vector.has_value()); - EXPECT_EQ(stored_vector.value(), values); - SearchQuery query; - query.topk_ = 1; - query.target_.field_name_ = vector; - query.target_.set_vector( - std::string(reinterpret_cast(stored_vector->data()), - stored_vector->size() * sizeof(float))); - query.filter_ = scalar + " = 42 AND " + second_scalar + " = 7"; - query.output_fields_ = std::vector{scalar, second_scalar}; - auto matches = collection_->query(query); - ASSERT_TRUE(matches.has_value()) << matches.error().message(); - ASSERT_EQ(matches.value().size(), 1u); - EXPECT_EQ(matches.value()[0]->pk(), doc.pk()); - EXPECT_EQ(matches.value()[0]->get(scalar), 42); - EXPECT_EQ(matches.value()[0]->get(second_scalar), 7); -} - -TEST_F(RelaxedValidationTest, MaximumLengthFieldsPersistWithBufferedStorage) { - options_.enable_mmap_ = false; - const std::string field(64, 'f'); - ASSERT_NO_FATAL_FAILURE(Create(MakeSchema("buffered", field))); - std::vector docs{MakeDoc("doc", 42, field)}; - ASSERT_NO_FATAL_FAILURE(ExpectWrite(collection_->insert(docs), 1)); - const auto status = collection_->flush(); - ASSERT_TRUE(status.ok()) << status.message(); - ASSERT_NO_FATAL_FAILURE(Reopen()); - EXPECT_TRUE(collection_->schema().value().has_field(field)); - ASSERT_NO_FATAL_FAILURE(ExpectValue("doc", 42, field)); -} - -TEST_F(RelaxedValidationTest, InvalidRenameLeavesSchemaAndDataUnchanged) { - ASSERT_NO_FATAL_FAILURE(Create(MakeSchema())); - std::vector docs{MakeDoc("id", 42)}; - ASSERT_NO_FATAL_FAILURE(ExpectWrite(collection_->insert(docs), 1)); - const auto before = collection_->schema().value(); - for (const auto &name : std::vector{ - "user name", "../value", u8"字段", std::string(65, 'f')}) { - SCOPED_TRACE(name); - auto status = collection_->alter_column("value", name); - ASSERT_EQ(status.code(), StatusCode::INVALID_ARGUMENT); - EXPECT_EQ(status.message().find("Invalid schema:"), 0u); - EXPECT_EQ(status.message().find("offset"), std::string::npos); - EXPECT_EQ(collection_->schema().value(), before); - ASSERT_NO_FATAL_FAILURE(ExpectValue("id", 42)); - } - ASSERT_NO_FATAL_FAILURE(Reopen()); - EXPECT_EQ(collection_->schema().value(), before); - ASSERT_NO_FATAL_FAILURE(ExpectValue("id", 42)); - - const std::string renamed(64, 'r'); - auto status = collection_->alter_column("value", renamed); - ASSERT_TRUE(status.ok()) << status.message(); - ASSERT_NO_FATAL_FAILURE(Reopen()); - EXPECT_FALSE(collection_->schema().value().has_field("value")); - EXPECT_TRUE(collection_->schema().value().has_field(renamed)); - ASSERT_NO_FATAL_FAILURE(ExpectValue("id", 42, renamed)); -} - -TEST_F(RelaxedValidationTest, InvalidIdRejectsWholeBatchBeforeWriting) { - ASSERT_NO_FATAL_FAILURE(Create(MakeSchema())); - std::vector initial{MakeDoc("existing", 1)}; - ASSERT_NO_FATAL_FAILURE(ExpectWrite(collection_->insert(initial), 1)); - for (int operation = 0; operation < 3; ++operation) { - SCOPED_TRACE(operation); - const std::string first_id = operation == 0 ? "new:id" : "existing"; - std::vector batch{MakeDoc(first_id, 99), - MakeDoc(std::string("bad\0id", 6), 100)}; - auto result = operation == 0 ? collection_->insert(batch) - : operation == 1 ? collection_->update(batch) - : collection_->upsert(batch); - ASSERT_FALSE(result.has_value()); - EXPECT_EQ(result.error().code(), StatusCode::INVALID_ARGUMENT); - EXPECT_EQ(result.error().message().find("Invalid doc:"), 0u); - EXPECT_NE(result.error().message().find("null character"), - std::string::npos); - EXPECT_NE(result.error().message().find("id[bad\\0id]"), std::string::npos); - EXPECT_EQ(result.error().message().find("offset"), std::string::npos); - ASSERT_NO_FATAL_FAILURE(ExpectValue("existing", 1)); - auto missing = collection_->fetch({"new:id"}); - ASSERT_TRUE(missing.has_value()) << missing.error().message(); - ASSERT_EQ(missing.value().size(), 1u); - EXPECT_EQ(missing.value().at("new:id"), nullptr); - } - ASSERT_NO_FATAL_FAILURE(Reopen()); - ASSERT_NO_FATAL_FAILURE(ExpectValue("existing", 1)); - EXPECT_EQ(collection_->stats().value().doc_count, 1u); -} - -TEST_F(RelaxedValidationTest, FetchAndDeleteKeepLookupSemantics) { - ASSERT_NO_FATAL_FAILURE(Create(MakeSchema())); - // These are invalid for a new document, but lookup must retain its existing - // missing-key behavior rather than introducing input validation errors. - const std::vector absent_ids{"", std::string("bad\0id", 6), - std::string(1025, 'x'), "\xff"}; - auto fetched = collection_->fetch(absent_ids); - ASSERT_TRUE(fetched.has_value()) << fetched.error().message(); - ASSERT_EQ(fetched.value().size(), absent_ids.size()); - for (const auto &id : absent_ids) { - EXPECT_EQ(fetched.value().at(id), nullptr); - } - auto deleted = collection_->delete_(absent_ids); - ASSERT_TRUE(deleted.has_value()) << deleted.error().message(); - ASSERT_EQ(deleted.value().size(), absent_ids.size()); - for (const auto &status : deleted.value()) { - EXPECT_EQ(status.code(), StatusCode::NOT_FOUND); - } -} - -TEST_F(RelaxedValidationTest, ReservedNamesAndDuplicatesFailBeforeCreation) { - for (const std::string name : - {"_zvec_uid_", "_zvec_g_doc_id_", "_zvec_row_id_", "_zvec_score", - "_zvec_group_id"}) { - SCOPED_TRACE(name); - auto result = - Collection::CreateAndOpen(path_, MakeSchema("x", name), options_); - ASSERT_FALSE(result.has_value()); - EXPECT_EQ(result.error().code(), StatusCode::INVALID_ARGUMENT); - EXPECT_NE(result.error().message().find("is reserved"), std::string::npos); - EXPECT_FALSE(ailego::FileHelper::IsExist(path_.c_str())); - } - auto scalar = std::make_shared("value", DataType::INT32); - auto other_scalar = std::make_shared("value", DataType::INT64); - auto vector = - std::make_shared("value", DataType::VECTOR_FP32, 4, false); - auto other_vector = - std::make_shared("value", DataType::VECTOR_FP32, 8, false); - const std::vector cases{ - {scalar, std::make_shared(*scalar)}, - {scalar, other_scalar}, - {scalar, scalar}, - {vector, other_vector}, - {scalar, vector}, - {vector, scalar}}; - for (size_t i = 0; i < cases.size(); ++i) { - SCOPED_TRACE(i); - CollectionSchema duplicate("x", cases[i]); - auto result = Collection::CreateAndOpen(path_, duplicate, options_); - ASSERT_FALSE(result.has_value()); - EXPECT_EQ(result.error().code(), StatusCode::INVALID_ARGUMENT); - EXPECT_NE(result.error().message().find("duplicate field name [value]"), - std::string::npos); - EXPECT_FALSE(ailego::FileHelper::IsExist(path_.c_str())); - } -} - -TEST_F(RelaxedValidationTest, DuplicateDdlTargetsLeaveSchemaAndDataUnchanged) { - auto schema = MakeSchema(); - ASSERT_TRUE(schema - .add_field(std::make_shared( - "other", DataType::INT32, true)) - .ok()); - ASSERT_TRUE(schema - .add_field(std::make_shared( - "embedding", DataType::VECTOR_FP32, 4, true)) - .ok()); - ASSERT_NO_FATAL_FAILURE(Create(schema)); - std::vector docs{MakeDoc("id", 42)}; - ASSERT_TRUE(docs[0].set>("embedding", {1, 2, 3, 4})); - ASSERT_NO_FATAL_FAILURE(ExpectWrite(collection_->insert(docs), 1)); - const auto before = collection_->schema().value(); - for (const std::string name : {"other", "embedding"}) { - SCOPED_TRACE(name); - auto field = std::make_shared(name, DataType::INT32, true); - auto status = collection_->add_column(field, ""); - EXPECT_EQ(status.code(), StatusCode::INVALID_ARGUMENT); - EXPECT_NE(status.message().find("already exists"), std::string::npos); - status = collection_->alter_column("value", name); - EXPECT_EQ(status.code(), StatusCode::INVALID_ARGUMENT); - EXPECT_NE(status.message().find("already exists"), std::string::npos); - status = collection_->alter_column("value", "", field); - EXPECT_EQ(status.code(), StatusCode::ALREADY_EXISTS); - EXPECT_NE(status.message().find("already exists"), std::string::npos); - EXPECT_EQ(collection_->schema().value(), before); - ASSERT_NO_FATAL_FAILURE(ExpectValue("id", 42)); - } - ASSERT_NO_FATAL_FAILURE(Reopen()); - EXPECT_EQ(collection_->schema().value(), before); - ASSERT_NO_FATAL_FAILURE(ExpectValue("id", 42)); -} - -TEST_F(RelaxedValidationTest, ReservedDdlTargetsLeaveTheCollectionUnchanged) { - ASSERT_NO_FATAL_FAILURE(Create(MakeSchema())); - std::vector docs{MakeDoc("id", 42)}; - ASSERT_NO_FATAL_FAILURE(ExpectWrite(collection_->insert(docs), 1)); - const auto before = collection_->schema().value(); - for (const std::string name : - {"_zvec_uid_", "_zvec_g_doc_id_", "_zvec_row_id_", "_zvec_score", - "_zvec_group_id"}) { - SCOPED_TRACE(name); - auto status = collection_->add_column( - std::make_shared(name, DataType::INT32, true), ""); - EXPECT_EQ(status.code(), StatusCode::INVALID_ARGUMENT); - EXPECT_NE(status.message().find("is reserved"), std::string::npos); - status = collection_->alter_column("value", name); - EXPECT_EQ(status.code(), StatusCode::INVALID_ARGUMENT); - EXPECT_NE(status.message().find("is reserved"), std::string::npos); - status = collection_->alter_column( - "value", "", std::make_shared(name, DataType::INT32)); - EXPECT_EQ(status.code(), StatusCode::INVALID_ARGUMENT); - EXPECT_NE(status.message().find("is reserved"), std::string::npos); - EXPECT_EQ(collection_->schema().value(), before); - ASSERT_NO_FATAL_FAILURE(ExpectValue("id", 42)); - } - ASSERT_NO_FATAL_FAILURE(Reopen()); - EXPECT_EQ(collection_->schema().value(), before); - ASSERT_NO_FATAL_FAILURE(ExpectValue("id", 42)); -} - -TEST_F(RelaxedValidationTest, DdlRetainsFieldCountInvariants) { - CollectionSchema schema("x"); - for (uint32_t i = 0; i < kMaxScalarFieldSize; ++i) { - ASSERT_TRUE(schema - .add_field(std::make_shared( - "f" + std::to_string(i), DataType::INT32, true)) - .ok()); - } - ASSERT_NO_FATAL_FAILURE(Create(schema)); - auto status = collection_->add_column( - std::make_shared("excess", DataType::INT32, true), ""); - EXPECT_EQ(status.code(), StatusCode::INVALID_ARGUMENT); - EXPECT_NE(status.message().find("1024 scalar fields"), std::string::npos); - EXPECT_EQ(collection_->schema().value(), schema); - EXPECT_FALSE(collection_->schema().value().has_field("excess")); - status = collection_->destroy(); - ASSERT_TRUE(status.ok()) << status.message(); - collection_.reset(); - - ASSERT_NO_FATAL_FAILURE(Create(MakeSchema())); - std::vector docs{MakeDoc("id", 42)}; - ASSERT_NO_FATAL_FAILURE(ExpectWrite(collection_->insert(docs), 1)); - status = collection_->drop_column("value"); - EXPECT_EQ(status.code(), StatusCode::INVALID_ARGUMENT); - EXPECT_NE(status.message().find("last field"), std::string::npos); - ASSERT_NO_FATAL_FAILURE(ExpectValue("id", 42)); - ASSERT_NO_FATAL_FAILURE(Reopen()); - EXPECT_TRUE(collection_->schema().value().has_field("value")); - ASSERT_NO_FATAL_FAILURE(ExpectValue("id", 42)); -} - -TEST_F(RelaxedValidationTest, DdlSnapshotsCallerOwnedFieldSchemas) { - ASSERT_NO_FATAL_FAILURE(Create(MakeSchema())); - std::vector initial{MakeDoc("original", 1)}; - ASSERT_NO_FATAL_FAILURE(ExpectWrite(collection_->insert(initial), 1)); - auto added = std::make_shared("extra", DataType::INT32, true); - auto status = collection_->add_column(added, ""); - ASSERT_TRUE(status.ok()) << status.message(); - added->set_name("bad name"); - added->set_data_type(DataType::STRING); - added->set_nullable(false); - - auto current = collection_->schema().value(); - ASSERT_TRUE(current.has_field("extra")); - EXPECT_EQ(current.get_field("extra")->data_type(), DataType::INT32); - EXPECT_TRUE(current.get_field("extra")->nullable()); - EXPECT_FALSE(current.has_field("bad name")); - Doc doc = MakeDoc("new", 2); - ASSERT_TRUE(doc.set("extra", 7)); - std::vector docs{doc}; - ASSERT_NO_FATAL_FAILURE(ExpectWrite(collection_->insert(docs), 1)); - - auto altered = std::make_shared("extra", DataType::INT64, true); - status = collection_->alter_column("extra", "", altered); - ASSERT_TRUE(status.ok()) << status.message(); - altered->set_name("another bad name"); - altered->set_data_type(DataType::STRING); - current = collection_->schema().value(); - ASSERT_TRUE(current.has_field("extra")); - EXPECT_EQ(current.get_field("extra")->data_type(), DataType::INT64); - EXPECT_FALSE(current.has_field("another bad name")); - ASSERT_NO_FATAL_FAILURE(Reopen()); - ASSERT_NO_FATAL_FAILURE(ExpectValue("original", 1)); - ASSERT_NO_FATAL_FAILURE(ExpectValue("new", 2)); - auto fetched = collection_->fetch({"new"}); - ASSERT_TRUE(fetched.has_value()) << fetched.error().message(); - ASSERT_NE(fetched.value().at("new"), nullptr); - EXPECT_EQ(fetched.value().at("new")->get("extra"), 7); -} - -} // namespace -} // namespace zvec From 8b48b41698f13f63903c35b6d352811acf992401 Mon Sep 17 00:00:00 2001 From: Qinren Zhou Date: Wed, 16 Sep 2026 15:49:49 +0800 Subject: [PATCH 07/11] refactor: simplify Python document conversion errors --- python/tests/test_name_validation.py | 27 +++++---------------------- python/zvec/model/collection.py | 14 ++++++++++---- python/zvec/model/convert.py | 19 ------------------- 3 files changed, 15 insertions(+), 45 deletions(-) diff --git a/python/tests/test_name_validation.py b/python/tests/test_name_validation.py index 97329ef10..b9d584d49 100644 --- a/python/tests/test_name_validation.py +++ b/python/tests/test_name_validation.py @@ -98,11 +98,7 @@ def test_invalid_id_rejects_batch_before_writing(collection, operation, doc_id, message = str(exc_info.value) assert message.startswith("Invalid doc:") assert reason in message - if doc_id == "\ud800": - # Conversion fails in Python before native validation runs. - assert "document at index 1" in message - else: - assert "document at index" not in message + assert "document at index" not in message assert "offset" not in message if doc_id: assert "id[" in message @@ -116,9 +112,7 @@ def test_invalid_id_rejects_batch_before_writing(collection, operation, doc_id, @pytest.mark.parametrize("operation", ["insert", "update", "upsert"]) -def test_conversion_type_error_identifies_document_before_writing( - collection, operation -): +def test_conversion_type_error_rejects_batch_before_writing(collection, operation): if operation == "update": assert collection.insert(zvec.Doc("valid", fields={"text": "before"})).ok() @@ -130,8 +124,8 @@ def test_conversion_type_error_identifies_document_before_writing( getattr(collection, operation)(docs) message = str(exc_info.value) - assert message.endswith(" (document at index 1)") - assert message.count("document at index") == 1 + assert "Field 'text': expected STRING" in message + assert "document at index" not in message fetched = collection.fetch("valid") if operation == "update": assert fetched["valid"].field("text") == "before" @@ -203,11 +197,7 @@ def test_long_field_name_and_rejected_rename_preserve_data(tmp_path): def test_surrogate_id_has_a_readable_encoding_error(collection, operation): with pytest.raises( ValueError, - match="^" - + re.escape( - r"Invalid doc: id['\ud800'] is not valid UTF-8 (document at index 0)" - ) - + "$", + match="^" + re.escape(r"Invalid doc: id['\ud800'] is not valid UTF-8") + "$", ): getattr(collection, operation)(zvec.Doc("\ud800", fields={"text": "value"})) assert collection.stats.doc_count == 0 @@ -271,10 +261,3 @@ def test_duplicate_name_errors_are_escaped_and_bounded(kind, name): assert "\\n" in message else: assert "..." in message - - -def test_native_schema_rejects_null_field_pointer(): - from zvec._zvec.schema import _CollectionSchema - - with pytest.raises(ValueError, match="^Invalid schema:"): - _CollectionSchema("fields", [None]) diff --git a/python/zvec/model/collection.py b/python/zvec/model/collection.py index 898d1f9c1..33620dd3f 100644 --- a/python/zvec/model/collection.py +++ b/python/zvec/model/collection.py @@ -25,7 +25,7 @@ from ..extension import ReRanker from ..typing import Status from ._validation import explain_utf8_conversion_error -from .convert import convert_to_cpp_docs, convert_to_py_doc +from .convert import convert_to_cpp_doc, convert_to_py_doc from .doc import Doc, DocList, GroupResult from .param import ( AddColumnOption, @@ -346,7 +346,9 @@ def insert(self, docs: Union[Doc, list[Doc]]) -> Union[Status, list[Status]]: """ is_single = isinstance(docs, Doc) doc_list = [docs] if is_single else docs - results = self._obj.Insert(convert_to_cpp_docs(doc_list, self.schema)) + results = self._obj.Insert( + [convert_to_cpp_doc(doc, self.schema) for doc in doc_list] + ) return results[0] if is_single else results @overload @@ -369,7 +371,9 @@ def upsert(self, docs: Union[Doc, list[Doc]]) -> Union[Status, list[Status]]: """ is_single = isinstance(docs, Doc) doc_list = [docs] if is_single else docs - results = self._obj.Upsert(convert_to_cpp_docs(doc_list, self.schema)) + results = self._obj.Upsert( + [convert_to_cpp_doc(doc, self.schema) for doc in doc_list] + ) return results[0] if is_single else results @overload @@ -394,7 +398,9 @@ def update(self, docs: Union[Doc, list[Doc]]) -> Union[Status, list[Status]]: """ is_single = isinstance(docs, Doc) doc_list = [docs] if is_single else docs - results = self._obj.Update(convert_to_cpp_docs(doc_list, self.schema)) + results = self._obj.Update( + [convert_to_cpp_doc(doc, self.schema) for doc in doc_list] + ) return results[0] if is_single else results @overload diff --git a/python/zvec/model/convert.py b/python/zvec/model/convert.py index e37ee06ec..8dc1a7c28 100644 --- a/python/zvec/model/convert.py +++ b/python/zvec/model/convert.py @@ -51,25 +51,6 @@ def convert_to_cpp_doc(doc: Doc, collection_schema: CollectionSchema) -> _Doc: return _doc -def convert_to_cpp_docs( - docs: list[Doc], collection_schema: CollectionSchema -) -> list[_Doc]: - converted = [] - for index, doc in enumerate(docs): - try: - converted.append(convert_to_cpp_doc(doc, collection_schema)) - except (TypeError, ValueError) as error: - # Preserve the original exception and cause. Unicode error subclasses - # carry structured arguments that must not be replaced with a string. - if type(error) in (TypeError, ValueError): - suffix = f" (document at index {index})" - message = str(error) - if not message.endswith(suffix): - error.args = (message + suffix,) - raise - return converted - - def convert_to_py_doc(doc: _Doc, collection_schema: CollectionSchema) -> Doc: if not doc or not collection_schema: return None From ccf6250fcd7662e3979473c870c206128cbe85e6 Mon Sep 17 00:00:00 2001 From: Qinren Zhou Date: Wed, 16 Sep 2026 21:40:14 +0800 Subject: [PATCH 08/11] fix(ci): align quantization schema errors with new prefix and skip hardlink test on Android --- python/tests/test_uniform_quantization_support_matrix.py | 2 +- src/db/index/common/schema.cc | 6 ++---- tests/db/collection_test.cc | 4 ++++ 3 files changed, 7 insertions(+), 5 deletions(-) diff --git a/python/tests/test_uniform_quantization_support_matrix.py b/python/tests/test_uniform_quantization_support_matrix.py index 8b851f664..a64281102 100644 --- a/python/tests/test_uniform_quantization_support_matrix.py +++ b/python/tests/test_uniform_quantization_support_matrix.py @@ -49,7 +49,7 @@ def test_uniform_rejects_unsupported_field_types( ) ], ) - with pytest.raises(ValueError, match="schema validate failed:.*quantiz"): + with pytest.raises(ValueError, match="Invalid schema:.*quantiz"): collection = zvec.create_and_open(str(tmp_path / "collection"), schema=schema) collection.close() diff --git a/src/db/index/common/schema.cc b/src/db/index/common/schema.cc index 9e59b62ec..dab3b402c 100644 --- a/src/db/index/common/schema.cc +++ b/src/db/index/common/schema.cc @@ -273,8 +273,7 @@ Status FieldSchema::validate() const { index_params_->type() == IndexType::IVF || index_params_->type() == IndexType::DISKANN)) { return Status::InvalidArgument( - "schema validate failed: ", - QuantizeTypeCodeBook::AsString(quantize_type), + "Invalid schema: ", QuantizeTypeCodeBook::AsString(quantize_type), " quantization is not supported with ", IndexTypeCodeBook::AsString(index_params_->type()), " index, field[", name_, "]"); @@ -282,8 +281,7 @@ Status FieldSchema::validate() const { if (is_uniform && vector_index_params->metric_type() != MetricType::L2) { return Status::InvalidArgument( - "schema validate failed: ", - QuantizeTypeCodeBook::AsString(quantize_type), + "Invalid schema: ", QuantizeTypeCodeBook::AsString(quantize_type), " quantize only supports L2 metric, but field[", name_, "]'s metric is ", MetricTypeCodeBook::AsString(vector_index_params->metric_type())); diff --git a/tests/db/collection_test.cc b/tests/db/collection_test.cc index aa663f6d5..a67788af4 100644 --- a/tests/db/collection_test.cc +++ b/tests/db/collection_test.cc @@ -4598,6 +4598,10 @@ TEST_F(CollectionTest, Feature_Query_Validate) { } TEST_F(CollectionTest, Feature_Query_MaximumLengthFieldNames) { +#ifdef __ANDROID__ + GTEST_SKIP() << "Skipped on Android: emulator filesystem lacks hardlink " + "support (needed by RocksDB checkpoint)"; +#endif auto check = [&](bool enable_mmap) { SCOPED_TRACE(enable_mmap); options_.enable_mmap_ = enable_mmap; From d37d21f7002eb88fd9e18080a1477a04ca390bde Mon Sep 17 00:00:00 2001 From: Qinren Zhou Date: Thu, 17 Sep 2026 10:56:39 +0800 Subject: [PATCH 09/11] refactor: simplify C binding errors and remove unused constants --- src/binding/c/c_api.cc | 13 +------------ src/db/sqlengine/common/util.h | 3 --- tests/c/c_api_test.c | 13 +++++++++++-- tests/db/index/common/identifier_validation_test.cc | 5 +---- 4 files changed, 13 insertions(+), 21 deletions(-) diff --git a/src/binding/c/c_api.cc b/src/binding/c/c_api.cc index 9e3fab7d4..bf9150875 100644 --- a/src/binding/c/c_api.cc +++ b/src/binding/c/c_api.cc @@ -26,7 +26,6 @@ #include #include #include -#include #include #include #include @@ -101,9 +100,6 @@ SET_LAST_ERROR(ZVEC_ERROR_RESOURCE_EXHAUSTED, \ std::string(msg) + ": " + e.what()); \ return nullptr; \ - } catch (const std::invalid_argument &e) { \ - SET_LAST_ERROR(ZVEC_ERROR_INVALID_ARGUMENT, e.what()); \ - return nullptr; \ } catch (const std::exception &e) { \ SET_LAST_ERROR(ZVEC_ERROR_INTERNAL_ERROR, \ std::string(msg) + ": " + e.what()); \ @@ -122,9 +118,6 @@ SET_LAST_ERROR(ZVEC_ERROR_RESOURCE_EXHAUSTED, \ std::string(msg) + ": " + e.what()); \ return ZVEC_ERROR_RESOURCE_EXHAUSTED; \ - } catch (const std::invalid_argument &e) { \ - SET_LAST_ERROR(ZVEC_ERROR_INVALID_ARGUMENT, e.what()); \ - return ZVEC_ERROR_INVALID_ARGUMENT; \ } catch (const std::exception &e) { \ SET_LAST_ERROR(ZVEC_ERROR_INTERNAL_ERROR, \ std::string(msg) + ": " + e.what()); \ @@ -143,9 +136,6 @@ SET_LAST_ERROR(ZVEC_ERROR_RESOURCE_EXHAUSTED, \ std::string(msg) + ": " + e.what()); \ return (error_val); \ - } catch (const std::invalid_argument &e) { \ - SET_LAST_ERROR(ZVEC_ERROR_INVALID_ARGUMENT, e.what()); \ - return (error_val); \ } catch (const std::exception &e) { \ SET_LAST_ERROR(ZVEC_ERROR_INTERNAL_ERROR, \ std::string(msg) + ": " + e.what()); \ @@ -3195,8 +3185,7 @@ static zvec::Result> convert_zvec_docs_to_internal( for (size_t i = 0; i < doc_count; ++i) { if (!zvec_docs[i]) { return tl::make_unexpected(zvec::Status::InvalidArgument( - "Invalid doc: document must not be null (document at index ", i, - ")")); + "Invalid doc: document must not be null")); } } std::vector docs; diff --git a/src/db/sqlengine/common/util.h b/src/db/sqlengine/common/util.h index 18f3644be..ef6979f87 100644 --- a/src/db/sqlengine/common/util.h +++ b/src/db/sqlengine/common/util.h @@ -19,9 +19,6 @@ namespace zvec::sqlengine { -static const constexpr char *kFieldVector = "_zvec_vector"; -static const constexpr char *kFieldSparseIndices = "_zvec_sindices"; -static const constexpr char *kFieldSparseValues = "_zvec_svalues"; static const constexpr char *kFieldIsValid = "_zvec_is_valid"; static const inline std::string kCheckNotFiltered = "check_not_filtered"; diff --git a/tests/c/c_api_test.c b/tests/c/c_api_test.c index a4402703f..885a171a0 100644 --- a/tests/c/c_api_test.c +++ b/tests/c/c_api_test.c @@ -1267,7 +1267,7 @@ void test_batch_validation_errors(void) { zvec_doc_set_pk(valid_doc, "valid_before_error"); zvec_doc_set_pk(invalid_doc, "\xff"); const zvec_doc_t *invalid_inputs[] = {NULL, invalid_doc}; - const char *reasons[] = {"document must not be null", + const char *reasons[] = {"Invalid doc: document must not be null", "id[\\xFF] is not valid UTF-8"}; for (size_t i = 0; i < 2; ++i) { const zvec_doc_t *docs[] = {valid_doc, invalid_inputs[i]}; @@ -1279,7 +1279,10 @@ void test_batch_validation_errors(void) { TEST_ASSERT(error_count == 2); check_last_error(ZVEC_ERROR_INVALID_ARGUMENT, reasons[i]); if (i == 0) { - check_last_error(ZVEC_ERROR_INVALID_ARGUMENT, "document at index 1"); + zvec_error_details_t details = {0}; + TEST_ASSERT(zvec_get_last_error_details(&details) == ZVEC_OK); + TEST_ASSERT(details.message && + strcmp(details.message, reasons[i]) == 0); } zvec_write_result_t *results = (zvec_write_result_t *)(uintptr_t)1; @@ -1290,6 +1293,12 @@ void test_batch_validation_errors(void) { TEST_ASSERT(results == NULL); TEST_ASSERT(result_count == 0); check_last_error(ZVEC_ERROR_INVALID_ARGUMENT, reasons[i]); + if (i == 0) { + zvec_error_details_t details = {0}; + TEST_ASSERT(zvec_get_last_error_details(&details) == ZVEC_OK); + TEST_ASSERT(details.message && + strcmp(details.message, reasons[i]) == 0); + } } } const char *ids[] = {"valid_before_error"}; diff --git a/tests/db/index/common/identifier_validation_test.cc b/tests/db/index/common/identifier_validation_test.cc index 99dcc2204..79a04d4a4 100644 --- a/tests/db/index/common/identifier_validation_test.cc +++ b/tests/db/index/common/identifier_validation_test.cc @@ -245,10 +245,7 @@ TEST(IdentifierValidationTest, RejectsExactInternalFieldNames) { EXPECT_TRUE(validate_collection_name(name).ok()); } EXPECT_TRUE(validate_field_name("_ZVEC_UID_").ok()); - for (const std::string name : - {"_zvec_vector", "_zvec_sindices", "_zvec_svalues", "_zvec_is_valid"}) { - EXPECT_TRUE(validate_field_name(name).ok()); - } + EXPECT_TRUE(validate_field_name("_zvec_is_valid").ok()); } TEST(IdentifierValidationTest, SharedErrorPreviewIsEscapedAndBounded) { From 98afeaed84f95de9aadffd3a9729726e982e35da Mon Sep 17 00:00:00 2001 From: Qinren Zhou Date: Thu, 17 Sep 2026 16:27:49 +0800 Subject: [PATCH 10/11] fix(test): avoid raw invalid UTF-8 in schema test output --- tests/db/index/common/schema_test.cc | 3 --- 1 file changed, 3 deletions(-) diff --git a/tests/db/index/common/schema_test.cc b/tests/db/index/common/schema_test.cc index fdba5698b..f9877d57a 100644 --- a/tests/db/index/common/schema_test.cc +++ b/tests/db/index/common/schema_test.cc @@ -999,9 +999,6 @@ TEST(CollectionSchemaTest, Validate) { for (const auto &name : invalid_names) { CollectionSchema c(name, {field}); s = c.validate(); - if (!s.ok()) { - std::cout << "Invalid name: " << name << std::endl; - } ASSERT_FALSE(s.ok()); ASSERT_EQ(s.code(), StatusCode::INVALID_ARGUMENT); } From e86a11c5efb4016197045d5af714f8df0f7e765e Mon Sep 17 00:00:00 2001 From: Qinren Zhou Date: Thu, 17 Sep 2026 17:00:20 +0800 Subject: [PATCH 11/11] fix: improve UTF-8 name previews in validation errors --- src/db/common/CMakeLists.txt | 1 + src/db/common/utils.cc | 34 +++++++++++++++++-- src/db/index/common/identifier_validation.cc | 6 ++-- .../common/identifier_validation_test.cc | 29 +++++++++++++--- 4 files changed, 60 insertions(+), 10 deletions(-) diff --git a/src/db/common/CMakeLists.txt b/src/db/common/CMakeLists.txt index cc2110aa3..089fd515f 100644 --- a/src/db/common/CMakeLists.txt +++ b/src/db/common/CMakeLists.txt @@ -9,6 +9,7 @@ cc_library( zvec_ailego roaring rocksdb + utf8proc INCS . VERSION "${PROXIMA_ZVEC_VERSION}" ) diff --git a/src/db/common/utils.cc b/src/db/common/utils.cc index 047f01e94..6524e5707 100644 --- a/src/db/common/utils.cc +++ b/src/db/common/utils.cc @@ -13,6 +13,7 @@ // limitations under the License. #include "utils.h" +#include #include @@ -28,8 +29,34 @@ std::string format_name(std::string_view value) { auto length = std::min(value.size(), kMaxPreviewBytes); std::string preview; preview.reserve(length); - for (size_t i = 0; i < length; ++i) { + size_t i = 0; + while (i < length) { auto byte = static_cast(value[i]); + if (byte >= 0x80) { + utf8proc_int32_t codepoint; + auto bytes = utf8proc_iterate( + reinterpret_cast(value.data() + i), + static_cast(std::min(value.size() - i, size_t{4})), + &codepoint); + if (bytes > 0) { + auto codepoint_bytes = static_cast(bytes); + if (i + codepoint_bytes > length) { + break; + } + switch (utf8proc_category(codepoint)) { + case UTF8PROC_CATEGORY_CC: + case UTF8PROC_CATEGORY_CF: + case UTF8PROC_CATEGORY_CN: + case UTF8PROC_CATEGORY_ZL: + case UTF8PROC_CATEGORY_ZP: + break; + default: + preview.append(value.data() + i, codepoint_bytes); + i += codepoint_bytes; + continue; + } + } + } switch (byte) { case '\0': preview += "\\0"; @@ -59,11 +86,12 @@ std::string format_name(std::string_view value) { } break; } + ++i; } - if (length < value.size()) { + if (i < value.size()) { preview += "..."; } return preview; } -} // namespace zvec \ No newline at end of file +} // namespace zvec diff --git a/src/db/index/common/identifier_validation.cc b/src/db/index/common/identifier_validation.cc index 1720c6278..0d06120be 100644 --- a/src/db/index/common/identifier_validation.cc +++ b/src/db/index/common/identifier_validation.cc @@ -103,9 +103,9 @@ Status validate_field_name(std::string_view name) { "Invalid schema: field name must not be empty"); } if (name.size() > kMaxFieldNameBytes) { - return Status::InvalidArgument("Invalid schema: field name exceeds ", - kMaxFieldNameBytes, " bytes (got ", - name.size(), ")"); + return Status::InvalidArgument("Invalid schema: field[", format_name(name), + "] exceeds ", kMaxFieldNameBytes, + " bytes (got ", name.size(), ")"); } for (unsigned char byte : name) { if ((byte >= 'A' && byte <= 'Z') || (byte >= 'a' && byte <= 'z') || diff --git a/tests/db/index/common/identifier_validation_test.cc b/tests/db/index/common/identifier_validation_test.cc index 79a04d4a4..2446b5e8d 100644 --- a/tests/db/index/common/identifier_validation_test.cc +++ b/tests/db/index/common/identifier_validation_test.cc @@ -201,6 +201,8 @@ TEST(IdentifierValidationTest, Utf8ErrorsIncludeEscapedAndBoundedPreviews) { prefix + "[order\\n\\[123\\]] contains a newline"); ExpectInvalid(validator.validate("order\xff"), prefix + "[order\\xFF] is not valid UTF-8"); + ExpectInvalid(validator.validate(u8"订单\n"), + prefix + u8"[订单\\n] contains a newline"); ExpectInvalid( validator.validate(std::string(40, 'a') + "\n"), prefix + "[" + std::string(32, 'a') + "...] contains a newline"); @@ -226,9 +228,11 @@ TEST(IdentifierValidationTest, RetainsTheFieldAsciiCharacterSet) { .ok()); EXPECT_TRUE(validate_field_name(std::string(kMaxFieldNameBytes, 'a')).ok()); ExpectInvalid(validate_field_name(std::string(kMaxFieldNameBytes + 1, 'a')), - "Invalid schema: field name exceeds 64 bytes (got 65)"); + "Invalid schema: field[" + std::string(32, 'a') + + "...] exceeds 64 bytes (got 65)"); ExpectInvalid(validate_field_name(std::string(10000, 'a')), - "Invalid schema: field name exceeds 64 bytes (got 10000)"); + "Invalid schema: field[" + std::string(32, 'a') + + "...] exceeds 64 bytes (got 10000)"); } TEST(IdentifierValidationTest, RejectsExactInternalFieldNames) { @@ -252,7 +256,24 @@ TEST(IdentifierValidationTest, SharedErrorPreviewIsEscapedAndBounded) { EXPECT_EQ(format_name(""), ""); EXPECT_EQ(format_name(std::string("a\0\n\r\t[]\\", 8)), "a\\0\\n\\r\\t\\[\\]\\\\"); - EXPECT_EQ(format_name(u8"中"), "\\xE4\\xB8\\xAD"); + EXPECT_EQ(format_name(u8"中文€😀e\u0301"), u8"中文€😀e\u0301"); + EXPECT_EQ(format_name(u8"\u0085\u2028\u2029\u202E"), + "\\xC2\\x85\\xE2\\x80\\xA8\\xE2\\x80\\xA9\\xE2\\x80\\xAE"); + EXPECT_EQ(format_name(std::string(u8"中") + "\xff\xe4\xb8"), + u8"中\\xFF\\xE4\\xB8"); + EXPECT_EQ(format_name(std::string("\xff") + u8"中"), u8"\\xFF中"); + for (const std::string character : {u8"é", u8"中", u8"😀"}) { + auto prefix = std::string(32 - character.size(), 'x'); + EXPECT_EQ(format_name(prefix + character), prefix + character); + EXPECT_EQ(format_name(prefix + character + "tail"), + prefix + character + "..."); + EXPECT_EQ(format_name(prefix + "x" + character), prefix + "x..."); + } + EXPECT_EQ(format_name(std::string(30, 'x') + "\xe4\xb8"), + std::string(30, 'x') + "\\xE4\\xB8"); + EXPECT_EQ(format_name(std::string(31, 'x') + "\xc2" + "a"), + std::string(31, 'x') + "\\xC2..."); EXPECT_EQ(format_name(std::string(10000, '\xff')), Repeat("\\xFF", 32) + "..."); EXPECT_EQ(format_name(std::string(10000, 'x')), std::string(32, 'x') + "..."); @@ -268,7 +289,7 @@ TEST(IdentifierValidationTest, validate_field_name("a.b"), "Invalid schema: field[a.b] contains an unsupported character" + rule); ExpectInvalid(validate_field_name(u8"中"), - "Invalid schema: field[\\xE4\\xB8\\xAD] contains a non-ASCII " + "Invalid schema: field[中] contains a non-ASCII " "character" + rule); ExpectInvalid(