test: use sqlite3 session in test_tool_label_manager (#38762)

This commit is contained in:
Asuka Minato 2026-07-16 21:56:08 +09:00 committed by GitHub
parent 415c0db22e
commit b820ccf086
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194

View File

@ -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():