From 88864ae086dd393e300392ecb1c4c697b3007736 Mon Sep 17 00:00:00 2001 From: FFXN <31929997+FFXN@users.noreply.github.com> Date: Thu, 2 Jul 2026 17:02:54 +0800 Subject: [PATCH] fix: enhance SQL query safety and add metadata key validation (#38307) --- .../src/dify_vdb_myscale/myscale_vector.py | 41 +++++++++++++------ .../tests/unit_tests/test_myscale_vector.py | 36 ++++++++++++++-- 2 files changed, 61 insertions(+), 16 deletions(-) diff --git a/api/providers/vdb/vdb-myscale/src/dify_vdb_myscale/myscale_vector.py b/api/providers/vdb/vdb-myscale/src/dify_vdb_myscale/myscale_vector.py index 941d14693a0..c43459ea44f 100644 --- a/api/providers/vdb/vdb-myscale/src/dify_vdb_myscale/myscale_vector.py +++ b/api/providers/vdb/vdb-myscale/src/dify_vdb_myscale/myscale_vector.py @@ -1,5 +1,6 @@ import json import logging +import re import uuid from enum import StrEnum from typing import Any, override @@ -17,6 +18,8 @@ from models.dataset import Dataset logger = logging.getLogger(__name__) +METADATA_KEY_PATTERN = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$") + class MyScaleConfig(BaseModel): host: str @@ -102,7 +105,10 @@ class MyScaleVector(BaseVector): @override def text_exists(self, id: str) -> bool: - results = self._client.query(f"SELECT id FROM {self._config.database}.{self._collection_name} WHERE id='{id}'") + results = self._client.query( + f"SELECT id FROM {self._config.database}.{self._collection_name} WHERE id={{id:String}}", + parameters={"id": id}, + ) return results.row_count > 0 @override @@ -110,20 +116,26 @@ class MyScaleVector(BaseVector): if not ids: return self._client.command( - f"DELETE FROM {self._config.database}.{self._collection_name} WHERE id IN {str(tuple(ids))}" + f"DELETE FROM {self._config.database}.{self._collection_name} WHERE id IN {{ids:Array(String)}}", + parameters={"ids": ids}, ) @override def get_ids_by_metadata_field(self, key: str, value: str): + self._validate_metadata_key(key) rows = self._client.query( - f"SELECT DISTINCT id FROM {self._config.database}.{self._collection_name} WHERE metadata.{key}='{value}'" + f"SELECT DISTINCT id FROM {self._config.database}.{self._collection_name} " + f"WHERE metadata.{key}={{value:String}}", + parameters={"value": value}, ).result_rows return [row[0] for row in rows] @override def delete_by_metadata_field(self, key: str, value: str): + self._validate_metadata_key(key) self._client.command( - f"DELETE FROM {self._config.database}.{self._collection_name} WHERE metadata.{key}='{value}'" + f"DELETE FROM {self._config.database}.{self._collection_name} WHERE metadata.{key}={{value:String}}", + parameters={"value": value}, ) @override @@ -146,15 +158,15 @@ class MyScaleVector(BaseVector): if not isinstance(top_k, int) or top_k <= 0: raise ValueError("top_k must be a positive integer") score_threshold = float(kwargs.get("score_threshold") or 0.0) - where_str = ( - f"WHERE dist < {1 - score_threshold}" - if self._metric.upper() == "COSINE" and order == SortOrder.ASC and score_threshold > 0.0 - else "" - ) + where_conditions = [] + query_parameters = dict(parameters or {}) + if self._metric.upper() == "COSINE" and order == SortOrder.ASC and score_threshold > 0.0: + where_conditions.append(f"dist < {1 - score_threshold}") document_ids_filter = kwargs.get("document_ids_filter") if document_ids_filter: - document_ids = ", ".join(f"'{id}'" for id in document_ids_filter) - where_str = f"{where_str} AND metadata['document_id'] in ({document_ids})" + where_conditions.append("metadata['document_id'] IN {document_ids_filter:Array(String)}") + query_parameters["document_ids_filter"] = document_ids_filter + where_str = f"WHERE {' AND '.join(where_conditions)}" if where_conditions else "" sql = f""" SELECT text, vector, metadata, {dist} as dist FROM {self._config.database}.{self._collection_name} {where_str} ORDER BY dist {order.value} LIMIT {top_k} @@ -166,12 +178,17 @@ class MyScaleVector(BaseVector): vector=r["vector"], metadata=r["metadata"], ) - for r in self._client.query(sql, parameters=parameters).named_results() + for r in self._client.query(sql, parameters=query_parameters).named_results() ] except Exception: logger.exception("Vector search operation failed") return [] + @staticmethod + def _validate_metadata_key(key: str) -> None: + if not METADATA_KEY_PATTERN.match(key): + raise ValueError("metadata key must be a valid identifier") + @override def delete(self): self._client.command(f"DROP TABLE IF EXISTS {self._config.database}.{self._collection_name}") diff --git a/api/providers/vdb/vdb-myscale/tests/unit_tests/test_myscale_vector.py b/api/providers/vdb/vdb-myscale/tests/unit_tests/test_myscale_vector.py index b1ae1f84a31..e2f4eee0769 100644 --- a/api/providers/vdb/vdb-myscale/tests/unit_tests/test_myscale_vector.py +++ b/api/providers/vdb/vdb-myscale/tests/unit_tests/test_myscale_vector.py @@ -181,14 +181,41 @@ def test_text_exists_and_metadata_operations(myscale_module): vector = myscale_module.MyScaleVector("collection_1", _config(myscale_module)) vector._client.query.return_value = SimpleNamespace(row_count=1, result_rows=[("id-1",), ("id-2",)]) - assert vector.text_exists("id-1") is True - assert vector.get_ids_by_metadata_field("document_id", "doc-1") == ["id-1", "id-2"] + assert vector.text_exists("id-1' OR '1'='1") is True + text_exists_call = vector._client.query.call_args + assert "id={id:String}" in text_exists_call.args[0] + assert "id-1' OR '1'='1" not in text_exists_call.args[0] + assert text_exists_call.kwargs["parameters"] == {"id": "id-1' OR '1'='1"} + + assert vector.get_ids_by_metadata_field("document_id", "doc-1' OR '1'='1") == ["id-1", "id-2"] + metadata_query_call = vector._client.query.call_args + assert "metadata.document_id={value:String}" in metadata_query_call.args[0] + assert "doc-1' OR '1'='1" not in metadata_query_call.args[0] + assert metadata_query_call.kwargs["parameters"] == {"value": "doc-1' OR '1'='1"} vector.delete_by_ids(["id-1", "id-2"]) - vector.delete_by_metadata_field("document_id", "doc-1") + delete_ids_call = vector._client.command.call_args + assert "id IN {ids:Array(String)}" in delete_ids_call.args[0] + assert delete_ids_call.kwargs["parameters"] == {"ids": ["id-1", "id-2"]} + + vector.delete_by_metadata_field("document_id", "doc-1' OR '1'='1") + delete_metadata_call = vector._client.command.call_args + assert "metadata.document_id={value:String}" in delete_metadata_call.args[0] + assert "doc-1' OR '1'='1" not in delete_metadata_call.args[0] + assert delete_metadata_call.kwargs["parameters"] == {"value": "doc-1' OR '1'='1"} assert vector._client.command.call_count >= 2 +def test_metadata_operations_reject_invalid_key(myscale_module): + vector = myscale_module.MyScaleVector("collection_1", _config(myscale_module)) + + with pytest.raises(ValueError, match="metadata key must be a valid identifier"): + vector.get_ids_by_metadata_field("document_id) OR 1=1 --", "doc-1") + + with pytest.raises(ValueError, match="metadata key must be a valid identifier"): + vector.delete_by_metadata_field("document_id) OR 1=1 --", "doc-1") + + def test_search_delegation_methods(myscale_module): vector = myscale_module.MyScaleVector("collection_1", _config(myscale_module)) vector._search = MagicMock(return_value=["result"]) @@ -237,7 +264,8 @@ def test_search_with_document_filter_and_exception(myscale_module): ) assert len(docs) == 1 sql = vector._client.query.call_args.args[0] - assert "metadata['document_id'] in ('doc-1', 'doc-2')" in sql + assert "WHERE metadata['document_id'] IN {document_ids_filter:Array(String)}" in sql + assert vector._client.query.call_args.kwargs["parameters"] == {"document_ids_filter": ["doc-1", "doc-2"]} vector._client.query.side_effect = RuntimeError("boom") assert vector._search("distance(vector, [0.1])", myscale_module.SortOrder.ASC, top_k=1) == []