mirror of
https://github.com/langgenius/dify.git
synced 2026-09-08 02:43:49 +08:00
test: use sqlite3 session in test_tool_label_manager (#38762)
This commit is contained in:
parent
415c0db22e
commit
b820ccf086
@ -1,15 +1,28 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from types import SimpleNamespace
|
|
||||||
from typing import Any, override
|
from typing import Any, override
|
||||||
from unittest.mock import MagicMock, PropertyMock, patch
|
from unittest.mock import PropertyMock, patch
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
from sqlalchemy import select
|
||||||
|
from sqlalchemy.engine import Engine
|
||||||
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
|
import core.tools.tool_label_manager as tool_label_manager_module
|
||||||
from core.tools.builtin_tool.provider import BuiltinToolProviderController
|
from core.tools.builtin_tool.provider import BuiltinToolProviderController
|
||||||
from core.tools.custom_tool.provider import ApiToolProviderController
|
from core.tools.custom_tool.provider import ApiToolProviderController
|
||||||
from core.tools.tool_label_manager import ToolLabelManager
|
from core.tools.tool_label_manager import ToolLabelManager
|
||||||
from core.tools.workflow_as_tool.provider import WorkflowToolProviderController
|
from core.tools.workflow_as_tool.provider import WorkflowToolProviderController
|
||||||
|
from models.tools import ToolLabelBinding
|
||||||
|
|
||||||
|
|
||||||
|
class _DatabaseBinding:
|
||||||
|
"""Expose the SQLite engine to code that owns its session lifecycle."""
|
||||||
|
|
||||||
|
engine: Engine
|
||||||
|
|
||||||
|
def __init__(self, engine: Engine) -> None:
|
||||||
|
self.engine = engine
|
||||||
|
|
||||||
|
|
||||||
# Create a mock class for testing abstract/base classes
|
# Create a mock class for testing abstract/base classes
|
||||||
@ -39,7 +52,8 @@ def test_tool_label_manager_filter_tool_labels():
|
|||||||
assert len(filtered) == 2
|
assert len(filtered) == 2
|
||||||
|
|
||||||
|
|
||||||
def test_tool_label_manager_update_tool_labels_db():
|
@pytest.mark.parametrize("sqlite_session", [(ToolLabelBinding,)], indirect=True)
|
||||||
|
def test_tool_label_manager_update_tool_labels_db(sqlite_session: Session):
|
||||||
"""
|
"""
|
||||||
Test the database update logic for tool labels.
|
Test the database update logic for tool labels.
|
||||||
Focus: Verify that labels are filtered, de-duplicated, and safely handled within a database session.
|
Focus: Verify that labels are filtered, de-duplicated, and safely handled within a database session.
|
||||||
@ -49,48 +63,18 @@ def test_tool_label_manager_update_tool_labels_db():
|
|||||||
expected_id = controller.provider_id
|
expected_id = controller.provider_id
|
||||||
expected_type = controller.provider_type
|
expected_type = controller.provider_type
|
||||||
|
|
||||||
# 2. Patching External Dependencies
|
sqlite_session.add(ToolLabelBinding(tool_id=expected_id, tool_type=expected_type, label_name="news"))
|
||||||
# - We patch 'db' to prevent Flask from trying to access a real database.
|
sqlite_session.commit()
|
||||||
# - We patch 'sessionmaker' to intercept and control the creation of SQLAlchemy sessions.
|
|
||||||
with (
|
|
||||||
patch("core.tools.tool_label_manager.db"),
|
|
||||||
patch("core.tools.tool_label_manager.sessionmaker") as mock_sessionmaker,
|
|
||||||
):
|
|
||||||
# 3. Constructing the "Mocking Chain"
|
|
||||||
# In the business logic, we use: with sessionmaker(db.engine).begin() as _session:
|
|
||||||
# We need to link our 'mock_session' to the end of this complex context manager chain:
|
|
||||||
# Step A: sessionmaker(db.engine) -> returns an object (mock_sessionmaker.return_value)
|
|
||||||
# Step B: .begin() -> returns a context manager (begin.return_value)
|
|
||||||
# Step C: with ... as _session: -> calls __enter__(), and _session gets the __enter__.return_value
|
|
||||||
mock_session = MagicMock()
|
|
||||||
mock_sessionmaker.return_value.begin.return_value.__enter__.return_value = mock_session
|
|
||||||
|
|
||||||
# 4. Trigger the logic under test
|
# Duplicate and unknown labels are filtered before the existing binding is replaced.
|
||||||
# Input: ["search", "search", "invalid"]
|
ToolLabelManager.update_tool_labels(controller, ["search", "search", "invalid"], session=sqlite_session)
|
||||||
# Logic:
|
sqlite_session.commit()
|
||||||
# - "invalid" should be filtered out (not in default_tool_label_name_list).
|
|
||||||
# - The duplicate "search" should be merged (unique labels).
|
|
||||||
ToolLabelManager.update_tool_labels(controller, ["search", "search", "invalid"])
|
|
||||||
|
|
||||||
# 5. Behavior Assertion: DELETE operation
|
bindings = list(sqlite_session.scalars(select(ToolLabelBinding)).all())
|
||||||
# Verify that the manager first attempts to clear existing labels for this specific tool.
|
assert len(bindings) == 1
|
||||||
# This ensures the update is idempotent.
|
assert bindings[0].label_name == "search"
|
||||||
mock_session.execute.assert_called_once()
|
assert bindings[0].tool_id == expected_id
|
||||||
|
assert bindings[0].tool_type == expected_type
|
||||||
# 6. Behavior Assertion: INSERT operation
|
|
||||||
# Verify that only ONE valid label ("search") was added after filtering and deduplication.
|
|
||||||
# If call_count == 1, it proves filter_tool_labels() worked as expected.
|
|
||||||
assert mock_session.add.call_count == 1
|
|
||||||
|
|
||||||
# 7. State Assertion: Data Integrity & Isolation
|
|
||||||
# Inspect the actual object passed to session.add() to ensure it has correct properties.
|
|
||||||
# This confirms that the data isolation (tool_id + tool_type) we refactored is active.
|
|
||||||
call_args = mock_session.add.call_args
|
|
||||||
added_label = call_args[0][0] # Retrieve the ToolLabelBinding instance
|
|
||||||
|
|
||||||
assert added_label.label_name == "search", "The label name should be 'search' after filtering."
|
|
||||||
assert added_label.tool_id == expected_id, "The tool_id must match the provider_id for correct binding."
|
|
||||||
assert added_label.tool_type == expected_type, "Isolation failed: tool_type must be verified during update."
|
|
||||||
|
|
||||||
|
|
||||||
# Test error handling
|
# Test error handling
|
||||||
@ -100,7 +84,10 @@ def test_tool_label_manager_update_tool_labels_unsupported():
|
|||||||
|
|
||||||
|
|
||||||
# Test retrieval logic
|
# Test retrieval logic
|
||||||
def test_tool_label_manager_get_tool_labels_for_builtin_and_db():
|
@pytest.mark.parametrize("sqlite_session", [(ToolLabelBinding,)], indirect=True)
|
||||||
|
def test_tool_label_manager_get_tool_labels_for_builtin_and_db(
|
||||||
|
monkeypatch: pytest.MonkeyPatch, sqlite_engine: Engine, sqlite_session: Session
|
||||||
|
):
|
||||||
# Mocking a property (@property) using PropertyMock
|
# Mocking a property (@property) using PropertyMock
|
||||||
with patch.object(
|
with patch.object(
|
||||||
_ConcreteBuiltinToolProviderController,
|
_ConcreteBuiltinToolProviderController,
|
||||||
@ -112,18 +99,17 @@ def test_tool_label_manager_get_tool_labels_for_builtin_and_db():
|
|||||||
assert ToolLabelManager.get_tool_labels(builtin) == ["search", "news"]
|
assert ToolLabelManager.get_tool_labels(builtin) == ["search", "news"]
|
||||||
|
|
||||||
api = _api_controller("api-1")
|
api = _api_controller("api-1")
|
||||||
with (
|
sqlite_session.add_all(
|
||||||
patch("core.tools.tool_label_manager.db"),
|
[
|
||||||
patch("core.tools.tool_label_manager.sessionmaker") as mock_sessionmaker,
|
ToolLabelBinding(tool_id=api.provider_id, tool_type=api.provider_type, label_name="search"),
|
||||||
):
|
ToolLabelBinding(tool_id=api.provider_id, tool_type=api.provider_type, label_name="news"),
|
||||||
mock_session = MagicMock()
|
]
|
||||||
mock_sessionmaker.return_value.begin.return_value.__enter__.return_value = mock_session
|
)
|
||||||
|
sqlite_session.commit()
|
||||||
|
monkeypatch.setattr(tool_label_manager_module, "db", _DatabaseBinding(sqlite_engine))
|
||||||
|
|
||||||
# Inject mock data into the query result: session.scalars(stmt).all()
|
labels = ToolLabelManager.get_tool_labels(api)
|
||||||
mock_session.scalars.return_value.all.return_value = ["search", "news"]
|
assert set(labels) == {"search", "news"}
|
||||||
|
|
||||||
labels = ToolLabelManager.get_tool_labels(api)
|
|
||||||
assert labels == ["search", "news"]
|
|
||||||
|
|
||||||
|
|
||||||
def test_tool_label_manager_get_tool_labels_unsupported():
|
def test_tool_label_manager_get_tool_labels_unsupported():
|
||||||
@ -137,33 +123,30 @@ def test_tool_label_manager_get_tool_labels_unsupported():
|
|||||||
|
|
||||||
|
|
||||||
# Test batch processing and mapping
|
# Test batch processing and mapping
|
||||||
def test_tool_label_manager_get_tools_labels_batch():
|
@pytest.mark.parametrize("sqlite_session", [(ToolLabelBinding,)], indirect=True)
|
||||||
|
def test_tool_label_manager_get_tools_labels_batch(
|
||||||
|
monkeypatch: pytest.MonkeyPatch, sqlite_engine: Engine, sqlite_session: Session
|
||||||
|
):
|
||||||
assert ToolLabelManager.get_tools_labels([]) == {}
|
assert ToolLabelManager.get_tools_labels([]) == {}
|
||||||
|
|
||||||
api = _api_controller("api-1")
|
api = _api_controller("api-1")
|
||||||
wf = _workflow_controller("wf-1")
|
wf = _workflow_controller("wf-1")
|
||||||
|
|
||||||
# SimpleNamespace is a quick way to simulate SQLAlchemy row objects
|
sqlite_session.add_all(
|
||||||
records = [
|
[
|
||||||
SimpleNamespace(tool_id="api-1", label_name="search"),
|
ToolLabelBinding(tool_id=api.provider_id, tool_type=api.provider_type, label_name="search"),
|
||||||
SimpleNamespace(tool_id="api-1", label_name="news"),
|
ToolLabelBinding(tool_id=api.provider_id, tool_type=api.provider_type, label_name="news"),
|
||||||
SimpleNamespace(tool_id="wf-1", label_name="utilities"),
|
ToolLabelBinding(tool_id=wf.provider_id, tool_type=wf.provider_type, label_name="utilities"),
|
||||||
]
|
]
|
||||||
|
)
|
||||||
|
sqlite_session.commit()
|
||||||
|
monkeypatch.setattr(tool_label_manager_module, "db", _DatabaseBinding(sqlite_engine))
|
||||||
|
|
||||||
with (
|
labels = ToolLabelManager.get_tools_labels([api, wf])
|
||||||
patch("core.tools.tool_label_manager.db"),
|
|
||||||
patch("core.tools.tool_label_manager.sessionmaker") as mock_sessionmaker,
|
|
||||||
):
|
|
||||||
mock_session = MagicMock()
|
|
||||||
mock_sessionmaker.return_value.begin.return_value.__enter__.return_value = mock_session
|
|
||||||
|
|
||||||
# Simulating the batch query result
|
assert labels.keys() == {"api-1", "wf-1"}
|
||||||
mock_session.scalars.return_value.all.return_value = records
|
assert set(labels["api-1"]) == {"search", "news"}
|
||||||
|
assert labels["wf-1"] == ["utilities"]
|
||||||
labels = ToolLabelManager.get_tools_labels([api, wf])
|
|
||||||
|
|
||||||
# Verify the final dictionary mapping
|
|
||||||
assert labels == {"api-1": ["search", "news"], "wf-1": ["utilities"]}
|
|
||||||
|
|
||||||
|
|
||||||
def test_tool_label_manager_get_tools_labels_unsupported():
|
def test_tool_label_manager_get_tools_labels_unsupported():
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user