Skip to content
Open
Show file tree
Hide file tree
Changes from 2 commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 3 additions & 0 deletions pymilvus/orm/schema.py
Original file line number Diff line number Diff line change
Expand Up @@ -657,6 +657,9 @@ def _parse_type_params(self):
def construct_from_dict(cls, raw: Dict):
kwargs = {}
kwargs.update(raw.get("params", {}))
for key in COMMON_TYPE_PARAMS:
Comment thread
yhmo marked this conversation as resolved.
Outdated
if key not in kwargs and raw.get(key) is not None:
kwargs[key] = raw[key]
kwargs["is_primary"] = raw.get("is_primary", False)
if raw.get("auto_id") is not None:
kwargs["auto_id"] = raw.get("auto_id")
Expand Down
116 changes: 116 additions & 0 deletions tests/unit/orm/test_schema.py
Original file line number Diff line number Diff line change
Expand Up @@ -337,6 +337,122 @@ def test_construct_from_dict_roundtrip(self, raw_dict):
assert result["name"] == raw_dict["name"]
assert result["type"] == raw_dict["type"]

def test_construct_from_dict_with_top_level_max_length(self):
raw_dict = {
"name": "text",
"type": DataType.VARCHAR,
"max_length": 256,
}

field = FieldSchema.construct_from_dict(raw_dict)

assert field.params == {"max_length": 256}
assert field.to_dict()["params"] == {"max_length": 256}

def test_construct_from_dict_with_top_level_dim(self):
raw_dict = {
"name": "vec",
"type": DataType.FLOAT_VECTOR,
"dim": 128,
}

field = FieldSchema.construct_from_dict(raw_dict)

assert field.params == {"dim": 128}
assert field.to_dict()["params"] == {"dim": 128}

@pytest.mark.parametrize(
"raw_dict,expected_params",
[
pytest.param(
{
"name": "tags",
"type": DataType.ARRAY,
"element_type": DataType.VARCHAR,
"max_capacity": 100,
},
{"max_capacity": 100},
id="max_capacity",
),
pytest.param(
{
"name": "text",
"type": DataType.TEXT,
"enable_match": True,
"enable_analyzer": False,
},
{"enable_match": True, "enable_analyzer": False},
id="analyzer_flags",
),
pytest.param(
{
"name": "text",
"type": DataType.TEXT,
"analyzer_params": {"type": "standard"},
},
{"analyzer_params": '{"type":"standard"}'},
id="analyzer_params",
),
pytest.param(
{
"name": "text",
"type": DataType.TEXT,
"multi_analyzer_params": {"analyzers": [{"type": "standard"}]},
},
{"multi_analyzer_params": '{"analyzers":[{"type":"standard"}]}'},
id="multi_analyzer_params",
),
],
)
def test_construct_from_dict_with_top_level_common_type_params(self, raw_dict, expected_params):
field = FieldSchema.construct_from_dict(raw_dict)

assert field.params == expected_params
assert field.to_dict()["params"] == expected_params

def test_construct_from_dict_ignores_none_top_level_param(self):
raw_dict = {
"name": "vec",
"type": DataType.FLOAT_VECTOR,
"dim": None,
}

field = FieldSchema.construct_from_dict(raw_dict)

assert field.params == {}
assert "params" not in field.to_dict()

@pytest.mark.parametrize(
"raw_dict,expected_params",
[
pytest.param(
{
"name": "text",
"type": DataType.VARCHAR,
"params": {"max_length": 256},
"max_length": 512,
},
{"max_length": 256},
id="max_length",
),
pytest.param(
{
"name": "vec",
"type": DataType.FLOAT_VECTOR,
"params": {"dim": 128},
"dim": 256,
},
{"dim": 128},
id="dim",
),
],
)
def test_construct_from_dict_nested_params_take_precedence(self, raw_dict, expected_params):
field = FieldSchema.construct_from_dict(raw_dict)

assert field.params == expected_params
assert field.to_dict()["params"] == expected_params


class TestFieldSchemaDeepCopy:
"""Tests for FieldSchema __deepcopy__ method."""
Expand Down