mirror of
https://github.com/langgenius/dify.git
synced 2026-09-08 11:04:27 +08:00
fix: sql injection (#38295)
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com> Co-authored-by: Copilot Autofix powered by AI <175728472+Copilot@users.noreply.github.com>
This commit is contained in:
parent
1ae2f95f7c
commit
d9884efaee
@ -132,9 +132,16 @@ class MyScaleVector(BaseVector):
|
|||||||
|
|
||||||
@override
|
@override
|
||||||
def search_by_full_text(self, query: str, **kwargs: Any) -> list[Document]:
|
def search_by_full_text(self, query: str, **kwargs: Any) -> list[Document]:
|
||||||
return self._search(f"TextSearch('enable_nlq=false')(text, '{query}')", SortOrder.DESC, **kwargs)
|
return self._search(
|
||||||
|
"TextSearch('enable_nlq=false')(text, {query:String})",
|
||||||
|
SortOrder.DESC,
|
||||||
|
parameters={"query": query},
|
||||||
|
**kwargs,
|
||||||
|
)
|
||||||
|
|
||||||
def _search(self, dist: str, order: SortOrder, **kwargs: Any) -> list[Document]:
|
def _search(
|
||||||
|
self, dist: str, order: SortOrder, parameters: dict[str, Any] | None = None, **kwargs: Any
|
||||||
|
) -> list[Document]:
|
||||||
top_k = kwargs.get("top_k", 4)
|
top_k = kwargs.get("top_k", 4)
|
||||||
if not isinstance(top_k, int) or top_k <= 0:
|
if not isinstance(top_k, int) or top_k <= 0:
|
||||||
raise ValueError("top_k must be a positive integer")
|
raise ValueError("top_k must be a positive integer")
|
||||||
@ -159,7 +166,7 @@ class MyScaleVector(BaseVector):
|
|||||||
vector=r["vector"],
|
vector=r["vector"],
|
||||||
metadata=r["metadata"],
|
metadata=r["metadata"],
|
||||||
)
|
)
|
||||||
for r in self._client.query(sql).named_results()
|
for r in self._client.query(sql, parameters=parameters).named_results()
|
||||||
]
|
]
|
||||||
except Exception:
|
except Exception:
|
||||||
logger.exception("Vector search operation failed")
|
logger.exception("Vector search operation failed")
|
||||||
|
|||||||
@ -199,6 +199,28 @@ def test_search_delegation_methods(myscale_module):
|
|||||||
assert result_vector == ["result"]
|
assert result_vector == ["result"]
|
||||||
assert result_text == ["result"]
|
assert result_text == ["result"]
|
||||||
assert vector._search.call_count == 2
|
assert vector._search.call_count == 2
|
||||||
|
vector._search.assert_any_call(
|
||||||
|
"TextSearch('enable_nlq=false')(text, {query:String})",
|
||||||
|
myscale_module.SortOrder.DESC,
|
||||||
|
parameters={"query": "hello"},
|
||||||
|
top_k=2,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_search_by_full_text_uses_query_parameters(myscale_module):
|
||||||
|
vector = myscale_module.MyScaleVector("collection_1", _config(myscale_module))
|
||||||
|
vector._client.query.return_value = SimpleNamespace(
|
||||||
|
named_results=lambda: [{"text": "doc", "vector": [0.1], "metadata": {"doc_id": "1"}}]
|
||||||
|
)
|
||||||
|
payload = "x') AS dist FROM dify.collection_1 UNION ALL SELECT secret FROM users --"
|
||||||
|
|
||||||
|
docs = vector.search_by_full_text(payload, top_k=2)
|
||||||
|
|
||||||
|
assert len(docs) == 1
|
||||||
|
sql = vector._client.query.call_args.args[0]
|
||||||
|
assert payload not in sql
|
||||||
|
assert "TextSearch('enable_nlq=false')(text, {query:String})" in sql
|
||||||
|
assert vector._client.query.call_args.kwargs["parameters"] == {"query": payload}
|
||||||
|
|
||||||
|
|
||||||
def test_search_with_document_filter_and_exception(myscale_module):
|
def test_search_with_document_filter_and_exception(myscale_module):
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user