chore: add Type to test (#35942)

Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
This commit is contained in:
Asuka Minato
2026-05-09 03:16:22 +00:00
committed by GitHub
co-authored by autofix-ci[bot]
parent e03eb3a76c
commit 140ad6ba4e
200 changed files with 1497 additions and 1264 deletions
@@ -13,7 +13,7 @@ from controllers.console.app import wraps
from libs.datetime_utils import naive_utc_now
from models import App, Tenant
from models.account import Account, TenantAccountJoin, TenantAccountRole
from models.enums import ConversationFromSource
from models.enums import AppStatus, ConversationFromSource
from models.model import AppMode
from services.app_generate_service import AppGenerateService
@@ -28,7 +28,7 @@ class TestChatMessageApiPermissions:
app.id = str(uuid.uuid4())
app.mode = AppMode.CHAT
app.tenant_id = str(uuid.uuid4())
app.status = "normal"
app.status = AppStatus.NORMAL
return app
@pytest.fixture
@@ -78,7 +78,7 @@ class TestChatMessageApiPermissions:
self,
test_client: FlaskClient,
auth_header,
monkeypatch,
monkeypatch: pytest.MonkeyPatch,
mock_app_model,
mock_account,
role: TenantAccountRole,
@@ -130,7 +130,7 @@ class TestChatMessageApiPermissions:
self,
test_client: FlaskClient,
auth_header,
monkeypatch,
monkeypatch: pytest.MonkeyPatch,
mock_app_model,
mock_account,
role: TenantAccountRole,
@@ -14,7 +14,7 @@ from controllers.console.app import wraps
from libs.datetime_utils import naive_utc_now
from models import App, Tenant
from models.account import Account, TenantAccountJoin, TenantAccountRole
from models.enums import FeedbackFromSource, FeedbackRating
from models.enums import AppStatus, FeedbackFromSource, FeedbackRating
from models.model import AppMode, MessageFeedback
from services.feedback_service import FeedbackService
@@ -29,7 +29,7 @@ class TestFeedbackExportApi:
app.id = str(uuid.uuid4())
app.mode = AppMode.CHAT
app.tenant_id = str(uuid.uuid4())
app.status = "normal"
app.status = AppStatus.NORMAL
app.name = "Test App"
return app
@@ -135,7 +135,7 @@ class TestFeedbackExportApi:
self,
test_client: FlaskClient,
auth_header,
monkeypatch,
monkeypatch: pytest.MonkeyPatch,
mock_app_model,
mock_account,
role: TenantAccountRole,
@@ -167,7 +167,13 @@ class TestFeedbackExportApi:
mock_export_feedbacks.assert_called_once()
def test_feedback_export_csv_format(
self, test_client: FlaskClient, auth_header, monkeypatch, mock_app_model, mock_account, sample_feedback_data
self,
test_client: FlaskClient,
auth_header,
monkeypatch: pytest.MonkeyPatch,
mock_app_model,
mock_account,
sample_feedback_data,
):
"""Test feedback export in CSV format."""
@@ -202,7 +208,13 @@ class TestFeedbackExportApi:
assert "text/csv" in response.content_type
def test_feedback_export_json_format(
self, test_client: FlaskClient, auth_header, monkeypatch, mock_app_model, mock_account, sample_feedback_data
self,
test_client: FlaskClient,
auth_header,
monkeypatch: pytest.MonkeyPatch,
mock_app_model,
mock_account,
sample_feedback_data,
):
"""Test feedback export in JSON format."""
@@ -246,7 +258,7 @@ class TestFeedbackExportApi:
assert "application/json" in response.content_type
def test_feedback_export_with_filters(
self, test_client: FlaskClient, auth_header, monkeypatch, mock_app_model, mock_account
self, test_client: FlaskClient, auth_header, monkeypatch: pytest.MonkeyPatch, mock_app_model, mock_account
):
"""Test feedback export with various filters."""
@@ -287,7 +299,7 @@ class TestFeedbackExportApi:
)
def test_feedback_export_invalid_date_format(
self, test_client: FlaskClient, auth_header, monkeypatch, mock_app_model, mock_account
self, test_client: FlaskClient, auth_header, monkeypatch: pytest.MonkeyPatch, mock_app_model, mock_account
):
"""Test feedback export with invalid date format."""
@@ -312,7 +324,7 @@ class TestFeedbackExportApi:
assert "Parameter validation error" in response_json["error"]
def test_feedback_export_server_error(
self, test_client: FlaskClient, auth_header, monkeypatch, mock_app_model, mock_account
self, test_client: FlaskClient, auth_header, monkeypatch: pytest.MonkeyPatch, mock_app_model, mock_account
):
"""Test feedback export with server error."""
@@ -11,6 +11,7 @@ from controllers.console.app import wraps
from libs.datetime_utils import naive_utc_now
from models import App, Tenant
from models.account import Account, TenantAccountJoin, TenantAccountRole
from models.enums import AppStatus
from models.model import AppMode
from services.app_model_config_service import AppModelConfigService
@@ -25,7 +26,7 @@ class TestModelConfigResourcePermissions:
app.id = str(uuid.uuid4())
app.mode = AppMode.CHAT
app.tenant_id = str(uuid.uuid4())
app.status = "normal"
app.status = AppStatus.NORMAL
app.app_model_config_id = str(uuid.uuid4())
return app
@@ -73,7 +74,7 @@ class TestModelConfigResourcePermissions:
self,
test_client: FlaskClient,
auth_header,
monkeypatch,
monkeypatch: pytest.MonkeyPatch,
mock_app_model,
mock_account,
role: TenantAccountRole,
@@ -1,5 +1,7 @@
from collections.abc import Generator
from pytest_mock import MockerFixture
from core.datasource.datasource_manager import DatasourceManager
from core.datasource.entities.datasource_entities import DatasourceMessage
from graphon.node_events import StreamCompletedEvent
@@ -19,7 +21,7 @@ def _gen_var_stream() -> Generator[DatasourceMessage, None, None]:
)
def test_stream_node_events_accumulates_variables(mocker):
def test_stream_node_events_accumulates_variables(mocker: MockerFixture):
mocker.patch.object(DatasourceManager, "stream_online_results", return_value=_gen_var_stream())
events = list(
DatasourceManager.stream_node_events(
@@ -1,3 +1,5 @@
from pytest_mock import MockerFixture
from core.app.entities.app_invoke_entities import DIFY_RUN_CONTEXT_KEY
from core.workflow.nodes.datasource.datasource_node import DatasourceNode
from core.workflow.nodes.datasource.entities import DatasourceNodeData
@@ -44,7 +46,7 @@ class _GP:
call_depth = 0
def test_node_integration_minimal_stream(mocker):
def test_node_integration_minimal_stream(mocker: MockerFixture):
sys_d = {
"sys": {
"datasource_type": "online_document",
@@ -2,6 +2,8 @@ import time
import uuid
from unittest.mock import MagicMock, patch
import pytest
from core.app.entities.app_invoke_entities import InvokeFrom, UserFrom
from core.tools.utils.configuration import ToolParameterConfigurationManager
from core.workflow.node_factory import DifyNodeFactory
@@ -71,7 +73,7 @@ def init_tool_node(config: dict):
return node
def test_tool_variable_invoke(monkeypatch):
def test_tool_variable_invoke(monkeypatch: pytest.MonkeyPatch):
node = init_tool_node(
config={
"id": "1",
@@ -106,7 +108,7 @@ def test_tool_variable_invoke(monkeypatch):
assert item.node_run_result.outputs.get("text") is not None
def test_tool_mixed_invoke(monkeypatch):
def test_tool_mixed_invoke(monkeypatch: pytest.MonkeyPatch):
node = init_tool_node(
config={
"id": "1",
@@ -11,7 +11,7 @@ from libs import helper as helper_module
@pytest.mark.usefixtures("flask_app_with_containers")
def test_rate_limiter_counts_multiple_attempts_in_same_second(monkeypatch):
def test_rate_limiter_counts_multiple_attempts_in_same_second(monkeypatch: pytest.MonkeyPatch):
prefix = f"test_rate_limit:{uuid.uuid4().hex}"
limiter = helper_module.RateLimiter(prefix=prefix, max_attempts=2, time_window=60)
key = limiter._get_key("203.0.113.10")
@@ -6,7 +6,7 @@ from faker import Faker
from sqlalchemy.orm import Session
from core.plugin.impl.exc import PluginDaemonClientSideError
from models import Account
from models import Account, CreatorUserRole
from models.enums import ConversationFromSource, MessageFileBelongsTo
from models.model import AppModelConfig, Conversation, EndUser, Message, MessageAgentThought
from services.account_service import AccountService, TenantService
@@ -246,7 +246,7 @@ class TestAgentService:
tool_input=json.dumps({"test_tool": {"input": "test_input"}}),
observation=json.dumps({"test_tool": {"output": "test_output"}}),
tokens=50,
created_by_role="account",
created_by_role=CreatorUserRole.ACCOUNT,
created_by=message.from_account_id,
)
db_session_with_containers.add(thought1)
@@ -294,7 +294,7 @@ class TestAgentService:
agent_thoughts = self._create_test_agent_thoughts(db_session_with_containers, message)
# Execute the method under test
result = AgentService.get_agent_logs(app, str(conversation.id), str(message.id))
result = AgentService.get_agent_logs(app, conversation.id, message.id)
# Verify the result structure
assert result is not None
@@ -370,7 +370,7 @@ class TestAgentService:
# Execute the method under test with non-existent message
with pytest.raises(ValueError, match="Message not found"):
AgentService.get_agent_logs(app, str(conversation.id), fake.uuid4())
AgentService.get_agent_logs(app, conversation.id, fake.uuid4())
def test_get_agent_logs_with_end_user(
self, db_session_with_containers: Session, mock_external_service_dependencies
@@ -451,7 +451,7 @@ class TestAgentService:
db_session_with_containers.commit()
# Execute the method under test
result = AgentService.get_agent_logs(app, str(conversation.id), str(message.id))
result = AgentService.get_agent_logs(app, conversation.id, message.id)
# Verify the result
assert result is not None
@@ -523,7 +523,7 @@ class TestAgentService:
db_session_with_containers.commit()
# Execute the method under test
result = AgentService.get_agent_logs(app, str(conversation.id), str(message.id))
result = AgentService.get_agent_logs(app, conversation.id, message.id)
# Verify the result
assert result is not None
@@ -561,14 +561,14 @@ class TestAgentService:
tool_input=json.dumps({"error_tool": {"input": "test_input"}}),
observation=json.dumps({"error_tool": {"output": "error_output"}}),
tokens=50,
created_by_role="account",
created_by_role=CreatorUserRole.ACCOUNT,
created_by=message.from_account_id,
)
db_session_with_containers.add(thought_with_error)
db_session_with_containers.commit()
# Execute the method under test
result = AgentService.get_agent_logs(app, str(conversation.id), str(message.id))
result = AgentService.get_agent_logs(app, conversation.id, message.id)
# Verify the result
assert result is not None
@@ -592,7 +592,7 @@ class TestAgentService:
conversation, message = self._create_test_conversation_and_message(db_session_with_containers, app, account)
# Execute the method under test
result = AgentService.get_agent_logs(app, str(conversation.id), str(message.id))
result = AgentService.get_agent_logs(app, conversation.id, message.id)
# Verify the result
assert result is not None
@@ -654,7 +654,7 @@ class TestAgentService:
# Execute the method under test
with pytest.raises(ValueError, match="App model config not found"):
AgentService.get_agent_logs(app, str(conversation.id), str(message.id))
AgentService.get_agent_logs(app, conversation.id, message.id)
def test_get_agent_logs_agent_config_not_found(
self, db_session_with_containers: Session, mock_external_service_dependencies
@@ -673,7 +673,7 @@ class TestAgentService:
# Execute the method under test
with pytest.raises(ValueError, match="Agent config not found"):
AgentService.get_agent_logs(app, str(conversation.id), str(message.id))
AgentService.get_agent_logs(app, conversation.id, message.id)
def test_list_agent_providers_success(
self, db_session_with_containers: Session, mock_external_service_dependencies
@@ -687,7 +687,7 @@ class TestAgentService:
app, account = self._create_test_app_and_account(db_session_with_containers, mock_external_service_dependencies)
# Execute the method under test
result = AgentService.list_agent_providers(str(account.id), str(app.tenant_id))
result = AgentService.list_agent_providers(account.id, app.tenant_id)
# Verify the result
assert result is not None
@@ -696,7 +696,7 @@ class TestAgentService:
# Verify the mock was called correctly
mock_plugin_client = mock_external_service_dependencies["plugin_agent_client"].return_value
mock_plugin_client.fetch_agent_strategy_providers.assert_called_once_with(str(app.tenant_id))
mock_plugin_client.fetch_agent_strategy_providers.assert_called_once_with(app.tenant_id)
def test_get_agent_provider_success(self, db_session_with_containers: Session, mock_external_service_dependencies):
"""
@@ -710,7 +710,7 @@ class TestAgentService:
provider_name = "test_provider"
# Execute the method under test
result = AgentService.get_agent_provider(str(account.id), str(app.tenant_id), provider_name)
result = AgentService.get_agent_provider(account.id, app.tenant_id, provider_name)
# Verify the result
assert result is not None
@@ -718,7 +718,7 @@ class TestAgentService:
# Verify the mock was called correctly
mock_plugin_client = mock_external_service_dependencies["plugin_agent_client"].return_value
mock_plugin_client.fetch_agent_strategy_provider.assert_called_once_with(str(app.tenant_id), provider_name)
mock_plugin_client.fetch_agent_strategy_provider.assert_called_once_with(app.tenant_id, provider_name)
def test_get_agent_provider_plugin_error(
self, db_session_with_containers: Session, mock_external_service_dependencies
@@ -740,7 +740,7 @@ class TestAgentService:
# Execute the method under test
with pytest.raises(ValueError, match=error_message):
AgentService.get_agent_provider(str(account.id), str(app.tenant_id), provider_name)
AgentService.get_agent_provider(account.id, app.tenant_id, provider_name)
def test_get_agent_logs_with_complex_tool_data(
self, db_session_with_containers: Session, mock_external_service_dependencies
@@ -796,14 +796,14 @@ class TestAgentService:
{"tool1": {"output1": "result1"}, "tool2": {"output2": "result2"}, "tool3": {"output3": "result3"}}
),
tokens=100,
created_by_role="account",
created_by_role=CreatorUserRole.ACCOUNT,
created_by=message.from_account_id,
)
db_session_with_containers.add(complex_thought)
db_session_with_containers.commit()
# Execute the method under test
result = AgentService.get_agent_logs(app, str(conversation.id), str(message.id))
result = AgentService.get_agent_logs(app, conversation.id, message.id)
# Verify the result
assert result is not None
@@ -891,14 +891,14 @@ class TestAgentService:
observation=json.dumps({"file_tool": {"output": "test_output"}}),
message_files=json.dumps(["file1", "file2"]),
tokens=50,
created_by_role="account",
created_by_role=CreatorUserRole.ACCOUNT,
created_by=message.from_account_id,
)
db_session_with_containers.add(thought_with_files)
db_session_with_containers.commit()
# Execute the method under test
result = AgentService.get_agent_logs(app, str(conversation.id), str(message.id))
result = AgentService.get_agent_logs(app, conversation.id, message.id)
# Verify the result
assert result is not None
@@ -926,7 +926,7 @@ class TestAgentService:
mock_external_service_dependencies["current_user"].timezone = "Asia/Shanghai"
# Execute the method under test
result = AgentService.get_agent_logs(app, str(conversation.id), str(message.id))
result = AgentService.get_agent_logs(app, conversation.id, message.id)
# Verify the result
assert result is not None
@@ -960,14 +960,14 @@ class TestAgentService:
tool_input="", # Empty input
observation="", # Empty observation
tokens=50,
created_by_role="account",
created_by_role=CreatorUserRole.ACCOUNT,
created_by=message.from_account_id,
)
db_session_with_containers.add(empty_thought)
db_session_with_containers.commit()
# Execute the method under test
result = AgentService.get_agent_logs(app, str(conversation.id), str(message.id))
result = AgentService.get_agent_logs(app, conversation.id, message.id)
# Verify the result
assert result is not None
@@ -1001,14 +1001,14 @@ class TestAgentService:
tool_input="invalid json", # Malformed JSON
observation="invalid json", # Malformed JSON
tokens=50,
created_by_role="account",
created_by_role=CreatorUserRole.ACCOUNT,
created_by=message.from_account_id,
)
db_session_with_containers.add(malformed_thought)
db_session_with_containers.commit()
# Execute the method under test
result = AgentService.get_agent_logs(app, str(conversation.id), str(message.id))
result = AgentService.get_agent_logs(app, conversation.id, message.id)
# Verify the result - should handle malformed JSON gracefully
assert result is not None
@@ -198,7 +198,7 @@ class TestAppDslService:
def test_check_version_compatibility_newer_version_returns_pending(self):
assert _check_version_compatibility("99.0.0") == ImportStatus.PENDING
def test_check_version_compatibility_major_older_returns_pending(self, monkeypatch):
def test_check_version_compatibility_major_older_returns_pending(self, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr(app_dsl_service, "CURRENT_DSL_VERSION", "1.0.0")
assert _check_version_compatibility("0.9.9") == ImportStatus.PENDING
@@ -272,7 +272,9 @@ class TestAppDslService:
assert result.status == ImportStatus.FAILED
assert "Missing app data" in result.error
def test_import_app_yaml_error_returns_failed(self, db_session_with_containers: Session, monkeypatch):
def test_import_app_yaml_error_returns_failed(
self, db_session_with_containers: Session, monkeypatch: pytest.MonkeyPatch
):
def bad_safe_load(_content: str):
raise yaml.YAMLError("bad")
@@ -287,7 +289,9 @@ class TestAppDslService:
assert result.status == ImportStatus.FAILED
assert result.error.startswith("Invalid YAML format:")
def test_import_app_unexpected_error_returns_failed(self, db_session_with_containers: Session, monkeypatch):
def test_import_app_unexpected_error_returns_failed(
self, db_session_with_containers: Session, monkeypatch: pytest.MonkeyPatch
):
monkeypatch.setattr(
AppDslService,
"_create_or_update_app",
@@ -305,7 +309,9 @@ class TestAppDslService:
# ── Import: YAML URL ──────────────────────────────────────────────
def test_import_app_yaml_url_fetch_error_returns_failed(self, db_session_with_containers: Session, monkeypatch):
def test_import_app_yaml_url_fetch_error_returns_failed(
self, db_session_with_containers: Session, monkeypatch: pytest.MonkeyPatch
):
monkeypatch.setattr(
app_dsl_service.ssrf_proxy,
"get",
@@ -321,7 +327,9 @@ class TestAppDslService:
assert result.status == ImportStatus.FAILED
assert "Error fetching YAML from URL: boom" in result.error
def test_import_app_yaml_url_empty_content_returns_failed(self, db_session_with_containers: Session, monkeypatch):
def test_import_app_yaml_url_empty_content_returns_failed(
self, db_session_with_containers: Session, monkeypatch: pytest.MonkeyPatch
):
response = MagicMock()
response.content = b""
response.raise_for_status.return_value = None
@@ -336,7 +344,9 @@ class TestAppDslService:
assert result.status == ImportStatus.FAILED
assert "Empty content" in result.error
def test_import_app_yaml_url_file_too_large_returns_failed(self, db_session_with_containers: Session, monkeypatch):
def test_import_app_yaml_url_file_too_large_returns_failed(
self, db_session_with_containers: Session, monkeypatch: pytest.MonkeyPatch
):
response = MagicMock()
response.content = b"x" * (DSL_MAX_SIZE + 1)
response.raise_for_status.return_value = None
@@ -379,7 +389,9 @@ class TestAppDslService:
assert result.imported_dsl_version == "99.0.0"
assert requested_urls == [yaml_url]
def test_import_app_yaml_url_github_blob_rewrites_to_raw(self, db_session_with_containers: Session, monkeypatch):
def test_import_app_yaml_url_github_blob_rewrites_to_raw(
self, db_session_with_containers: Session, monkeypatch: pytest.MonkeyPatch
):
yaml_url = "https://github.com/acme/repo/blob/main/app.yml"
raw_url = "https://raw.githubusercontent.com/acme/repo/main/app.yml"
yaml_bytes = _pending_yaml_content()
@@ -491,7 +503,7 @@ class TestAppDslService:
@pytest.mark.parametrize("has_workflow", [True, False])
def test_import_app_legacy_versions_extract_dependencies(
self, db_session_with_containers: Session, monkeypatch, has_workflow: bool
self, db_session_with_containers: Session, monkeypatch: pytest.MonkeyPatch, has_workflow: bool
):
monkeypatch.setattr(
AppDslService,
@@ -554,7 +566,9 @@ class TestAppDslService:
assert result.status == ImportStatus.FAILED
assert "expired" in result.error
def test_confirm_import_success_deletes_redis_key(self, db_session_with_containers: Session, monkeypatch):
def test_confirm_import_success_deletes_redis_key(
self, db_session_with_containers: Session, monkeypatch: pytest.MonkeyPatch
):
import_id = str(uuid4())
redis_key = f"{IMPORT_INFO_REDIS_KEY_PREFIX}{import_id}"
@@ -614,7 +628,9 @@ class TestAppDslService:
result = service.check_dependencies(app_model=app_model)
assert result.leaked_dependencies == []
def test_check_dependencies_calls_analysis_service(self, db_session_with_containers: Session, monkeypatch):
def test_check_dependencies_calls_analysis_service(
self, db_session_with_containers: Session, monkeypatch: pytest.MonkeyPatch
):
app_id = str(uuid4())
pending = CheckDependenciesPendingData(dependencies=[], app_id=app_id)
redis_client.setex(
@@ -665,7 +681,9 @@ class TestAppDslService:
with pytest.raises(ValueError, match="loss app mode"):
service._create_or_update_app(app=None, data={"app": {}}, account=_account_mock())
def test_create_or_update_app_existing_app_updates_fields(self, db_session_with_containers: Session, monkeypatch):
def test_create_or_update_app_existing_app_updates_fields(
self, db_session_with_containers: Session, monkeypatch: pytest.MonkeyPatch
):
fixed_now = object()
monkeypatch.setattr(app_dsl_service, "naive_utc_now", lambda: fixed_now)
@@ -778,8 +796,8 @@ class TestAppDslService:
service = AppDslService(db_session_with_containers)
with pytest.raises(ValueError, match="Missing model_config"):
service._create_or_update_app(
app=_app_stub(mode=AppMode.CHAT.value),
data={"app": {"mode": AppMode.CHAT.value}},
app=_app_stub(mode=AppMode.CHAT),
data={"app": {"mode": AppMode.CHAT}},
account=_account_mock(),
)
@@ -794,7 +812,7 @@ class TestAppDslService:
service._create_or_update_app(
app=app,
data={
"app": {"mode": AppMode.CHAT.value},
"app": {"mode": AppMode.CHAT},
"model_config": {"model": {"provider": "openai"}},
},
account=account,
@@ -807,14 +825,14 @@ class TestAppDslService:
service = AppDslService(db_session_with_containers)
with pytest.raises(ValueError, match="Invalid app mode"):
service._create_or_update_app(
app=_app_stub(mode=AppMode.RAG_PIPELINE.value),
data={"app": {"mode": AppMode.RAG_PIPELINE.value}},
app=_app_stub(mode=AppMode.RAG_PIPELINE),
data={"app": {"mode": AppMode.RAG_PIPELINE}},
account=_account_mock(),
)
# ── Export ─────────────────────────────────────────────────────────
def test_export_dsl_delegates_by_mode(self, monkeypatch):
def test_export_dsl_delegates_by_mode(self, monkeypatch: pytest.MonkeyPatch):
workflow_calls: list[bool] = []
model_calls: list[bool] = []
monkeypatch.setattr(
@@ -836,14 +854,14 @@ class TestAppDslService:
assert workflow_calls == [True]
chat_app = _app_stub(
mode=AppMode.CHAT.value,
mode=AppMode.CHAT,
icon_type="emoji",
app_model_config=SimpleNamespace(to_dict=lambda: {"agent_mode": {"tools": []}}),
)
AppDslService.export_dsl(chat_app)
assert model_calls == [True]
def test_export_dsl_preserves_icon_and_icon_type(self, monkeypatch):
def test_export_dsl_preserves_icon_and_icon_type(self, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr(
AppDslService,
"_append_workflow_export_data",
@@ -1011,7 +1029,7 @@ class TestAppDslService:
# ── Workflow Export Data ───────────────────────────────────────────
def test_append_workflow_export_data_filters_and_overrides(self, monkeypatch):
def test_append_workflow_export_data_filters_and_overrides(self, monkeypatch: pytest.MonkeyPatch):
workflow_dict = {
"graph": {
"nodes": [
@@ -1111,7 +1129,7 @@ class TestAppDslService:
assert nodes[5]["data"]["subscription_id"] == ""
assert export_data["dependencies"] == [{"tenant": _DEFAULT_TENANT_ID, "dep": "dep-1"}]
def test_append_workflow_export_data_missing_workflow_raises(self, monkeypatch):
def test_append_workflow_export_data_missing_workflow_raises(self, monkeypatch: pytest.MonkeyPatch):
workflow_service = MagicMock()
workflow_service.get_draft_workflow.return_value = None
monkeypatch.setattr(app_dsl_service, "WorkflowService", lambda: workflow_service)
@@ -1126,7 +1144,7 @@ class TestAppDslService:
# ── Model Config Export Data ──────────────────────────────────────
def test_append_model_config_export_data_filters_credential_id(self, monkeypatch):
def test_append_model_config_export_data_filters_credential_id(self, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr(
AppDslService,
"_extract_dependencies_from_model_config",
@@ -1160,7 +1178,7 @@ class TestAppDslService:
# ── Dependency Extraction ─────────────────────────────────────────
def test_extract_dependencies_from_workflow_graph_covers_all_node_types(self, monkeypatch):
def test_extract_dependencies_from_workflow_graph_covers_all_node_types(self, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr(
app_dsl_service.DependenciesAnalysisService,
"analyze_tool_dependency",
@@ -1230,7 +1248,7 @@ class TestAppDslService:
"model:m4",
]
def test_extract_dependencies_from_workflow_graph_handles_exceptions(self, monkeypatch):
def test_extract_dependencies_from_workflow_graph_handles_exceptions(self, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr(
app_dsl_service.ToolNodeData,
"model_validate",
@@ -1241,7 +1259,7 @@ class TestAppDslService:
)
assert deps == []
def test_extract_dependencies_from_model_config_parses_providers(self, monkeypatch):
def test_extract_dependencies_from_model_config_parses_providers(self, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr(
app_dsl_service.DependenciesAnalysisService,
"analyze_model_provider_dependency",
@@ -1264,7 +1282,7 @@ class TestAppDslService:
)
assert deps == ["model:p1", "model:p2", "tool:t1"]
def test_extract_dependencies_from_model_config_handles_exceptions(self, monkeypatch):
def test_extract_dependencies_from_model_config_handles_exceptions(self, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr(
app_dsl_service.DependenciesAnalysisService,
"analyze_model_provider_dependency",
@@ -1278,7 +1296,7 @@ class TestAppDslService:
def test_get_leaked_dependencies_empty_returns_empty(self):
assert AppDslService.get_leaked_dependencies(_DEFAULT_TENANT_ID, []) == []
def test_get_leaked_dependencies_delegates(self, monkeypatch):
def test_get_leaked_dependencies_delegates(self, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr(
app_dsl_service.DependenciesAnalysisService,
"get_leaked_dependencies",
@@ -1289,7 +1307,7 @@ class TestAppDslService:
# ── Encryption/Decryption ─────────────────────────────────────────
def test_encrypt_decrypt_dataset_id_respects_config(self, monkeypatch):
def test_encrypt_decrypt_dataset_id_respects_config(self, monkeypatch: pytest.MonkeyPatch):
tenant_id = _DEFAULT_TENANT_ID
dataset_uuid = "00000000-0000-0000-0000-000000000000"
@@ -1314,7 +1332,7 @@ class TestAppDslService:
value = "00000000-0000-0000-0000-000000000000"
assert AppDslService.decrypt_dataset_id(encrypted_data=value, tenant_id=_DEFAULT_TENANT_ID) == value
def test_decrypt_dataset_id_returns_none_on_invalid_data(self, monkeypatch):
def test_decrypt_dataset_id_returns_none_on_invalid_data(self, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr(
app_dsl_service.dify_config,
"DSL_EXPORT_ENCRYPT_DATASET_ID",
@@ -1322,7 +1340,7 @@ class TestAppDslService:
)
assert AppDslService.decrypt_dataset_id(encrypted_data="not-base64", tenant_id=_DEFAULT_TENANT_ID) is None
def test_decrypt_dataset_id_returns_none_when_decrypted_is_not_uuid(self, monkeypatch):
def test_decrypt_dataset_id_returns_none_when_decrypted_is_not_uuid(self, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr(
app_dsl_service.dify_config,
"DSL_EXPORT_ENCRYPT_DATASET_ID",
@@ -6,6 +6,7 @@ from sqlalchemy.orm import Session
from constants.model_template import default_app_templates
from models import Account
from models.enums import AppStatus, CustomizeTokenStrategy
from models.model import App, IconType, Site
from services.account_service import AccountService, TenantService
from tests.test_containers_integration_tests.helpers import generate_valid_password
@@ -1079,9 +1080,9 @@ class TestAppService:
site.app_id = app.id
site.code = fake.postalcode()
site.title = fake.company()
site.status = "normal"
site.status = AppStatus.NORMAL
site.default_language = "en-US"
site.customize_token_strategy = "uuid"
site.customize_token_strategy = CustomizeTokenStrategy.UUID
db_session_with_containers.add(site)
db_session_with_containers.commit()
@@ -10,6 +10,7 @@ from sqlalchemy import select
from sqlalchemy.orm import Session
from core.app.entities.app_invoke_entities import InvokeFrom
from models import TenantAccountRole
from models.account import Account, Tenant, TenantAccountJoin
from models.enums import ConversationFromSource
from models.model import App, Conversation, EndUser, Message, MessageAnnotation
@@ -22,7 +23,7 @@ from services.message_service import MessageService
class ConversationServiceIntegrationTestDataFactory:
@staticmethod
def create_app_and_account(db_session_with_containers):
def create_app_and_account(db_session_with_containers: Session):
tenant = Tenant(name=f"Tenant {uuid4()}")
db_session_with_containers.add(tenant)
db_session_with_containers.flush()
@@ -41,7 +42,7 @@ class ConversationServiceIntegrationTestDataFactory:
tenant_join = TenantAccountJoin(
tenant_id=tenant.id,
account_id=account.id,
role="owner",
role=TenantAccountRole.OWNER,
current=True,
)
db_session_with_containers.add(tenant_join)
@@ -155,7 +156,7 @@ class ConversationServiceIntegrationTestDataFactory:
total_price=Decimal(0),
currency="USD",
status="normal",
invoke_from=InvokeFrom.WEB_APP.value,
invoke_from=InvokeFrom.WEB_APP,
from_source=ConversationFromSource.API if isinstance(user, EndUser) else ConversationFromSource.CONSOLE,
from_end_user_id=user.id if isinstance(user, EndUser) else None,
from_account_id=user.id if isinstance(user, Account) else None,
@@ -25,7 +25,7 @@ from services.errors.conversation import (
class ConversationServiceVariableIntegrationFactory:
@staticmethod
def create_app_and_account(db_session_with_containers):
def create_app_and_account(db_session_with_containers: Session):
tenant = Tenant(name=f"Tenant {uuid4()}")
db_session_with_containers.add(tenant)
db_session_with_containers.flush()
@@ -6,6 +6,7 @@ from unittest.mock import create_autospec, patch
from uuid import uuid4
import pytest
from sqlalchemy.orm import Session
from werkzeug.exceptions import Forbidden, NotFound
from core.rag.index_processor.constant.index_type import IndexStructureType
@@ -119,13 +120,13 @@ def current_user_mock():
yield current_user
def test_get_document_returns_none_when_document_id_is_missing(db_session_with_containers):
def test_get_document_returns_none_when_document_id_is_missing(db_session_with_containers: Session):
dataset = DocumentServiceIntegrationFactory.create_dataset(db_session_with_containers)
assert DocumentService.get_document(dataset.id, None) is None
def test_get_document_queries_by_dataset_and_document_id(db_session_with_containers):
def test_get_document_queries_by_dataset_and_document_id(db_session_with_containers: Session):
dataset = DocumentServiceIntegrationFactory.create_dataset(db_session_with_containers)
document = DocumentServiceIntegrationFactory.create_document(db_session_with_containers, dataset=dataset)
@@ -135,7 +136,7 @@ def test_get_document_queries_by_dataset_and_document_id(db_session_with_contain
assert result.id == document.id
def test_get_documents_by_ids_returns_empty_for_empty_input(db_session_with_containers):
def test_get_documents_by_ids_returns_empty_for_empty_input(db_session_with_containers: Session):
dataset = DocumentServiceIntegrationFactory.create_dataset(db_session_with_containers)
result = DocumentService.get_documents_by_ids(dataset.id, [])
@@ -143,7 +144,7 @@ def test_get_documents_by_ids_returns_empty_for_empty_input(db_session_with_cont
assert result == []
def test_get_documents_by_ids_uses_single_batch_query(db_session_with_containers):
def test_get_documents_by_ids_uses_single_batch_query(db_session_with_containers: Session):
dataset = DocumentServiceIntegrationFactory.create_dataset(db_session_with_containers)
doc_a = DocumentServiceIntegrationFactory.create_document(db_session_with_containers, dataset=dataset, name="a.txt")
doc_b = DocumentServiceIntegrationFactory.create_document(
@@ -158,13 +159,13 @@ def test_get_documents_by_ids_uses_single_batch_query(db_session_with_containers
assert {document.id for document in result} == {doc_a.id, doc_b.id}
def test_update_documents_need_summary_returns_zero_for_empty_input(db_session_with_containers):
def test_update_documents_need_summary_returns_zero_for_empty_input(db_session_with_containers: Session):
dataset = DocumentServiceIntegrationFactory.create_dataset(db_session_with_containers)
assert DocumentService.update_documents_need_summary(dataset.id, []) == 0
def test_update_documents_need_summary_updates_matching_non_qa_documents(db_session_with_containers):
def test_update_documents_need_summary_updates_matching_non_qa_documents(db_session_with_containers: Session):
dataset = DocumentServiceIntegrationFactory.create_dataset(db_session_with_containers)
paragraph_doc = DocumentServiceIntegrationFactory.create_document(
db_session_with_containers,
@@ -195,7 +196,7 @@ def test_update_documents_need_summary_updates_matching_non_qa_documents(db_sess
assert refreshed_qa.need_summary is True
def test_get_document_download_url_uses_signed_url_helper(db_session_with_containers):
def test_get_document_download_url_uses_signed_url_helper(db_session_with_containers: Session):
dataset = DocumentServiceIntegrationFactory.create_dataset(db_session_with_containers)
upload_file = DocumentServiceIntegrationFactory.create_upload_file(
db_session_with_containers,
@@ -215,7 +216,7 @@ def test_get_document_download_url_uses_signed_url_helper(db_session_with_contai
get_url.assert_called_once_with(upload_file_id=upload_file.id, as_attachment=True)
def test_get_upload_file_id_for_upload_file_document_rejects_invalid_source_type(db_session_with_containers):
def test_get_upload_file_id_for_upload_file_document_rejects_invalid_source_type(db_session_with_containers: Session):
dataset = DocumentServiceIntegrationFactory.create_dataset(db_session_with_containers)
document = DocumentServiceIntegrationFactory.create_document(
db_session_with_containers,
@@ -232,7 +233,9 @@ def test_get_upload_file_id_for_upload_file_document_rejects_invalid_source_type
)
def test_get_upload_file_id_for_upload_file_document_rejects_missing_upload_file_id(db_session_with_containers):
def test_get_upload_file_id_for_upload_file_document_rejects_missing_upload_file_id(
db_session_with_containers: Session,
):
dataset = DocumentServiceIntegrationFactory.create_dataset(db_session_with_containers)
document = DocumentServiceIntegrationFactory.create_document(
db_session_with_containers,
@@ -248,7 +251,7 @@ def test_get_upload_file_id_for_upload_file_document_rejects_missing_upload_file
)
def test_get_upload_file_id_for_upload_file_document_returns_string_id(db_session_with_containers):
def test_get_upload_file_id_for_upload_file_document_returns_string_id(db_session_with_containers: Session):
dataset = DocumentServiceIntegrationFactory.create_dataset(db_session_with_containers)
document = DocumentServiceIntegrationFactory.create_document(
db_session_with_containers,
@@ -265,7 +268,9 @@ def test_get_upload_file_id_for_upload_file_document_returns_string_id(db_sessio
assert result == "99"
def test_get_upload_file_for_upload_file_document_raises_when_file_service_returns_nothing(db_session_with_containers):
def test_get_upload_file_for_upload_file_document_raises_when_file_service_returns_nothing(
db_session_with_containers: Session,
):
dataset = DocumentServiceIntegrationFactory.create_dataset(db_session_with_containers)
document = DocumentServiceIntegrationFactory.create_document(
db_session_with_containers,
@@ -278,7 +283,7 @@ def test_get_upload_file_for_upload_file_document_raises_when_file_service_retur
DocumentService._get_upload_file_for_upload_file_document(document)
def test_get_upload_file_for_upload_file_document_returns_upload_file(db_session_with_containers):
def test_get_upload_file_for_upload_file_document_returns_upload_file(db_session_with_containers: Session):
dataset = DocumentServiceIntegrationFactory.create_dataset(db_session_with_containers)
upload_file = DocumentServiceIntegrationFactory.create_upload_file(
db_session_with_containers,
@@ -296,7 +301,9 @@ def test_get_upload_file_for_upload_file_document_returns_upload_file(db_session
assert result.id == upload_file.id
def test_get_upload_files_by_document_id_for_zip_download_raises_for_missing_documents(db_session_with_containers):
def test_get_upload_files_by_document_id_for_zip_download_raises_for_missing_documents(
db_session_with_containers: Session,
):
dataset = DocumentServiceIntegrationFactory.create_dataset(db_session_with_containers)
with pytest.raises(NotFound, match="Document not found"):
@@ -307,7 +314,9 @@ def test_get_upload_files_by_document_id_for_zip_download_raises_for_missing_doc
)
def test_get_upload_files_by_document_id_for_zip_download_rejects_cross_tenant_access(db_session_with_containers):
def test_get_upload_files_by_document_id_for_zip_download_rejects_cross_tenant_access(
db_session_with_containers: Session,
):
dataset = DocumentServiceIntegrationFactory.create_dataset(db_session_with_containers)
upload_file = DocumentServiceIntegrationFactory.create_upload_file(
db_session_with_containers,
@@ -329,7 +338,9 @@ def test_get_upload_files_by_document_id_for_zip_download_rejects_cross_tenant_a
)
def test_get_upload_files_by_document_id_for_zip_download_rejects_missing_upload_files(db_session_with_containers):
def test_get_upload_files_by_document_id_for_zip_download_rejects_missing_upload_files(
db_session_with_containers: Session,
):
dataset = DocumentServiceIntegrationFactory.create_dataset(db_session_with_containers)
document = DocumentServiceIntegrationFactory.create_document(
db_session_with_containers,
@@ -345,7 +356,9 @@ def test_get_upload_files_by_document_id_for_zip_download_rejects_missing_upload
)
def test_get_upload_files_by_document_id_for_zip_download_returns_document_keyed_mapping(db_session_with_containers):
def test_get_upload_files_by_document_id_for_zip_download_returns_document_keyed_mapping(
db_session_with_containers: Session,
):
dataset = DocumentServiceIntegrationFactory.create_dataset(db_session_with_containers)
upload_file_a = DocumentServiceIntegrationFactory.create_upload_file(
db_session_with_containers,
@@ -395,7 +408,7 @@ def test_prepare_document_batch_download_zip_raises_not_found_for_missing_datase
def test_prepare_document_batch_download_zip_translates_permission_error_to_forbidden(
db_session_with_containers,
db_session_with_containers: Session,
current_user_mock,
):
dataset = DocumentServiceIntegrationFactory.create_dataset(
@@ -418,7 +431,7 @@ def test_prepare_document_batch_download_zip_translates_permission_error_to_forb
def test_prepare_document_batch_download_zip_returns_upload_files_in_requested_order(
db_session_with_containers,
db_session_with_containers: Session,
current_user_mock,
):
dataset = DocumentServiceIntegrationFactory.create_dataset(
@@ -461,7 +474,7 @@ def test_prepare_document_batch_download_zip_returns_upload_files_in_requested_o
assert download_name.endswith(".zip")
def test_get_document_by_dataset_id_returns_enabled_documents(db_session_with_containers):
def test_get_document_by_dataset_id_returns_enabled_documents(db_session_with_containers: Session):
dataset = DocumentServiceIntegrationFactory.create_dataset(db_session_with_containers)
enabled_document = DocumentServiceIntegrationFactory.create_document(
db_session_with_containers,
@@ -480,7 +493,9 @@ def test_get_document_by_dataset_id_returns_enabled_documents(db_session_with_co
assert [document.id for document in result] == [enabled_document.id]
def test_get_working_documents_by_dataset_id_returns_completed_enabled_unarchived_documents(db_session_with_containers):
def test_get_working_documents_by_dataset_id_returns_completed_enabled_unarchived_documents(
db_session_with_containers: Session,
):
dataset = DocumentServiceIntegrationFactory.create_dataset(db_session_with_containers)
available_document = DocumentServiceIntegrationFactory.create_document(
db_session_with_containers,
@@ -501,7 +516,7 @@ def test_get_working_documents_by_dataset_id_returns_completed_enabled_unarchive
assert [document.id for document in result] == [available_document.id]
def test_get_error_documents_by_dataset_id_returns_error_and_paused_documents(db_session_with_containers):
def test_get_error_documents_by_dataset_id_returns_error_and_paused_documents(db_session_with_containers: Session):
dataset = DocumentServiceIntegrationFactory.create_dataset(db_session_with_containers)
error_document = DocumentServiceIntegrationFactory.create_document(
db_session_with_containers,
@@ -526,7 +541,7 @@ def test_get_error_documents_by_dataset_id_returns_error_and_paused_documents(db
assert {document.id for document in result} == {error_document.id, paused_document.id}
def test_get_batch_documents_filters_by_current_user_tenant(db_session_with_containers):
def test_get_batch_documents_filters_by_current_user_tenant(db_session_with_containers: Session):
dataset = DocumentServiceIntegrationFactory.create_dataset(db_session_with_containers)
batch = f"batch-{uuid4()}"
matching_document = DocumentServiceIntegrationFactory.create_document(
@@ -549,7 +564,7 @@ def test_get_batch_documents_filters_by_current_user_tenant(db_session_with_cont
assert [document.id for document in result] == [matching_document.id]
def test_get_document_file_detail_returns_upload_file(db_session_with_containers):
def test_get_document_file_detail_returns_upload_file(db_session_with_containers: Session):
dataset = DocumentServiceIntegrationFactory.create_dataset(db_session_with_containers)
upload_file = DocumentServiceIntegrationFactory.create_upload_file(
db_session_with_containers,
@@ -563,7 +578,7 @@ def test_get_document_file_detail_returns_upload_file(db_session_with_containers
assert result.id == upload_file.id
def test_delete_document_emits_signal_and_commits(db_session_with_containers):
def test_delete_document_emits_signal_and_commits(db_session_with_containers: Session):
dataset = DocumentServiceIntegrationFactory.create_dataset(db_session_with_containers)
upload_file = DocumentServiceIntegrationFactory.create_upload_file(
db_session_with_containers,
@@ -588,7 +603,7 @@ def test_delete_document_emits_signal_and_commits(db_session_with_containers):
)
def test_delete_documents_ignores_empty_input(db_session_with_containers):
def test_delete_documents_ignores_empty_input(db_session_with_containers: Session):
dataset = DocumentServiceIntegrationFactory.create_dataset(db_session_with_containers)
with patch("services.dataset_service.batch_clean_document_task.delay") as delay:
@@ -597,7 +612,7 @@ def test_delete_documents_ignores_empty_input(db_session_with_containers):
delay.assert_not_called()
def test_delete_documents_deletes_rows_and_dispatches_cleanup_task(db_session_with_containers):
def test_delete_documents_deletes_rows_and_dispatches_cleanup_task(db_session_with_containers: Session):
dataset = DocumentServiceIntegrationFactory.create_dataset(db_session_with_containers)
dataset.chunk_structure = IndexStructureType.PARAGRAPH_INDEX
db_session_with_containers.commit()
@@ -637,14 +652,14 @@ def test_delete_documents_deletes_rows_and_dispatches_cleanup_task(db_session_wi
assert set(args[3]) == {upload_file_a.id, upload_file_b.id}
def test_get_documents_position_returns_next_position_when_documents_exist(db_session_with_containers):
def test_get_documents_position_returns_next_position_when_documents_exist(db_session_with_containers: Session):
dataset = DocumentServiceIntegrationFactory.create_dataset(db_session_with_containers)
DocumentServiceIntegrationFactory.create_document(db_session_with_containers, dataset=dataset, position=3)
assert DocumentService.get_documents_position(dataset.id) == 4
def test_get_documents_position_defaults_to_one_when_dataset_is_empty(db_session_with_containers):
def test_get_documents_position_defaults_to_one_when_dataset_is_empty(db_session_with_containers: Session):
dataset = DocumentServiceIntegrationFactory.create_dataset(db_session_with_containers)
assert DocumentService.get_documents_position(dataset.id) == 1
@@ -2,6 +2,7 @@ import datetime
from uuid import uuid4
from sqlalchemy import select
from sqlalchemy.orm import Session
from core.rag.index_processor.constant.index_type import IndexStructureType
from models.dataset import Dataset, Document
@@ -58,7 +59,7 @@ def _create_document(
return document
def test_build_display_status_filters_available(db_session_with_containers):
def test_build_display_status_filters_available(db_session_with_containers: Session):
dataset = _create_dataset(db_session_with_containers)
available_doc = _create_document(
db_session_with_containers,
@@ -97,7 +98,7 @@ def test_build_display_status_filters_available(db_session_with_containers):
assert [row.id for row in rows] == [available_doc.id]
def test_apply_display_status_filter_applies_when_status_present(db_session_with_containers):
def test_apply_display_status_filter_applies_when_status_present(db_session_with_containers: Session):
dataset = _create_dataset(db_session_with_containers)
waiting_doc = _create_document(
db_session_with_containers,
@@ -121,7 +122,7 @@ def test_apply_display_status_filter_applies_when_status_present(db_session_with
assert [row.id for row in rows] == [waiting_doc.id]
def test_apply_display_status_filter_returns_same_when_invalid(db_session_with_containers):
def test_apply_display_status_filter_returns_same_when_invalid(db_session_with_containers: Session):
dataset = _create_dataset(db_session_with_containers)
doc1 = _create_document(
db_session_with_containers,
@@ -7,6 +7,7 @@ import pytest
from sqlalchemy.orm import Session
from core.app.entities.app_invoke_entities import InvokeFrom
from models import TenantAccountRole
from models.account import Account, Tenant, TenantAccountJoin
from models.model import App, DefaultEndUserSessionID, EndUser
from services.end_user_service import EndUserService
@@ -16,7 +17,7 @@ class TestEndUserServiceFactory:
"""Factory class for creating test data and mock objects for end user service tests."""
@staticmethod
def create_app_and_account(db_session_with_containers):
def create_app_and_account(db_session_with_containers: Session):
tenant = Tenant(name=f"Tenant {uuid4()}")
db_session_with_containers.add(tenant)
db_session_with_containers.flush()
@@ -35,7 +36,7 @@ class TestEndUserServiceFactory:
tenant_join = TenantAccountJoin(
tenant_id=tenant.id,
account_id=account.id,
role="owner",
role=TenantAccountRole.OWNER,
current=True,
)
db_session_with_containers.add(tenant_join)
@@ -644,7 +644,7 @@ class TestFeatureService:
assert result.max_plugin_package_size == 15728640
# Verify default license status
assert result.license.status.value == "none"
assert result.license.status == "none"
assert result.license.expired_at == ""
assert result.license.workspaces.enabled is False
@@ -23,7 +23,7 @@ class TestFeedbackService:
"""Test FeedbackService methods."""
@pytest.fixture
def mock_db_session(self, monkeypatch):
def mock_db_session(self, monkeypatch: pytest.MonkeyPatch):
"""Mock database session."""
mock_session = mock.Mock()
monkeypatch.setattr(db, "session", mock_session)
@@ -122,7 +122,7 @@ class TestEmailDeliveryTestHandler:
with pytest.raises(DeliveryTestUnsupportedError):
handler.send_test(context=MagicMock(), method=MagicMock())
def test_send_test_feature_disabled(self, monkeypatch):
def test_send_test_feature_disabled(self, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr(
service_module.FeatureService,
"get_features",
@@ -137,7 +137,7 @@ class TestEmailDeliveryTestHandler:
with pytest.raises(DeliveryTestError, match="Email delivery is not available"):
handler.send_test(context=context, method=method)
def test_send_test_mail_not_inited(self, monkeypatch):
def test_send_test_mail_not_inited(self, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr(
service_module.FeatureService,
"get_features",
@@ -154,7 +154,7 @@ class TestEmailDeliveryTestHandler:
with pytest.raises(DeliveryTestError, match="Mail client is not initialized."):
handler.send_test(context=context, method=method)
def test_send_test_no_recipients(self, monkeypatch):
def test_send_test_no_recipients(self, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr(
service_module.FeatureService,
"get_features",
@@ -173,7 +173,7 @@ class TestEmailDeliveryTestHandler:
with pytest.raises(DeliveryTestError, match="No recipients configured"):
handler.send_test(context=context, method=method)
def test_send_test_success(self, monkeypatch):
def test_send_test_success(self, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr(
service_module.FeatureService,
"get_features",
@@ -209,7 +209,7 @@ class TestEmailDeliveryTestHandler:
assert kwargs["to"] == "test@example.com"
assert "RENDERED_Subj" in kwargs["subject"]
def test_send_test_sanitizes_subject(self, monkeypatch):
def test_send_test_sanitizes_subject(self, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr(
service_module.FeatureService,
"get_features",
@@ -1,6 +1,7 @@
from __future__ import annotations
import pytest
from sqlalchemy.orm import Session
from services.message_service import MessageService
from tests.test_containers_integration_tests.helpers.execution_extra_content import (
@@ -9,7 +10,7 @@ from tests.test_containers_integration_tests.helpers.execution_extra_content imp
@pytest.mark.usefixtures("flask_req_ctx_with_containers")
def test_pagination_returns_extra_contents(db_session_with_containers):
def test_pagination_returns_extra_contents(db_session_with_containers: Session):
fixture = create_human_input_message_fixture(db_session_with_containers)
pagination = MessageService.pagination_by_first_id(
@@ -16,7 +16,7 @@ from sqlalchemy.orm import Session
from core.rag.index_processor.constant.index_type import IndexStructureType, IndexTechniqueType
from extensions.ext_redis import redis_client
from models import Account, Tenant, TenantAccountJoin, TenantAccountRole
from models import Account, AccountStatus, Tenant, TenantAccountJoin, TenantAccountRole, TenantStatus
from models.dataset import Dataset, Document, DocumentSegment
from models.enums import DataSourceType, DocumentCreatedFrom, IndexingStatus, SegmentStatus
from tasks.create_segment_to_index_task import create_segment_to_index_task
@@ -73,7 +73,7 @@ class TestCreateSegmentToIndexTask:
email=fake.email(),
name=fake.name(),
interface_language="en-US",
status="active",
status=AccountStatus.ACTIVE,
)
db_session_with_containers.add(account)
@@ -82,7 +82,7 @@ class TestCreateSegmentToIndexTask:
# Create tenant
tenant = Tenant(
name=fake.company(),
status="normal",
status=TenantStatus.NORMAL,
plan="basic",
)
db_session_with_containers.add(tenant)
@@ -12,7 +12,7 @@ from sqlalchemy.orm import Session
from core.indexing_runner import DocumentIsPausedError
from core.rag.index_processor.constant.index_type import IndexTechniqueType
from enums.cloud_plan import CloudPlan
from models import Account, Tenant, TenantAccountJoin, TenantAccountRole
from models import Account, AccountStatus, Tenant, TenantAccountJoin, TenantAccountRole, TenantStatus
from models.dataset import Dataset, Document
from models.enums import DataSourceType, DocumentCreatedFrom, IndexingStatus
from tasks.document_indexing_task import (
@@ -54,7 +54,7 @@ class _TrackedSessionContext:
@pytest.fixture(autouse=True)
def _ensure_testcontainers_db(db_session_with_containers):
def _ensure_testcontainers_db(db_session_with_containers: Session):
"""Ensure this suite always runs on testcontainers infrastructure."""
return db_session_with_containers
@@ -121,12 +121,12 @@ class TestDatasetIndexingTaskIntegration:
email=fake.email(),
name=fake.name(),
interface_language="en-US",
status="active",
status=AccountStatus.ACTIVE,
)
db_session_with_containers.add(account)
db_session_with_containers.flush()
tenant = Tenant(name=fake.company(), status="normal")
tenant = Tenant(name=fake.company(), status=TenantStatus.NORMAL)
db_session_with_containers.add(tenant)
db_session_with_containers.flush()
@@ -5,6 +5,7 @@ from faker import Faker
from sqlalchemy.orm import Session
from libs.email_i18n import EmailType
from models import TenantStatus
from models.account import Account, Tenant, TenantAccountJoin, TenantAccountRole
from tasks.mail_account_deletion_task import send_account_deletion_verification_code, send_deletion_success_task
@@ -55,7 +56,7 @@ class TestMailAccountDeletionTask:
# Create tenant
tenant = Tenant(
name=fake.company(),
status="normal",
status=TenantStatus.NORMAL,
)
db_session_with_containers.add(tenant)
db_session_with_containers.commit()
@@ -18,6 +18,7 @@ from sqlalchemy import delete
from sqlalchemy.orm import Session
from libs.email_i18n import EmailType
from models import AccountStatus, TenantStatus
from models.account import Account, Tenant, TenantAccountJoin, TenantAccountRole
from tasks.mail_email_code_login import send_email_code_login_mail_task
@@ -91,7 +92,7 @@ class TestSendEmailCodeLoginMailTask:
email=fake.email(),
name=fake.name(),
interface_language="en-US",
status="active",
status=AccountStatus.ACTIVE,
)
db_session_with_containers.add(account)
@@ -120,7 +121,7 @@ class TestSendEmailCodeLoginMailTask:
tenant = Tenant(
name=fake.company(),
plan="basic",
status="normal",
status=TenantStatus.NORMAL,
)
db_session_with_containers.add(tenant)
@@ -31,7 +31,7 @@ from tasks.mail_human_input_delivery_task import dispatch_human_input_email_task
@pytest.fixture(autouse=True)
def cleanup_database(db_session_with_containers):
def cleanup_database(db_session_with_containers: Session):
db_session_with_containers.execute(delete(HumanInputFormRecipient))
db_session_with_containers.execute(delete(HumanInputDelivery))
db_session_with_containers.execute(delete(HumanInputForm))
@@ -43,7 +43,7 @@ def cleanup_database(db_session_with_containers):
db_session_with_containers.commit()
def _create_workspace_member(db_session_with_containers):
def _create_workspace_member(db_session_with_containers: Session):
account = Account(
email="owner@example.com",
name="Owner",
@@ -21,7 +21,7 @@ from tasks.remove_app_and_related_data_task import (
@pytest.fixture(autouse=True)
def cleanup_database(db_session_with_containers):
def cleanup_database(db_session_with_containers: Session):
db_session_with_containers.execute(delete(WorkflowDraftVariable))
db_session_with_containers.execute(delete(WorkflowDraftVariableFile))
db_session_with_containers.execute(delete(UploadFile))
@@ -30,7 +30,7 @@ def cleanup_database(db_session_with_containers):
db_session_with_containers.commit()
def _create_tenant_and_app(db_session_with_containers):
def _create_tenant_and_app(db_session_with_containers: Session):
tenant = Tenant(name=f"test_tenant_{uuid.uuid4()}")
db_session_with_containers.add(tenant)
db_session_with_containers.flush()
@@ -57,7 +57,7 @@ class TestGuessFileInfoFromResponse:
(False, "bin"),
],
)
def test_generated_filename_when_missing(self, monkeypatch, magic_available, expected_ext):
def test_generated_filename_when_missing(self, monkeypatch: pytest.MonkeyPatch, magic_available, expected_ext):
if magic_available:
if helpers.magic is None:
pytest.skip("python-magic is not installed, cannot run 'magic_available=True' test variant")
@@ -155,7 +155,7 @@ class TestMagicImportWarnings:
)
def test_magic_import_warning_per_platform(
self,
monkeypatch,
monkeypatch: pytest.MonkeyPatch,
platform_name,
expected_message,
):
@@ -101,7 +101,7 @@ def test_register_schema_models_registers_multiple_models():
assert called_names == ["UserModel", "ProductModel"]
def test_register_schema_models_calls_register_schema_model(monkeypatch):
def test_register_schema_models_calls_register_schema_model(monkeypatch: pytest.MonkeyPatch):
from controllers.common.schema import register_schema_models
namespace = MagicMock(spec=Namespace)
@@ -68,7 +68,7 @@ def _segment():
)
def test_get_segment_with_summary(monkeypatch):
def test_get_segment_with_summary(monkeypatch: pytest.MonkeyPatch):
segment = _segment()
summary = SimpleNamespace(summary_content="summary")
@@ -3,6 +3,7 @@ from unittest.mock import MagicMock, PropertyMock, patch
import pytest
from flask import Flask
from pytest_mock import MockerFixture
from werkzeug.exceptions import NotFound
from controllers.console import console_ns
@@ -35,7 +36,7 @@ def dataset():
@pytest.fixture(autouse=True)
def bypass_decorators(mocker):
def bypass_decorators(mocker: MockerFixture):
"""Bypass all decorators on the API method."""
mocker.patch(
"controllers.console.datasets.hit_testing.setup_required",
@@ -56,7 +57,7 @@ def bypass_decorators(mocker):
class TestHitTestingApi:
def test_hit_testing_success(self, app, dataset, dataset_id):
def test_hit_testing_success(self, app: Flask, dataset, dataset_id):
api = HitTestingApi()
method = unwrap(api.post)
@@ -99,7 +100,7 @@ class TestHitTestingApi:
assert "records" in result
assert result["records"] == []
def test_hit_testing_success_with_optional_record_fields(self, app, dataset, dataset_id):
def test_hit_testing_success_with_optional_record_fields(self, app: Flask, dataset, dataset_id):
api = HitTestingApi()
method = unwrap(api.post)
@@ -150,7 +151,7 @@ class TestHitTestingApi:
assert result["query"] == payload["query"]
assert result["records"] == records
def test_hit_testing_dataset_not_found(self, app, dataset_id):
def test_hit_testing_dataset_not_found(self, app: Flask, dataset_id):
api = HitTestingApi()
method = unwrap(api.post)
@@ -175,7 +176,7 @@ class TestHitTestingApi:
with pytest.raises(NotFound, match="Dataset not found"):
method(api, dataset_id)
def test_hit_testing_invalid_args(self, app, dataset, dataset_id):
def test_hit_testing_invalid_args(self, app: Flask, dataset, dataset_id):
api = HitTestingApi()
method = unwrap(api.post)
@@ -3,6 +3,7 @@ from unittest.mock import MagicMock, PropertyMock, patch
import pytest
from flask import Flask
from pytest_mock import MockerFixture
from werkzeug.exceptions import NotFound
from controllers.console import console_ns
@@ -60,7 +61,7 @@ def metadata_id():
@pytest.fixture(autouse=True)
def bypass_decorators(mocker):
def bypass_decorators(mocker: MockerFixture):
"""Bypass setup/login/license decorators."""
mocker.patch(
"controllers.console.datasets.metadata.setup_required",
@@ -2,6 +2,7 @@ from unittest.mock import Mock, PropertyMock, patch
import pytest
from flask import Flask
from pytest_mock import MockerFixture
from controllers.console import console_ns
from controllers.console.datasets.error import WebsiteCrawlError
@@ -31,7 +32,7 @@ def app():
@pytest.fixture(autouse=True)
def bypass_auth_and_setup(mocker):
def bypass_auth_and_setup(mocker: MockerFixture):
"""Bypass setup/login/account decorators."""
mocker.patch(
"controllers.console.datasets.website.login_required",
@@ -48,7 +49,7 @@ def bypass_auth_and_setup(mocker):
class TestWebsiteCrawlApi:
def test_crawl_success(self, app, mocker):
def test_crawl_success(self, app, mocker: MockerFixture):
api = WebsiteCrawlApi()
method = unwrap(api.post)
@@ -85,7 +86,7 @@ class TestWebsiteCrawlApi:
assert status == 200
assert result["job_id"] == "job-1"
def test_crawl_invalid_payload(self, app, mocker):
def test_crawl_invalid_payload(self, app, mocker: MockerFixture):
api = WebsiteCrawlApi()
method = unwrap(api.post)
@@ -113,7 +114,7 @@ class TestWebsiteCrawlApi:
with pytest.raises(WebsiteCrawlError, match="invalid payload"):
method(api)
def test_crawl_service_error(self, app, mocker):
def test_crawl_service_error(self, app, mocker: MockerFixture):
api = WebsiteCrawlApi()
method = unwrap(api.post)
@@ -150,7 +151,7 @@ class TestWebsiteCrawlApi:
class TestWebsiteCrawlStatusApi:
def test_get_status_success(self, app, mocker):
def test_get_status_success(self, app, mocker: MockerFixture):
api = WebsiteCrawlStatusApi()
method = unwrap(api.get)
@@ -181,7 +182,7 @@ class TestWebsiteCrawlStatusApi:
assert status == 200
assert result["status"] == "completed"
def test_get_status_invalid_provider(self, app, mocker):
def test_get_status_invalid_provider(self, app, mocker: MockerFixture):
api = WebsiteCrawlStatusApi()
method = unwrap(api.get)
@@ -203,7 +204,7 @@ class TestWebsiteCrawlStatusApi:
with pytest.raises(WebsiteCrawlError, match="invalid provider"):
method(api, job_id)
def test_get_status_service_error(self, app, mocker):
def test_get_status_service_error(self, app, mocker: MockerFixture):
api = WebsiteCrawlStatusApi()
method = unwrap(api.get)
@@ -1,6 +1,7 @@
from unittest.mock import Mock
import pytest
from pytest_mock import MockerFixture
from controllers.console.datasets.error import PipelineNotFoundError
from controllers.console.datasets.wraps import get_rag_pipeline
@@ -16,7 +17,7 @@ class TestGetRagPipeline:
with pytest.raises(ValueError, match="missing pipeline_id"):
dummy_view()
def test_pipeline_not_found(self, mocker):
def test_pipeline_not_found(self, mocker: MockerFixture):
@get_rag_pipeline
def dummy_view(**kwargs):
return "ok"
@@ -34,7 +35,7 @@ class TestGetRagPipeline:
with pytest.raises(PipelineNotFoundError):
dummy_view(pipeline_id="pipeline-1")
def test_pipeline_found_and_injected(self, mocker):
def test_pipeline_found_and_injected(self, mocker: MockerFixture):
pipeline = Mock(spec=Pipeline)
pipeline.id = "pipeline-1"
pipeline.tenant_id = "tenant-1"
@@ -57,7 +58,7 @@ class TestGetRagPipeline:
assert result is pipeline
def test_pipeline_id_removed_from_kwargs(self, mocker):
def test_pipeline_id_removed_from_kwargs(self, mocker: MockerFixture):
pipeline = Mock(spec=Pipeline)
@get_rag_pipeline
@@ -79,7 +80,7 @@ class TestGetRagPipeline:
assert result == "ok"
def test_pipeline_id_cast_to_string(self, mocker):
def test_pipeline_id_cast_to_string(self, mocker: MockerFixture):
pipeline = Mock(spec=Pipeline)
@get_rag_pipeline
@@ -4,6 +4,7 @@ import uuid
from unittest.mock import Mock, PropertyMock, patch
import pytest
from pytest_mock import MockerFixture
from werkzeug.exceptions import NotFound, Unauthorized
from controllers.console.admin import (
@@ -18,7 +19,7 @@ from models.model import App, InstalledApp, RecommendedApp
@pytest.fixture(autouse=True)
def bypass_only_edition_cloud(mocker):
def bypass_only_edition_cloud(mocker: MockerFixture):
"""
Bypass only_edition_cloud decorator by setting EDITION to "CLOUD".
"""
@@ -29,7 +30,7 @@ def bypass_only_edition_cloud(mocker):
@pytest.fixture
def mock_admin_auth(mocker):
def mock_admin_auth(mocker: MockerFixture):
"""
Provide valid admin authentication for controller tests.
"""
@@ -44,7 +45,7 @@ def mock_admin_auth(mocker):
@pytest.fixture
def mock_console_payload(mocker):
def mock_console_payload(mocker: MockerFixture):
payload = {
"app_id": str(uuid.uuid4()),
"language": "en-US",
@@ -62,7 +63,7 @@ def mock_console_payload(mocker):
@pytest.fixture
def mock_banner_payload(mocker):
def mock_banner_payload(mocker: MockerFixture):
mocker.patch(
"flask_restx.namespace.Namespace.payload",
new_callable=PropertyMock,
@@ -78,7 +79,7 @@ def mock_banner_payload(mocker):
@pytest.fixture
def mock_session_factory(mocker):
def mock_session_factory(mocker: MockerFixture):
mock_session = Mock()
mock_session.execute = Mock()
mock_session.add = Mock()
@@ -97,7 +98,7 @@ class TestDeleteExploreBannerApi:
def setup_method(self):
self.api = DeleteExploreBannerApi()
def test_delete_banner_not_found(self, mocker, mock_admin_auth):
def test_delete_banner_not_found(self, mocker: MockerFixture, mock_admin_auth):
mocker.patch(
"controllers.console.admin.db.session.execute",
return_value=Mock(scalar_one_or_none=lambda: None),
@@ -106,7 +107,7 @@ class TestDeleteExploreBannerApi:
with pytest.raises(NotFound, match="is not found"):
self.api.delete(uuid.uuid4())
def test_delete_banner_success(self, mocker, mock_admin_auth):
def test_delete_banner_success(self, mocker: MockerFixture, mock_admin_auth):
mock_banner = Mock()
mocker.patch(
@@ -126,7 +127,7 @@ class TestInsertExploreBannerApi:
def setup_method(self):
self.api = InsertExploreBannerApi()
def test_insert_banner_success(self, mocker, mock_admin_auth, mock_banner_payload):
def test_insert_banner_success(self, mocker: MockerFixture, mock_admin_auth, mock_banner_payload):
mocker.patch("controllers.console.admin.db.session.add")
mocker.patch("controllers.console.admin.db.session.commit")
@@ -168,7 +169,7 @@ class TestInsertExploreAppApiDelete:
def setup_method(self):
self.api = InsertExploreAppApi()
def test_delete_when_not_in_explore(self, mocker, mock_admin_auth):
def test_delete_when_not_in_explore(self, mocker: MockerFixture, mock_admin_auth):
mocker.patch(
"controllers.console.admin.session_factory.create_session",
return_value=Mock(
@@ -183,7 +184,7 @@ class TestInsertExploreAppApiDelete:
assert status == 204
assert response["result"] == "success"
def test_delete_when_in_explore_with_trial_app(self, mocker, mock_admin_auth):
def test_delete_when_in_explore_with_trial_app(self, mocker: MockerFixture, mock_admin_auth):
"""Test deleting an app from explore that has a trial app."""
app_id = uuid.uuid4()
@@ -225,7 +226,7 @@ class TestInsertExploreAppApiDelete:
assert response["result"] == "success"
assert mock_app.is_public is False
def test_delete_with_installed_apps(self, mocker, mock_admin_auth):
def test_delete_with_installed_apps(self, mocker: MockerFixture, mock_admin_auth):
"""Test deleting an app that has installed apps in other tenants."""
app_id = uuid.uuid4()
@@ -270,7 +271,7 @@ class TestInsertExploreAppListApi:
def setup_method(self):
self.api = InsertExploreAppListApi()
def test_app_not_found(self, mocker, mock_admin_auth, mock_console_payload):
def test_app_not_found(self, mocker: MockerFixture, mock_admin_auth, mock_console_payload):
mocker.patch(
"controllers.console.admin.db.session.execute",
return_value=Mock(scalar_one_or_none=lambda: None),
@@ -281,7 +282,7 @@ class TestInsertExploreAppListApi:
def test_create_recommended_app(
self,
mocker,
mocker: MockerFixture,
mock_admin_auth,
mock_console_payload,
):
@@ -318,7 +319,9 @@ class TestInsertExploreAppListApi:
assert response["result"] == "success"
assert mock_app.is_public is True
def test_update_recommended_app(self, mocker, mock_admin_auth, mock_console_payload, mock_session_factory):
def test_update_recommended_app(
self, mocker: MockerFixture, mock_admin_auth, mock_console_payload, mock_session_factory
):
mock_app = Mock(spec=App)
mock_app.id = "app-id"
mock_app.site = None
@@ -344,7 +347,7 @@ class TestInsertExploreAppListApi:
def test_site_data_overrides_payload(
self,
mocker,
mocker: MockerFixture,
mock_admin_auth,
mock_console_payload,
mock_session_factory,
@@ -381,7 +384,7 @@ class TestInsertExploreAppListApi:
def test_create_trial_app_when_can_trial_enabled(
self,
mocker,
mocker: MockerFixture,
mock_admin_auth,
mock_console_payload,
mock_session_factory,
@@ -413,7 +416,7 @@ class TestInsertExploreAppListApi:
def test_update_recommended_app_with_trial(
self,
mocker,
mocker: MockerFixture,
mock_admin_auth,
mock_console_payload,
mock_session_factory,
@@ -450,7 +453,7 @@ class TestInsertExploreAppListApi:
def test_update_recommended_app_without_trial(
self,
mocker,
mocker: MockerFixture,
mock_admin_auth,
mock_console_payload,
mock_session_factory,
@@ -1,3 +1,4 @@
from pytest_mock import MockerFixture
from werkzeug.exceptions import Unauthorized
@@ -11,7 +12,7 @@ def unwrap(func):
class TestFeatureApi:
def test_get_tenant_features_success(self, mocker):
def test_get_tenant_features_success(self, mocker: MockerFixture):
from controllers.console.feature import FeatureApi
mocker.patch(
@@ -32,7 +33,7 @@ class TestFeatureApi:
class TestSystemFeatureApi:
def test_get_system_features_authenticated(self, mocker):
def test_get_system_features_authenticated(self, mocker: MockerFixture):
"""
current_user.is_authenticated == True
"""
@@ -56,7 +57,7 @@ class TestSystemFeatureApi:
assert result == {"features": {"sys_feature": True}}
def test_get_system_features_unauthenticated(self, mocker):
def test_get_system_features_unauthenticated(self, mocker: MockerFixture):
"""
current_user.is_authenticated raises Unauthorized
"""
@@ -32,7 +32,7 @@ class TestDefaultModelApi:
with (
app.test_request_context(
"/",
query_string={"model_type": ModelType.LLM.value},
query_string={"model_type": ModelType.LLM},
),
patch(
"controllers.console.workspace.models.current_account_with_tenant",
@@ -53,7 +53,7 @@ class TestDefaultModelApi:
payload = {
"model_settings": [
{
"model_type": ModelType.LLM.value,
"model_type": ModelType.LLM,
"provider": "openai",
"model": "gpt-4",
}
@@ -77,7 +77,7 @@ class TestDefaultModelApi:
method = unwrap(api.get)
with (
app.test_request_context("/", query_string={"model_type": ModelType.LLM.value}),
app.test_request_context("/", query_string={"model_type": ModelType.LLM}),
patch("controllers.console.workspace.models.current_account_with_tenant", return_value=(MagicMock(), "t1")),
patch("controllers.console.workspace.models.ModelProviderService") as service,
):
@@ -113,7 +113,7 @@ class TestModelProviderModelApi:
payload = {
"model": "gpt-4",
"model_type": ModelType.LLM.value,
"model_type": ModelType.LLM,
"load_balancing": {
"configs": [{"weight": 1}],
"enabled": True,
@@ -139,7 +139,7 @@ class TestModelProviderModelApi:
payload = {
"model": "gpt-4",
"model_type": ModelType.LLM.value,
"model_type": ModelType.LLM,
}
with (
@@ -180,7 +180,7 @@ class TestModelProviderModelCredentialApi:
"/",
query_string={
"model": "gpt-4",
"model_type": ModelType.LLM.value,
"model_type": ModelType.LLM,
},
),
patch(
@@ -208,7 +208,7 @@ class TestModelProviderModelCredentialApi:
payload = {
"model": "gpt-4",
"model_type": ModelType.LLM.value,
"model_type": ModelType.LLM,
"credentials": {"key": "val"},
}
@@ -229,7 +229,7 @@ class TestModelProviderModelCredentialApi:
method = unwrap(api.get)
with (
app.test_request_context("/", query_string={"model": "gpt", "model_type": ModelType.LLM.value}),
app.test_request_context("/", query_string={"model": "gpt", "model_type": ModelType.LLM}),
patch("controllers.console.workspace.models.current_account_with_tenant", return_value=(MagicMock(), "t1")),
patch("controllers.console.workspace.models.ModelProviderService") as service,
patch("controllers.console.workspace.models.ModelLoadBalancingService") as lb,
@@ -248,7 +248,7 @@ class TestModelProviderModelCredentialApi:
payload = {
"model": "gpt",
"model_type": ModelType.LLM.value,
"model_type": ModelType.LLM,
"credential_id": "123e4567-e89b-12d3-a456-426614174000",
}
@@ -269,7 +269,7 @@ class TestModelProviderModelCredentialSwitchApi:
payload = {
"model": "gpt-4",
"model_type": ModelType.LLM.value,
"model_type": ModelType.LLM,
"credential_id": "abc",
}
@@ -293,7 +293,7 @@ class TestModelEnableDisableApis:
payload = {
"model": "gpt-4",
"model_type": ModelType.LLM.value,
"model_type": ModelType.LLM,
}
with (
@@ -314,7 +314,7 @@ class TestModelEnableDisableApis:
payload = {
"model": "gpt-4",
"model_type": ModelType.LLM.value,
"model_type": ModelType.LLM,
}
with (
@@ -337,7 +337,7 @@ class TestModelProviderModelValidateApi:
payload = {
"model": "gpt-4",
"model_type": ModelType.LLM.value,
"model_type": ModelType.LLM,
"credentials": {"key": "val"},
}
@@ -360,7 +360,7 @@ class TestModelProviderModelValidateApi:
payload = {
"model": model_name,
"model_type": ModelType.LLM.value,
"model_type": ModelType.LLM,
"credentials": {},
}
@@ -412,7 +412,7 @@ class TestParameterAndAvailableModels:
):
service_mock.return_value.get_models_by_model_type.return_value = []
result = method(api, ModelType.LLM.value)
result = method(api, ModelType.LLM)
assert "data" in result
@@ -442,6 +442,6 @@ class TestParameterAndAvailableModels:
):
service.return_value.get_models_by_model_type.return_value = []
result = method(api, ModelType.LLM.value)
result = method(api, ModelType.LLM)
assert result["data"] == []
@@ -189,7 +189,7 @@ class TestGetUserTenant:
"""Test get_user_tenant decorator"""
@patch("controllers.inner_api.plugin.wraps.Tenant")
def test_should_inject_tenant_and_user_models(self, mock_tenant_class, app: Flask, monkeypatch):
def test_should_inject_tenant_and_user_models(self, mock_tenant_class, app: Flask, monkeypatch: pytest.MonkeyPatch):
"""Test that decorator injects tenant_model and user_model into kwargs"""
# Arrange
@@ -244,7 +244,9 @@ class TestGetUserTenant:
protected_view()
@patch("controllers.inner_api.plugin.wraps.Tenant")
def test_should_use_default_session_id_when_user_id_empty(self, mock_tenant_class, app: Flask, monkeypatch):
def test_should_use_default_session_id_when_user_id_empty(
self, mock_tenant_class, app: Flask, monkeypatch: pytest.MonkeyPatch
):
"""Test that default session ID is used when user_id is empty string"""
# Arrange
@@ -340,7 +340,7 @@ class TestConversationAppModeValidation:
@pytest.mark.parametrize(
"mode",
[
AppMode.CHAT.value,
AppMode.CHAT,
AppMode.AGENT_CHAT.value,
AppMode.ADVANCED_CHAT.value,
],
@@ -365,7 +365,7 @@ class TestConversationAppModeValidation:
app raises NotChatAppError.
"""
app = Mock(spec=App)
app.mode = AppMode.COMPLETION.value
app.mode = AppMode.COMPLETION
app_mode = AppMode.value_of(app.mode)
assert app_mode not in {AppMode.CHAT, AppMode.AGENT_CHAT, AppMode.ADVANCED_CHAT}
@@ -498,7 +498,7 @@ class TestConversationApiController:
def test_list_not_chat(self, app) -> None:
api = ConversationApi()
handler = _unwrap(api.get)
app_model = SimpleNamespace(mode=AppMode.COMPLETION.value)
app_model = SimpleNamespace(mode=AppMode.COMPLETION)
end_user = SimpleNamespace()
with app.test_request_context("/conversations", method="GET"):
@@ -531,7 +531,7 @@ class TestConversationApiController:
api = ConversationApi()
handler = _unwrap(api.get)
app_model = SimpleNamespace(mode=AppMode.CHAT.value)
app_model = SimpleNamespace(mode=AppMode.CHAT)
end_user = SimpleNamespace()
with app.test_request_context(
@@ -546,7 +546,7 @@ class TestConversationDetailApiController:
def test_delete_not_chat(self, app) -> None:
api = ConversationDetailApi()
handler = _unwrap(api.delete)
app_model = SimpleNamespace(mode=AppMode.COMPLETION.value)
app_model = SimpleNamespace(mode=AppMode.COMPLETION)
end_user = SimpleNamespace()
with app.test_request_context("/conversations/1", method="DELETE"):
@@ -562,7 +562,7 @@ class TestConversationDetailApiController:
api = ConversationDetailApi()
handler = _unwrap(api.delete)
app_model = SimpleNamespace(mode=AppMode.CHAT.value)
app_model = SimpleNamespace(mode=AppMode.CHAT)
end_user = SimpleNamespace()
with app.test_request_context("/conversations/1", method="DELETE"):
@@ -580,7 +580,7 @@ class TestConversationRenameApiController:
api = ConversationRenameApi()
handler = _unwrap(api.post)
app_model = SimpleNamespace(mode=AppMode.CHAT.value)
app_model = SimpleNamespace(mode=AppMode.CHAT)
end_user = SimpleNamespace()
with app.test_request_context(
@@ -596,7 +596,7 @@ class TestConversationVariablesApiController:
def test_not_chat(self, app) -> None:
api = ConversationVariablesApi()
handler = _unwrap(api.get)
app_model = SimpleNamespace(mode=AppMode.COMPLETION.value)
app_model = SimpleNamespace(mode=AppMode.COMPLETION)
end_user = SimpleNamespace()
with app.test_request_context("/conversations/1/variables", method="GET"):
@@ -612,7 +612,7 @@ class TestConversationVariablesApiController:
api = ConversationVariablesApi()
handler = _unwrap(api.get)
app_model = SimpleNamespace(mode=AppMode.CHAT.value)
app_model = SimpleNamespace(mode=AppMode.CHAT)
end_user = SimpleNamespace()
with app.test_request_context(
@@ -645,7 +645,7 @@ class TestConversationVariablesApiController:
api = ConversationVariablesApi()
handler = _unwrap(api.get)
app_model = SimpleNamespace(mode=AppMode.CHAT.value)
app_model = SimpleNamespace(mode=AppMode.CHAT)
end_user = SimpleNamespace()
with app.test_request_context(
@@ -671,7 +671,7 @@ class TestConversationVariableDetailApiController:
api = ConversationVariableDetailApi()
handler = _unwrap(api.put)
app_model = SimpleNamespace(mode=AppMode.CHAT.value)
app_model = SimpleNamespace(mode=AppMode.CHAT)
end_user = SimpleNamespace()
with app.test_request_context(
@@ -697,7 +697,7 @@ class TestConversationVariableDetailApiController:
api = ConversationVariableDetailApi()
handler = _unwrap(api.put)
app_model = SimpleNamespace(mode=AppMode.CHAT.value)
app_model = SimpleNamespace(mode=AppMode.CHAT)
end_user = SimpleNamespace()
with app.test_request_context(
@@ -731,7 +731,7 @@ class TestConversationVariableDetailApiController:
api = ConversationVariableDetailApi()
handler = _unwrap(api.put)
app_model = SimpleNamespace(mode=AppMode.CHAT.value)
app_model = SimpleNamespace(mode=AppMode.CHAT)
end_user = SimpleNamespace()
with app.test_request_context(
@@ -3,6 +3,7 @@ from unittest.mock import Mock
from uuid import UUID, uuid4
import pytest
from pytest_mock import MockerFixture
from controllers.service_api.end_user.end_user import EndUserApi
from controllers.service_api.end_user.error import EndUserNotFoundError
@@ -21,7 +22,9 @@ class TestEndUserApi:
app.tenant_id = str(uuid4())
return app
def test_get_end_user_returns_all_attributes(self, mocker, resource: EndUserApi, app_model: App) -> None:
def test_get_end_user_returns_all_attributes(
self, mocker: MockerFixture, resource: EndUserApi, app_model: App
) -> None:
end_user = Mock(spec=EndUser)
end_user.id = str(uuid4())
end_user.tenant_id = app_model.tenant_id
@@ -54,7 +57,7 @@ class TestEndUserApi:
assert result["created_at"].startswith("2024-01-01T00:00:00")
assert result["updated_at"].startswith("2024-01-02T00:00:00")
def test_get_end_user_not_found(self, mocker, resource: EndUserApi, app_model: App) -> None:
def test_get_end_user_not_found(self, mocker: MockerFixture, resource: EndUserApi, app_model: App) -> None:
mocker.patch("controllers.service_api.end_user.end_user.EndUserService.get_end_user_by_id", return_value=None)
with pytest.raises(EndUserNotFoundError):
@@ -12,12 +12,13 @@ from types import SimpleNamespace
from unittest.mock import MagicMock
import pytest
from pytest_mock import MockerFixture
from core.agent.output_parser.cot_output_parser import CotAgentOutputParser
@pytest.fixture
def mock_action_class(mocker):
def mock_action_class(mocker: MockerFixture):
mock_action = MagicMock()
mocker.patch(
"core.agent.output_parser.cot_output_parser.AgentScratchpadUnit.Action",
@@ -3,6 +3,7 @@
from unittest.mock import MagicMock
import pytest
from pytest_mock import MockerFixture
from core.agent.strategy.plugin import PluginAgentStrategy
@@ -213,7 +214,9 @@ class TestInvoke:
(None, None, "msg"),
],
)
def test_invoke_optional_arguments(self, strategy, mocker, conversation_id, app_id, message_id) -> None:
def test_invoke_optional_arguments(
self, strategy, mocker: MockerFixture, conversation_id, app_id, message_id
) -> None:
mock_manager = MagicMock()
mock_manager.invoke = MagicMock(return_value=iter([]))
@@ -3,6 +3,7 @@ from decimal import Decimal
from unittest.mock import MagicMock
import pytest
from pytest_mock import MockerFixture
import core.agent.base_agent_runner as module
from core.agent.base_agent_runner import BaseAgentRunner
@@ -13,7 +14,7 @@ from core.agent.base_agent_runner import BaseAgentRunner
@pytest.fixture
def mock_db_session(mocker):
def mock_db_session(mocker: MockerFixture):
session = mocker.MagicMock()
mocker.patch.object(module.db, "session", session)
return session
@@ -41,13 +42,13 @@ def runner(mocker, mock_db_session):
class TestRepack:
def test_sets_empty_if_none(self, runner, mocker):
def test_sets_empty_if_none(self, runner, mocker: MockerFixture):
entity = mocker.MagicMock()
entity.app_config.prompt_template.simple_prompt_template = None
result = runner._repack_app_generate_entity(entity)
assert result.app_config.prompt_template.simple_prompt_template == ""
def test_keeps_existing(self, runner, mocker):
def test_keeps_existing(self, runner, mocker: MockerFixture):
entity = mocker.MagicMock()
entity.app_config.prompt_template.simple_prompt_template = "abc"
result = runner._repack_app_generate_entity(entity)
@@ -60,7 +61,7 @@ class TestRepack:
class TestUpdatePromptTool:
def build_param(self, mocker, **kwargs):
def build_param(self, mocker: MockerFixture, **kwargs):
p = mocker.MagicMock()
p.form = kwargs.get("form")
@@ -75,7 +76,7 @@ class TestUpdatePromptTool:
p.required = kwargs.get("required", False)
return p
def test_skip_non_llm(self, runner, mocker):
def test_skip_non_llm(self, runner, mocker: MockerFixture):
tool = mocker.MagicMock()
param = self.build_param(mocker, form="NOT_LLM")
tool.get_runtime_parameters.return_value = [param]
@@ -86,7 +87,7 @@ class TestUpdatePromptTool:
result = runner.update_prompt_message_tool(tool, prompt_tool)
assert result.parameters["properties"] == {}
def test_enum_and_required(self, runner, mocker):
def test_enum_and_required(self, runner, mocker: MockerFixture):
option = mocker.MagicMock(value="opt1")
param = self.build_param(
mocker,
@@ -104,7 +105,7 @@ class TestUpdatePromptTool:
result = runner.update_prompt_message_tool(tool, prompt_tool)
assert "p1" in result.parameters["required"]
def test_skip_file_type_param(self, runner, mocker):
def test_skip_file_type_param(self, runner, mocker: MockerFixture):
tool = mocker.MagicMock()
param = self.build_param(mocker, form=module.ToolParameter.ToolParameterForm.LLM)
param.type = module.ToolParameter.ToolParameterType.FILE
@@ -116,7 +117,7 @@ class TestUpdatePromptTool:
result = runner.update_prompt_message_tool(tool, prompt_tool)
assert result.parameters["properties"] == {}
def test_duplicate_required_not_duplicated(self, runner, mocker):
def test_duplicate_required_not_duplicated(self, runner, mocker: MockerFixture):
tool = mocker.MagicMock()
param = self.build_param(
@@ -141,7 +142,7 @@ class TestUpdatePromptTool:
class TestCreateAgentThought:
def test_with_files(self, runner, mock_db_session, mocker):
def test_with_files(self, runner, mock_db_session, mocker: MockerFixture):
mock_thought = mocker.MagicMock(id=10)
mocker.patch.object(module, "MessageAgentThought", return_value=mock_thought)
@@ -149,7 +150,7 @@ class TestCreateAgentThought:
assert result == "10"
assert runner.agent_thought_count == 1
def test_without_files(self, runner, mock_db_session, mocker):
def test_without_files(self, runner, mock_db_session, mocker: MockerFixture):
mock_thought = mocker.MagicMock(id=11)
mocker.patch.object(module, "MessageAgentThought", return_value=mock_thought)
@@ -163,7 +164,7 @@ class TestCreateAgentThought:
class TestSaveAgentThought:
def setup_agent(self, mocker):
def setup_agent(self, mocker: MockerFixture):
agent = mocker.MagicMock()
agent.tool = "tool1;tool2"
agent.tool_labels = {}
@@ -175,7 +176,7 @@ class TestSaveAgentThought:
with pytest.raises(ValueError):
runner.save_agent_thought("id", None, None, None, None, None, None, [], None)
def test_full_update(self, runner, mock_db_session, mocker):
def test_full_update(self, runner, mock_db_session, mocker: MockerFixture):
agent = self.setup_agent(mocker)
mock_db_session.scalar.return_value = agent
@@ -210,7 +211,7 @@ class TestSaveAgentThought:
assert agent.tokens == 3
assert "tool1" in json.loads(agent.tool_labels_str)
def test_label_fallback_when_none(self, runner, mock_db_session, mocker):
def test_label_fallback_when_none(self, runner, mock_db_session, mocker: MockerFixture):
agent = self.setup_agent(mocker)
agent.tool = "unknown_tool"
mock_db_session.scalar.return_value = agent
@@ -220,7 +221,7 @@ class TestSaveAgentThought:
labels = json.loads(agent.tool_labels_str)
assert "unknown_tool" in labels
def test_json_failure_paths(self, runner, mock_db_session, mocker):
def test_json_failure_paths(self, runner, mock_db_session, mocker: MockerFixture):
agent = self.setup_agent(mocker)
mock_db_session.scalar.return_value = agent
@@ -241,13 +242,13 @@ class TestSaveAgentThought:
assert mock_db_session.commit.called
def test_messages_ids_none(self, runner, mock_db_session, mocker):
def test_messages_ids_none(self, runner, mock_db_session, mocker: MockerFixture):
agent = self.setup_agent(mocker)
mock_db_session.scalar.return_value = agent
runner.save_agent_thought("id", None, None, None, None, None, None, None, None)
assert mock_db_session.commit.called
def test_success_dict_serialization(self, runner, mock_db_session, mocker):
def test_success_dict_serialization(self, runner, mock_db_session, mocker: MockerFixture):
agent = self.setup_agent(mocker)
mock_db_session.scalar.return_value = agent
@@ -273,19 +274,19 @@ class TestSaveAgentThought:
class TestOrganizeUserPrompt:
def test_no_files(self, runner, mock_db_session, mocker):
def test_no_files(self, runner, mock_db_session, mocker: MockerFixture):
mock_db_session.scalars.return_value.all.return_value = []
msg = mocker.MagicMock(id="1", query="hello", app_model_config=None)
result = runner.organize_agent_user_prompt(msg)
assert result.content == "hello"
def test_with_files_no_config(self, runner, mock_db_session, mocker):
def test_with_files_no_config(self, runner, mock_db_session, mocker: MockerFixture):
mock_db_session.scalars.return_value.all.return_value = [mocker.MagicMock()]
msg = mocker.MagicMock(id="1", query="hello", app_model_config=None)
result = runner.organize_agent_user_prompt(msg)
assert result.content == "hello"
def test_image_detail_low_fallback(self, runner, mock_db_session, mocker):
def test_image_detail_low_fallback(self, runner, mock_db_session, mocker: MockerFixture):
mock_db_session.scalars.return_value.all.return_value = [mocker.MagicMock()]
file_config = mocker.MagicMock()
file_config.image_config = mocker.MagicMock(detail=None)
@@ -305,27 +306,27 @@ class TestOrganizeUserPrompt:
class TestOrganizeHistory:
def test_empty(self, runner, mock_db_session, mocker):
def test_empty(self, runner, mock_db_session, mocker: MockerFixture):
mock_db_session.execute.return_value.scalars.return_value.all.return_value = []
mocker.patch.object(module, "extract_thread_messages", return_value=[])
result = runner.organize_agent_history([])
assert result == []
def test_with_answer_only(self, runner, mock_db_session, mocker):
def test_with_answer_only(self, runner, mock_db_session, mocker: MockerFixture):
msg = mocker.MagicMock(id="m1", answer="ans", agent_thoughts=[], app_model_config=None)
mock_db_session.execute.return_value.scalars.return_value.all.return_value = [msg]
mocker.patch.object(module, "extract_thread_messages", return_value=[msg])
result = runner.organize_agent_history([])
assert any(isinstance(x, module.AssistantPromptMessage) for x in result)
def test_skip_current_message(self, runner, mock_db_session, mocker):
def test_skip_current_message(self, runner, mock_db_session, mocker: MockerFixture):
msg = mocker.MagicMock(id="msg_current", agent_thoughts=[], answer="ans", app_model_config=None)
mock_db_session.execute.return_value.scalars.return_value.all.return_value = [msg]
mocker.patch.object(module, "extract_thread_messages", return_value=[msg])
result = runner.organize_agent_history([])
assert result == []
def test_with_tool_calls_invalid_json(self, runner, mock_db_session, mocker):
def test_with_tool_calls_invalid_json(self, runner, mock_db_session, mocker: MockerFixture):
thought = mocker.MagicMock(
tool="tool1",
tool_input="invalid",
@@ -341,7 +342,7 @@ class TestOrganizeHistory:
result = runner.organize_agent_history([])
assert isinstance(result, list)
def test_empty_tool_name_split(self, runner, mock_db_session, mocker):
def test_empty_tool_name_split(self, runner, mock_db_session, mocker: MockerFixture):
thought = mocker.MagicMock(tool=";", thought="thinking")
msg = mocker.MagicMock(id="m5", agent_thoughts=[thought], answer=None, app_model_config=None)
@@ -350,7 +351,7 @@ class TestOrganizeHistory:
result = runner.organize_agent_history([])
assert isinstance(result, list)
def test_valid_json_tool_flow(self, runner, mock_db_session, mocker):
def test_valid_json_tool_flow(self, runner, mock_db_session, mocker: MockerFixture):
thought = mocker.MagicMock(
tool="tool1",
tool_input=json.dumps({"tool1": {"x": 1}}),
@@ -379,7 +380,7 @@ class TestOrganizeHistory:
class TestConvertToolToPromptMessageTool:
def test_basic_conversion(self, runner, mocker):
def test_basic_conversion(self, runner, mocker: MockerFixture):
tool = mocker.MagicMock(tool_name="tool1")
runtime_param = mocker.MagicMock()
@@ -404,7 +405,7 @@ class TestConvertToolToPromptMessageTool:
prompt_tool, entity = runner._convert_tool_to_prompt_message_tool(tool)
assert entity == tool_entity
def test_full_conversion_multiple_params(self, runner, mocker):
def test_full_conversion_multiple_params(self, runner, mocker: MockerFixture):
tool = mocker.MagicMock(tool_name="tool1")
# LLM param with input_schema override
@@ -441,7 +442,7 @@ class TestConvertToolToPromptMessageTool:
class TestInitPromptToolsExtended:
def test_agent_tool_branch(self, runner, mocker):
def test_agent_tool_branch(self, runner, mocker: MockerFixture):
agent_tool = mocker.MagicMock(tool_name="agent_tool")
runner.app_config.agent = mocker.MagicMock(tools=[agent_tool])
mocker.patch.object(runner, "_convert_tool_to_prompt_message_tool", return_value=(MagicMock(), "entity"))
@@ -449,7 +450,7 @@ class TestInitPromptToolsExtended:
tools, prompts = runner._init_prompt_tools()
assert "agent_tool" in tools
def test_exception_in_conversion(self, runner, mocker):
def test_exception_in_conversion(self, runner, mocker: MockerFixture):
agent_tool = mocker.MagicMock(tool_name="bad_tool")
runner.app_config.agent = mocker.MagicMock(tools=[agent_tool])
mocker.patch.object(runner, "_convert_tool_to_prompt_message_tool", side_effect=Exception)
@@ -464,7 +465,7 @@ class TestInitPromptToolsExtended:
class TestAdditionalCoverage:
def test_update_prompt_with_input_schema(self, runner, mocker):
def test_update_prompt_with_input_schema(self, runner, mocker: MockerFixture):
tool = mocker.MagicMock()
param = mocker.MagicMock()
@@ -487,7 +488,7 @@ class TestAdditionalCoverage:
result = runner.update_prompt_message_tool(tool, prompt_tool)
assert result.parameters["properties"]["p1"]["type"] == "number"
def test_save_agent_thought_existing_labels(self, runner, mock_db_session, mocker):
def test_save_agent_thought_existing_labels(self, runner, mock_db_session, mocker: MockerFixture):
agent = mocker.MagicMock()
agent.tool = "tool1"
agent.tool_labels = {"tool1": {"en_US": "existing"}}
@@ -498,7 +499,7 @@ class TestAdditionalCoverage:
labels = json.loads(agent.tool_labels_str)
assert labels["tool1"]["en_US"] == "existing"
def test_save_agent_thought_tool_meta_string(self, runner, mock_db_session, mocker):
def test_save_agent_thought_tool_meta_string(self, runner, mock_db_session, mocker: MockerFixture):
agent = mocker.MagicMock()
agent.tool = "tool1"
agent.tool_labels = {}
@@ -508,7 +509,7 @@ class TestAdditionalCoverage:
runner.save_agent_thought("id", None, None, None, None, "meta_string", None, [], None)
assert agent.tool_meta_str == "meta_string"
def test_convert_dataset_retriever_tool(self, runner, mocker):
def test_convert_dataset_retriever_tool(self, runner, mocker: MockerFixture):
ds_tool = mocker.MagicMock()
ds_tool.entity.identity.name = "ds"
ds_tool.entity.description.llm = "desc"
@@ -525,7 +526,7 @@ class TestAdditionalCoverage:
prompt = runner._convert_dataset_retriever_tool_to_prompt_message_tool(ds_tool)
assert prompt is not None
def test_organize_user_prompt_with_file_objects(self, runner, mock_db_session, mocker):
def test_organize_user_prompt_with_file_objects(self, runner, mock_db_session, mocker: MockerFixture):
mock_db_session.scalars.return_value.all.return_value = [mocker.MagicMock()]
file_config = mocker.MagicMock()
@@ -544,7 +545,7 @@ class TestAdditionalCoverage:
result = runner.organize_agent_user_prompt(msg)
assert result is not None
def test_organize_history_without_tool_names(self, runner, mock_db_session, mocker):
def test_organize_history_without_tool_names(self, runner, mock_db_session, mocker: MockerFixture):
thought = mocker.MagicMock(tool=None, thought="thinking")
msg = mocker.MagicMock(id="m3", agent_thoughts=[thought], answer=None, app_model_config=None)
@@ -554,7 +555,7 @@ class TestAdditionalCoverage:
result = runner.organize_agent_history([])
assert isinstance(result, list)
def test_organize_history_multiple_tools_split(self, runner, mock_db_session, mocker):
def test_organize_history_multiple_tools_split(self, runner, mock_db_session, mocker: MockerFixture):
thought = mocker.MagicMock(
tool="tool1;tool2",
tool_input=json.dumps({"tool1": {}, "tool2": {}}),
@@ -572,7 +573,7 @@ class TestAdditionalCoverage:
# ================= Additional Surgical Coverage =================
def test_convert_tool_select_enum_branch(self, runner, mocker):
def test_convert_tool_select_enum_branch(self, runner, mocker: MockerFixture):
tool = mocker.MagicMock(tool_name="tool1")
param = mocker.MagicMock()
@@ -599,7 +600,7 @@ class TestAdditionalCoverage:
class TestConvertDatasetRetrieverTool:
def test_required_param_added(self, runner, mocker):
def test_required_param_added(self, runner, mocker: MockerFixture):
ds_tool = mocker.MagicMock()
ds_tool.entity.identity.name = "ds"
ds_tool.entity.description.llm = "desc"
@@ -619,7 +620,7 @@ class TestConvertDatasetRetrieverTool:
class TestBaseAgentRunnerInit:
def test_init_sets_stream_tool_call_and_files(self, mocker):
def test_init_sets_stream_tool_call_and_files(self, mocker: MockerFixture):
session = mocker.MagicMock()
session.scalar.return_value = 2
mocker.patch.object(module.db, "session", session)
@@ -662,7 +663,7 @@ class TestBaseAgentRunnerInit:
class TestBaseAgentRunnerCoverage:
def test_convert_tool_skips_non_llm_param(self, runner, mocker):
def test_convert_tool_skips_non_llm_param(self, runner, mocker: MockerFixture):
tool = mocker.MagicMock(tool_name="tool1")
param = mocker.MagicMock()
@@ -680,7 +681,7 @@ class TestBaseAgentRunnerCoverage:
assert prompt_tool.parameters["properties"] == {}
def test_init_prompt_tools_adds_dataset_tools(self, runner, mocker):
def test_init_prompt_tools_adds_dataset_tools(self, runner, mocker: MockerFixture):
dataset_tool = mocker.MagicMock()
dataset_tool.entity.identity.name = "ds"
runner.dataset_tools = [dataset_tool]
@@ -692,7 +693,7 @@ class TestBaseAgentRunnerCoverage:
assert tools["ds"] == dataset_tool
assert len(prompt_tools) == 1
def test_update_prompt_message_tool_select_enum(self, runner, mocker):
def test_update_prompt_message_tool_select_enum(self, runner, mocker: MockerFixture):
tool = mocker.MagicMock()
option1 = mocker.MagicMock(value="A")
@@ -716,7 +717,7 @@ class TestBaseAgentRunnerCoverage:
assert result.parameters["properties"]["select_param"]["enum"] == ["A", "B"]
def test_save_agent_thought_json_dumps_fallbacks(self, runner, mock_db_session, mocker):
def test_save_agent_thought_json_dumps_fallbacks(self, runner, mock_db_session, mocker: MockerFixture):
agent = mocker.MagicMock()
agent.tool = "tool1"
agent.tool_labels = {}
@@ -754,7 +755,7 @@ class TestBaseAgentRunnerCoverage:
assert isinstance(agent.observation, str)
assert isinstance(agent.tool_meta_str, str)
def test_save_agent_thought_skips_empty_tool_name(self, runner, mock_db_session, mocker):
def test_save_agent_thought_skips_empty_tool_name(self, runner, mock_db_session, mocker: MockerFixture):
agent = mocker.MagicMock()
agent.tool = "tool1;;"
agent.tool_labels = {}
@@ -768,7 +769,7 @@ class TestBaseAgentRunnerCoverage:
labels = json.loads(agent.tool_labels_str)
assert "" not in labels
def test_organize_history_includes_system_prompt(self, runner, mock_db_session, mocker):
def test_organize_history_includes_system_prompt(self, runner, mock_db_session, mocker: MockerFixture):
mock_db_session.execute.return_value.scalars.return_value.all.return_value = []
mocker.patch.object(module, "extract_thread_messages", return_value=[])
@@ -778,7 +779,7 @@ class TestBaseAgentRunnerCoverage:
assert system_message in result
def test_organize_history_tool_inputs_and_observation_none(self, runner, mock_db_session, mocker):
def test_organize_history_tool_inputs_and_observation_none(self, runner, mock_db_session, mocker: MockerFixture):
thought = mocker.MagicMock(
tool="tool1",
tool_input=None,
@@ -2,6 +2,7 @@ import json
from unittest.mock import MagicMock
import pytest
from pytest_mock import MockerFixture
from core.agent.cot_agent_runner import CotAgentRunner
from core.agent.entities import AgentScratchpadUnit
@@ -25,7 +26,7 @@ class DummyRunner(CotAgentRunner):
@pytest.fixture
def runner(mocker):
def runner(mocker: MockerFixture):
# Prevent BaseAgentRunner __init__ from hitting database
mocker.patch(
"core.agent.base_agent_runner.BaseAgentRunner.organize_agent_history",
@@ -165,7 +166,7 @@ class TestHandleInvokeAction:
response, meta = runner._handle_invoke_action(action, {}, [])
assert "there is not a tool named" in response
def test_tool_with_json_string_args(self, runner, mocker):
def test_tool_with_json_string_args(self, runner, mocker: MockerFixture):
action = AgentScratchpadUnit.Action(action_name="tool", action_input=json.dumps({"a": 1}))
tool_instance = MagicMock()
tool_instances = {"tool": tool_instance}
@@ -180,7 +181,7 @@ class TestHandleInvokeAction:
class TestOrganizeHistoricPromptMessages:
def test_empty_history(self, runner, mocker):
def test_empty_history(self, runner, mocker: MockerFixture):
mocker.patch(
"core.agent.cot_agent_runner.AgentHistoryPromptTransform.get_prompt",
return_value=[],
@@ -190,7 +191,7 @@ class TestOrganizeHistoricPromptMessages:
class TestRun:
def test_run_handles_empty_parser_output(self, runner, mocker):
def test_run_handles_empty_parser_output(self, runner, mocker: MockerFixture):
message = MagicMock()
message.id = "msg-id"
@@ -202,7 +203,7 @@ class TestRun:
results = list(runner.run(message, "query", {}))
assert isinstance(results, list)
def test_run_with_action_and_tool_invocation(self, runner, mocker):
def test_run_with_action_and_tool_invocation(self, runner, mocker: MockerFixture):
message = MagicMock()
message.id = "msg-id"
@@ -223,7 +224,7 @@ class TestRun:
with pytest.raises(AgentMaxIterationError):
list(runner.run(message, "query", {"tool": MagicMock()}))
def test_run_respects_max_iteration_boundary(self, runner, mocker):
def test_run_respects_max_iteration_boundary(self, runner, mocker: MockerFixture):
runner.app_config.agent.max_iteration = 1
message = MagicMock()
message.id = "msg-id"
@@ -245,7 +246,7 @@ class TestRun:
with pytest.raises(AgentMaxIterationError):
list(runner.run(message, "query", {"tool": MagicMock()}))
def test_run_basic_flow(self, runner, mocker):
def test_run_basic_flow(self, runner, mocker: MockerFixture):
message = MagicMock()
message.id = "msg-id"
@@ -257,7 +258,7 @@ class TestRun:
results = list(runner.run(message, "query", {"name": "John"}))
assert results
def test_run_max_iteration_error(self, runner, mocker):
def test_run_max_iteration_error(self, runner, mocker: MockerFixture):
runner.app_config.agent.max_iteration = 0
message = MagicMock()
message.id = "msg-id"
@@ -272,7 +273,7 @@ class TestRun:
with pytest.raises(AgentMaxIterationError):
list(runner.run(message, "query", {}))
def test_run_increase_usage_aggregation(self, runner, mocker):
def test_run_increase_usage_aggregation(self, runner, mocker: MockerFixture):
message = MagicMock()
message.id = "msg-id"
runner.app_config.agent.max_iteration = 2
@@ -329,7 +330,7 @@ class TestRun:
assert final_usage.completion_price == 2
assert final_usage.total_price == 4
def test_run_when_no_action_branch(self, runner, mocker):
def test_run_when_no_action_branch(self, runner, mocker: MockerFixture):
message = MagicMock()
message.id = "msg-id"
@@ -341,7 +342,7 @@ class TestRun:
results = list(runner.run(message, "query", {}))
assert results[-1].delta.message.content == ""
def test_run_usage_missing_key_branch(self, runner, mocker):
def test_run_usage_missing_key_branch(self, runner, mocker: MockerFixture):
message = MagicMock()
message.id = "msg-id"
@@ -354,7 +355,7 @@ class TestRun:
list(runner.run(message, "query", {}))
def test_run_prompt_tool_update_branch(self, runner, mocker):
def test_run_prompt_tool_update_branch(self, runner, mocker: MockerFixture):
message = MagicMock()
message.id = "msg-id"
@@ -410,7 +411,7 @@ class TestRun:
class TestInitReactState:
def test_init_react_state_resets_state(self, runner, mocker):
def test_init_react_state_resets_state(self, runner, mocker: MockerFixture):
mocker.patch.object(runner, "_organize_historic_prompt_messages", return_value=["historic"])
runner._agent_scratchpad = ["old"]
runner._query = "old"
@@ -423,7 +424,7 @@ class TestInitReactState:
class TestHandleInvokeActionExtended:
def test_tool_with_invalid_json_string_args(self, runner, mocker):
def test_tool_with_invalid_json_string_args(self, runner, mocker: MockerFixture):
action = AgentScratchpadUnit.Action(action_name="tool", action_input="not-json")
tool_instance = MagicMock()
tool_instances = {"tool": tool_instance}
@@ -457,7 +458,7 @@ class TestFillInputsEdgeCases:
class TestOrganizeHistoricPromptMessagesExtended:
def test_user_message_flushes_scratchpad(self, runner, mocker):
def test_user_message_flushes_scratchpad(self, runner, mocker: MockerFixture):
from graphon.model_runtime.entities.message_entities import UserPromptMessage
user_message = UserPromptMessage(content="Hi")
@@ -480,7 +481,7 @@ class TestOrganizeHistoricPromptMessagesExtended:
with pytest.raises(NotImplementedError):
runner._organize_historic_prompt_messages([])
def test_agent_history_transform_invocation(self, runner, mocker):
def test_agent_history_transform_invocation(self, runner, mocker: MockerFixture):
mock_transform = MagicMock()
mock_transform.get_prompt.return_value = []
@@ -495,7 +496,7 @@ class TestOrganizeHistoricPromptMessagesExtended:
class TestRunAdditionalBranches:
def test_run_with_no_action_final_answer_empty(self, runner, mocker):
def test_run_with_no_action_final_answer_empty(self, runner, mocker: MockerFixture):
message = MagicMock()
message.id = "msg-id"
@@ -507,7 +508,7 @@ class TestRunAdditionalBranches:
results = list(runner.run(message, "query", {}))
assert any(hasattr(r, "delta") for r in results)
def test_run_with_final_answer_action_string(self, runner, mocker):
def test_run_with_final_answer_action_string(self, runner, mocker: MockerFixture):
message = MagicMock()
message.id = "msg-id"
@@ -521,7 +522,7 @@ class TestRunAdditionalBranches:
results = list(runner.run(message, "query", {}))
assert results[-1].delta.message.content == "done"
def test_run_with_final_answer_action_dict(self, runner, mocker):
def test_run_with_final_answer_action_dict(self, runner, mocker: MockerFixture):
message = MagicMock()
message.id = "msg-id"
@@ -535,7 +536,7 @@ class TestRunAdditionalBranches:
results = list(runner.run(message, "query", {}))
assert json.loads(results[-1].delta.message.content) == {"a": 1}
def test_run_with_string_final_answer(self, runner, mocker):
def test_run_with_string_final_answer(self, runner, mocker: MockerFixture):
message = MagicMock()
message.id = "msg-id"
@@ -1,6 +1,7 @@
from unittest.mock import MagicMock, patch
import pytest
from pytest_mock import MockerFixture
from core.agent.cot_chat_agent_runner import CotChatAgentRunner
from graphon.model_runtime.entities.message_entities import TextPromptMessageContent
@@ -55,7 +56,7 @@ def runner():
class TestOrganizeSystemPrompt:
def test_organize_system_prompt_success(self, runner, mocker):
def test_organize_system_prompt_success(self, runner, mocker: MockerFixture):
first_prompt = "Instruction: {{instruction}}, Tools: {{tools}}, Names: {{tool_names}}"
runner.app_config = DummyAppConfig(DummyAgentConfig(DummyPrompt(first_prompt)))
@@ -154,7 +155,7 @@ class TestOrganizeUserQuery:
class TestOrganizePromptMessages:
def test_no_scratchpad(self, runner, mocker):
def test_no_scratchpad(self, runner, mocker: MockerFixture):
runner.app_config = DummyAppConfig(DummyAgentConfig(DummyPrompt("{{instruction}}")))
runner._organize_system_prompt = MagicMock(return_value="system")
runner._organize_user_query = MagicMock(return_value=["query"])
@@ -164,7 +165,7 @@ class TestOrganizePromptMessages:
assert "query" in result
runner._organize_historic_prompt_messages.assert_called_once()
def test_with_final_scratchpad(self, runner, mocker):
def test_with_final_scratchpad(self, runner, mocker: MockerFixture):
runner.app_config = DummyAppConfig(DummyAgentConfig(DummyPrompt("{{instruction}}")))
runner._organize_system_prompt = MagicMock(return_value="system")
runner._organize_user_query = MagicMock(return_value=["query"])
@@ -177,7 +178,7 @@ class TestOrganizePromptMessages:
combined = "".join([m.content for m in assistant_msgs if isinstance(m.content, str)])
assert "Final Answer: done" in combined
def test_with_thought_action_observation(self, runner, mocker):
def test_with_thought_action_observation(self, runner, mocker: MockerFixture):
runner.app_config = DummyAppConfig(DummyAgentConfig(DummyPrompt("{{instruction}}")))
runner._organize_system_prompt = MagicMock(return_value="system")
runner._organize_user_query = MagicMock(return_value=["query"])
@@ -197,7 +198,7 @@ class TestOrganizePromptMessages:
assert "Action: action" in combined
assert "Observation: observe" in combined
def test_multiple_units_mixed(self, runner, mocker):
def test_multiple_units_mixed(self, runner, mocker: MockerFixture):
runner.app_config = DummyAppConfig(DummyAgentConfig(DummyPrompt("{{instruction}}")))
runner._organize_system_prompt = MagicMock(return_value="system")
runner._organize_user_query = MagicMock(return_value=["query"])
@@ -1,6 +1,7 @@
import json
import pytest
from pytest_mock import MockerFixture
from core.agent.cot_completion_agent_runner import CotCompletionAgentRunner
from graphon.model_runtime.entities.message_entities import (
@@ -74,7 +75,7 @@ class TestOrganizeInstructionPrompt:
class TestOrganizeHistoricPrompt:
def test_with_user_and_assistant_string(self, runner, mocker):
def test_with_user_and_assistant_string(self, runner, mocker: MockerFixture):
user_msg = UserPromptMessage(content="Hello")
assistant_msg = AssistantPromptMessage(content="Hi there")
@@ -89,7 +90,7 @@ class TestOrganizeHistoricPrompt:
assert "Question: Hello" in result
assert "Hi there" in result
def test_assistant_list_with_text_content(self, runner, mocker):
def test_assistant_list_with_text_content(self, runner, mocker: MockerFixture):
text_content = TextPromptMessageContent(data="Partial answer")
assistant_msg = AssistantPromptMessage(content=[text_content])
@@ -103,7 +104,7 @@ class TestOrganizeHistoricPrompt:
assert "Partial answer" in result
def test_assistant_list_with_non_text_content_ignored(self, runner, mocker):
def test_assistant_list_with_non_text_content_ignored(self, runner, mocker: MockerFixture):
non_text_content = ImagePromptMessageContent(format="url", mime_type="image/png")
assistant_msg = AssistantPromptMessage(content=[non_text_content])
@@ -116,7 +117,7 @@ class TestOrganizeHistoricPrompt:
result = runner._organize_historic_prompt()
assert result == ""
def test_empty_history(self, runner, mocker):
def test_empty_history(self, runner, mocker: MockerFixture):
mocker.patch.object(
runner,
"_organize_historic_prompt_messages",
@@ -136,7 +137,7 @@ class TestOrganizePromptMessages:
def test_full_flow_with_scratchpad(
self,
runner,
mocker,
mocker: MockerFixture,
dummy_app_config_factory,
dummy_agent_config_factory,
dummy_prompt_entity_factory,
@@ -171,7 +172,12 @@ class TestOrganizePromptMessages:
assert "Question: What is Python?" in content
def test_no_scratchpad(
self, runner, mocker, dummy_app_config_factory, dummy_agent_config_factory, dummy_prompt_entity_factory
self,
runner,
mocker: MockerFixture,
dummy_app_config_factory,
dummy_agent_config_factory,
dummy_prompt_entity_factory,
):
template = "SYS {{historic_messages}} {{agent_scratchpad}} {{query}}"
@@ -198,7 +204,7 @@ class TestOrganizePromptMessages:
def test_partial_scratchpad_units(
self,
runner,
mocker,
mocker: MockerFixture,
thought,
action,
observation,
@@ -3,6 +3,7 @@ from typing import Any
from unittest.mock import MagicMock
import pytest
from pytest_mock import MockerFixture
from core.agent.errors import AgentMaxIterationError
from core.agent.fc_agent_runner import FunctionCallAgentRunner
@@ -68,7 +69,7 @@ class DummyResult:
@pytest.fixture
def runner(mocker):
def runner(mocker: MockerFixture):
# Completely bypass BaseAgentRunner __init__ to avoid DB / Flask context
mocker.patch(
"core.agent.base_agent_runner.BaseAgentRunner.__init__",
@@ -230,7 +231,7 @@ class TestOrganizeUserQuery:
result = runner._organize_user_query(None, [])
assert len(result) == 1
def test_with_files_uses_image_detail_config(self, runner, mocker):
def test_with_files_uses_image_detail_config(self, runner, mocker: MockerFixture):
file_content = TextPromptMessageContent(data="file-content")
mock_to_prompt = mocker.patch(
"core.agent.fc_agent_runner.file_manager.to_prompt_message_content",
@@ -352,7 +353,7 @@ class TestRunMethod:
assert len(outputs) == 1
assert runner.save_agent_thought.call_args.kwargs["thought"] == "hi"
def test_run_streaming_tool_call_inputs_type_error(self, runner, mocker):
def test_run_streaming_tool_call_inputs_type_error(self, runner, mocker: MockerFixture):
message = MagicMock(id="m1")
runner.stream_tool_call = True
@@ -398,7 +399,7 @@ class TestRunMethod:
outputs = list(runner.run(message, "query"))
assert len(outputs) >= 1
def test_run_with_tool_instance_and_files(self, runner, mocker):
def test_run_with_tool_instance_and_files(self, runner, mocker: MockerFixture):
message = MagicMock(id="m1")
tool_call = MagicMock()
@@ -9,6 +9,7 @@ mocking; ensure entity invariants and validation rules remain stable.
import pytest
from pydantic import ValidationError
from pytest_mock import MockerFixture
from core.agent.plugin_entities import (
AgentFeature,
@@ -28,12 +29,12 @@ from core.tools.entities.tool_entities import ToolIdentity, ToolProviderIdentity
@pytest.fixture
def mock_identity(mocker):
def mock_identity(mocker: MockerFixture):
return mocker.MagicMock(spec=AgentStrategyIdentity)
@pytest.fixture
def mock_provider_identity(mocker):
def mock_provider_identity(mocker: MockerFixture):
return mocker.MagicMock(spec=AgentStrategyProviderIdentity)
@@ -47,7 +48,7 @@ class TestAgentStrategyParameterType:
"enum_member",
list(AgentStrategyParameter.AgentStrategyParameterType),
)
def test_as_normal_type_calls_external_function(self, mocker, enum_member) -> None:
def test_as_normal_type_calls_external_function(self, mocker: MockerFixture, enum_member) -> None:
mock_func = mocker.patch(
"core.agent.plugin_entities.as_normal_type",
return_value="normalized",
@@ -58,7 +59,7 @@ class TestAgentStrategyParameterType:
mock_func.assert_called_once_with(enum_member)
assert result == "normalized"
def test_as_normal_type_propagates_exception(self, mocker) -> None:
def test_as_normal_type_propagates_exception(self, mocker: MockerFixture) -> None:
enum_member = AgentStrategyParameter.AgentStrategyParameterType.STRING
mocker.patch(
"core.agent.plugin_entities.as_normal_type",
@@ -79,7 +80,7 @@ class TestAgentStrategyParameterType:
(AgentStrategyParameter.AgentStrategyParameterType.FILES, []),
],
)
def test_cast_value_calls_external_function(self, mocker, enum_member, value) -> None:
def test_cast_value_calls_external_function(self, mocker: MockerFixture, enum_member, value) -> None:
mock_func = mocker.patch(
"core.agent.plugin_entities.cast_parameter_value",
return_value="casted",
@@ -90,7 +91,7 @@ class TestAgentStrategyParameterType:
mock_func.assert_called_once_with(enum_member, value)
assert result == "casted"
def test_cast_value_propagates_exception(self, mocker) -> None:
def test_cast_value_propagates_exception(self, mocker: MockerFixture) -> None:
enum_member = AgentStrategyParameter.AgentStrategyParameterType.STRING
mocker.patch(
"core.agent.plugin_entities.cast_parameter_value",
@@ -136,7 +137,7 @@ class TestAgentStrategyParameter:
assert any(error["loc"] == ("type",) for error in exc_info.value.errors())
def test_init_frontend_parameter_calls_external(self, mocker) -> None:
def test_init_frontend_parameter_calls_external(self, mocker: MockerFixture) -> None:
mock_func = mocker.patch(
"core.agent.plugin_entities.init_frontend_parameter",
return_value="frontend",
@@ -153,7 +154,7 @@ class TestAgentStrategyParameter:
mock_func.assert_called_once_with(param, param.type, "value")
assert result == "frontend"
def test_init_frontend_parameter_propagates_exception(self, mocker) -> None:
def test_init_frontend_parameter_propagates_exception(self, mocker: MockerFixture) -> None:
mocker.patch(
"core.agent.plugin_entities.init_frontend_parameter",
side_effect=RuntimeError("error"),
@@ -10,7 +10,7 @@ class TestGetParametersFromFeatureDict:
"""Test suite for get_parameters_from_feature_dict"""
@pytest.fixture
def mock_config(self, monkeypatch):
def mock_config(self, monkeypatch: pytest.MonkeyPatch):
"""Mock dify_config values"""
mock = MagicMock()
mock.UPLOAD_IMAGE_FILE_SIZE_LIMIT = 1
@@ -23,7 +23,7 @@ class TestGetParametersFromFeatureDict:
return mock
@pytest.fixture
def mock_default_file_limits(self, monkeypatch):
def mock_default_file_limits(self, monkeypatch: pytest.MonkeyPatch):
"""Mock DEFAULT_FILE_NUMBER_LIMITS constant"""
monkeypatch.setattr(parameters_mapping, "DEFAULT_FILE_NUMBER_LIMITS", 99)
return 99
@@ -1,6 +1,7 @@
from unittest.mock import MagicMock
import pytest
from pytest_mock import MockerFixture
from core.app.app_config.common.sensitive_word_avoidance.manager import (
SensitiveWordAvoidanceConfigManager,
@@ -26,7 +27,7 @@ class TestSensitiveWordAvoidanceConfigManagerConvert:
# Assert
assert result is None
def test_convert_returns_entity_when_enabled(self, mocker):
def test_convert_returns_entity_when_enabled(self, mocker: MockerFixture):
# Arrange
mock_entity = MagicMock()
mocker.patch(
@@ -48,7 +49,7 @@ class TestSensitiveWordAvoidanceConfigManagerConvert:
# Assert
assert result == mock_entity
def test_convert_enabled_without_type_or_config(self, mocker):
def test_convert_enabled_without_type_or_config(self, mocker: MockerFixture):
# Arrange
mock_entity = MagicMock()
patched = mocker.patch(
@@ -135,7 +136,7 @@ class TestSensitiveWordAvoidanceConfigManagerValidateAndSetDefaults:
with pytest.raises(ValueError, match="must be a dict"):
SensitiveWordAvoidanceConfigManager.validate_and_set_defaults(tenant_id="tenant1", config=config)
def test_validate_calls_moderation_factory(self, mocker):
def test_validate_calls_moderation_factory(self, mocker: MockerFixture):
# Arrange
mock_validate = mocker.patch(
"core.app.app_config.common.sensitive_word_avoidance.manager.ModerationFactory.validate_config"
@@ -159,7 +160,7 @@ class TestSensitiveWordAvoidanceConfigManagerValidateAndSetDefaults:
assert result_config["sensitive_word_avoidance"]["enabled"] is True
assert fields == ["sensitive_word_avoidance"]
def test_validate_sets_empty_dict_when_config_none(self, mocker):
def test_validate_sets_empty_dict_when_config_none(self, mocker: MockerFixture):
# Arrange
mock_validate = mocker.patch(
"core.app.app_config.common.sensitive_word_avoidance.manager.ModerationFactory.validate_config"
@@ -179,7 +180,7 @@ class TestSensitiveWordAvoidanceConfigManagerValidateAndSetDefaults:
# Assert
mock_validate.assert_called_once_with(name="mock_type", tenant_id="tenant1", config={})
def test_validate_only_structure_validate_skips_factory(self, mocker):
def test_validate_only_structure_validate_skips_factory(self, mocker: MockerFixture):
# Arrange
mock_validate = mocker.patch(
"core.app.app_config.common.sensitive_word_avoidance.manager.ModerationFactory.validate_config"
@@ -1,6 +1,7 @@
from unittest.mock import MagicMock
import pytest
from pytest_mock import MockerFixture
from core.app.app_config.easy_ui_based_app.agent.manager import AgentConfigManager
@@ -84,7 +85,7 @@ class TestAgentConfigManagerConvert:
assert result.strategy.name == "CHAIN_OF_THOUGHT"
def test_convert_skips_disabled_tools(self, mocker, base_config):
def test_convert_skips_disabled_tools(self, mocker: MockerFixture, base_config):
# Patch AgentEntity to bypass pydantic validation
mock_agent_entity = mocker.patch(
"core.app.app_config.easy_ui_based_app.agent.manager.AgentEntity",
@@ -128,7 +129,7 @@ class TestAgentConfigManagerConvert:
mock_validate.assert_called_once()
mock_agent_entity.assert_called_once()
def test_convert_tool_requires_minimum_keys(self, mocker, base_config):
def test_convert_tool_requires_minimum_keys(self, mocker: MockerFixture, base_config):
mock_validate = mocker.patch(
"core.app.app_config.easy_ui_based_app.agent.manager.AgentToolEntity.model_validate",
return_value=MagicMock(),
@@ -2,6 +2,7 @@ import uuid
from unittest.mock import MagicMock
import pytest
from pytest_mock import MockerFixture
from core.app.app_config.easy_ui_based_app.dataset.manager import DatasetConfigManager
from core.entities.agent_entities import PlanningStrategy
@@ -69,7 +70,7 @@ class TestDatasetConfigManagerConvert:
assert result.dataset_ids == [valid_uuid]
assert result.retrieve_config.query_variable == "query"
def test_convert_single_with_metadata_configs(self, valid_uuid, mocker):
def test_convert_single_with_metadata_configs(self, valid_uuid, mocker: MockerFixture):
mock_retrieve_config = MagicMock()
mock_entity = MagicMock()
mock_entity.dataset_ids = [valid_uuid]
@@ -258,7 +259,7 @@ class TestExtractDatasetConfig:
with pytest.raises(ValueError):
DatasetConfigManager.extract_dataset_config_for_legacy_compatibility("tenant1", AppMode.CHAT, config)
def test_extract_invalid_uuid(self, mocker):
def test_extract_invalid_uuid(self, mocker: MockerFixture):
invalid_uuid = "not-a-uuid"
config = {
"agent_mode": {
@@ -270,7 +271,7 @@ class TestExtractDatasetConfig:
with pytest.raises(ValueError):
DatasetConfigManager.extract_dataset_config_for_legacy_compatibility("tenant1", AppMode.CHAT, config)
def test_extract_dataset_not_exists(self, valid_uuid, mocker):
def test_extract_dataset_not_exists(self, valid_uuid, mocker: MockerFixture):
mocker.patch(
"core.app.app_config.easy_ui_based_app.dataset.manager.DatasetService.get_dataset",
return_value=None,
@@ -292,7 +293,7 @@ class TestExtractDatasetConfig:
class TestIsDatasetExists:
def test_dataset_exists_true(self, mocker, valid_uuid):
def test_dataset_exists_true(self, mocker: MockerFixture, valid_uuid):
mock_dataset = MagicMock()
mock_dataset.tenant_id = "tenant1"
mocker.patch(
@@ -302,14 +303,14 @@ class TestIsDatasetExists:
assert DatasetConfigManager.is_dataset_exists("tenant1", valid_uuid)
def test_dataset_exists_false_when_not_found(self, mocker, valid_uuid):
def test_dataset_exists_false_when_not_found(self, mocker: MockerFixture, valid_uuid):
mocker.patch(
"core.app.app_config.easy_ui_based_app.dataset.manager.DatasetService.get_dataset",
return_value=None,
)
assert not DatasetConfigManager.is_dataset_exists("tenant1", valid_uuid)
def test_dataset_exists_false_when_tenant_mismatch(self, mocker, valid_uuid):
def test_dataset_exists_false_when_tenant_mismatch(self, mocker: MockerFixture, valid_uuid):
mock_dataset = MagicMock()
mock_dataset.tenant_id = "other"
mocker.patch(
@@ -2,6 +2,7 @@ from types import SimpleNamespace
from unittest.mock import MagicMock
import pytest
from pytest_mock import MockerFixture
from core.app.app_config.easy_ui_based_app.model_config.converter import ModelConfigConverter
from core.entities.model_entities import ModelStatus
@@ -16,7 +17,7 @@ from graphon.model_runtime.entities.model_entities import ModelPropertyKey
class TestModelConfigConverter:
@pytest.fixture(autouse=True)
def patch_response_entity(self, mocker):
def patch_response_entity(self, mocker: MockerFixture):
"""
Patch ModelConfigWithCredentialsEntity to bypass Pydantic validation
and return a simple namespace object instead.
@@ -69,7 +70,7 @@ class TestModelConfigConverter:
return bundle
@pytest.fixture
def patch_provider_manager(self, mocker, mock_provider_bundle):
def patch_provider_manager(self, mocker: MockerFixture, mock_provider_bundle):
mock_manager = MagicMock()
mock_manager.get_provider_model_bundle.return_value = mock_provider_bundle
mocker.patch(
@@ -99,7 +100,7 @@ class TestModelConfigConverter:
assert result.parameters == {"temperature": 0.7}
assert result.stop == ["\n"]
def test_convert_mode_from_schema_valid(self, mock_app_config, mock_provider_bundle, mocker):
def test_convert_mode_from_schema_valid(self, mock_app_config, mock_provider_bundle, mocker: MockerFixture):
mock_app_config.model.mode = None
mock_provider_bundle.model_type_instance.get_model_schema.return_value.model_properties = {
@@ -116,7 +117,9 @@ class TestModelConfigConverter:
result = ModelConfigConverter.convert(mock_app_config)
assert result.mode == LLMMode.COMPLETION
def test_convert_mode_from_schema_invalid_fallback(self, mock_app_config, mock_provider_bundle, mocker):
def test_convert_mode_from_schema_invalid_fallback(
self, mock_app_config, mock_provider_bundle, mocker: MockerFixture
):
mock_provider_bundle.model_type_instance.get_model_schema.return_value.model_properties = {
ModelPropertyKey.MODE: "invalid"
}
@@ -135,7 +138,7 @@ class TestModelConfigConverter:
# Credential Errors
# =============================
def test_convert_credentials_none_raises(self, mock_app_config, mock_provider_bundle, mocker):
def test_convert_credentials_none_raises(self, mock_app_config, mock_provider_bundle, mocker: MockerFixture):
mock_provider_bundle.configuration.get_current_credentials.return_value = None
mock_manager = MagicMock()
@@ -152,7 +155,7 @@ class TestModelConfigConverter:
# Provider Model Errors
# =============================
def test_convert_provider_model_none_raises(self, mock_app_config, mock_provider_bundle, mocker):
def test_convert_provider_model_none_raises(self, mock_app_config, mock_provider_bundle, mocker: MockerFixture):
mock_provider_bundle.configuration.get_provider_model.return_value = None
mock_manager = MagicMock()
@@ -174,7 +177,7 @@ class TestModelConfigConverter:
],
)
def test_convert_provider_model_status_errors(
self, mock_app_config, mock_provider_bundle, mocker, status, expected_exception
self, mock_app_config, mock_provider_bundle, mocker: MockerFixture, status, expected_exception
):
mock_provider = MagicMock()
mock_provider.status = status
@@ -194,7 +197,7 @@ class TestModelConfigConverter:
# Schema Errors
# =============================
def test_convert_model_schema_none_raises(self, mock_app_config, mock_provider_bundle, mocker):
def test_convert_model_schema_none_raises(self, mock_app_config, mock_provider_bundle, mocker: MockerFixture):
mock_provider_bundle.model_type_instance.get_model_schema.return_value = None
mock_manager = MagicMock()
@@ -1,6 +1,7 @@
from unittest.mock import MagicMock
import pytest
from pytest_mock import MockerFixture
# Target
from core.app.app_config.easy_ui_based_app.model_config.manager import ModelConfigManager
@@ -107,7 +108,9 @@ class TestModelConfigManager:
# validate_and_set_defaults
# ==========================================================
def test_validate_and_set_defaults_success(self, mocker, valid_config, provider_entities, valid_model_list):
def test_validate_and_set_defaults_success(
self, mocker: MockerFixture, valid_config, provider_entities, valid_model_list
):
self._patch_model_assembly(
mocker,
provider_entities=provider_entities,
@@ -127,35 +130,37 @@ class TestModelConfigManager:
with pytest.raises(ValueError, match="object type"):
ModelConfigManager.validate_and_set_defaults("tenant1", {"model": "invalid"})
def test_validate_and_set_defaults_missing_provider(self, mocker, provider_entities):
def test_validate_and_set_defaults_missing_provider(self, mocker: MockerFixture, provider_entities):
config = {"model": {"name": "gpt-4", "completion_params": {}}}
self._patch_model_assembly(mocker, provider_entities=provider_entities, model_list=[])
with pytest.raises(ValueError, match="model.provider is required"):
ModelConfigManager.validate_and_set_defaults("tenant1", config)
def test_validate_and_set_defaults_invalid_provider(self, mocker, provider_entities):
def test_validate_and_set_defaults_invalid_provider(self, mocker: MockerFixture, provider_entities):
config = {"model": {"provider": "invalid/provider", "name": "gpt-4", "completion_params": {}}}
self._patch_model_assembly(mocker, provider_entities=provider_entities, model_list=[])
with pytest.raises(ValueError, match="model.provider is required"):
ModelConfigManager.validate_and_set_defaults("tenant1", config)
def test_validate_and_set_defaults_missing_name(self, mocker, provider_entities):
def test_validate_and_set_defaults_missing_name(self, mocker: MockerFixture, provider_entities):
config = {"model": {"provider": "openai/gpt", "completion_params": {}}}
self._patch_model_assembly(mocker, provider_entities=provider_entities, model_list=[])
with pytest.raises(ValueError, match="model.name is required"):
ModelConfigManager.validate_and_set_defaults("tenant1", config)
def test_validate_and_set_defaults_empty_models(self, mocker, provider_entities):
def test_validate_and_set_defaults_empty_models(self, mocker: MockerFixture, provider_entities):
config = {"model": {"provider": "openai/gpt", "name": "gpt-4", "completion_params": {}}}
self._patch_model_assembly(mocker, provider_entities=provider_entities, model_list=[])
with pytest.raises(ValueError, match="must be in the specified model list"):
ModelConfigManager.validate_and_set_defaults("tenant1", config)
def test_validate_and_set_defaults_invalid_model_name(self, mocker, provider_entities, valid_model_list):
def test_validate_and_set_defaults_invalid_model_name(
self, mocker: MockerFixture, provider_entities, valid_model_list
):
config = {"model": {"provider": "openai/gpt", "name": "invalid", "completion_params": {}}}
self._patch_model_assembly(
mocker,
@@ -166,7 +171,7 @@ class TestModelConfigManager:
with pytest.raises(ValueError, match="must be in the specified model list"):
ModelConfigManager.validate_and_set_defaults("tenant1", config)
def test_validate_and_set_defaults_default_mode_when_missing(self, mocker, provider_entities):
def test_validate_and_set_defaults_default_mode_when_missing(self, mocker: MockerFixture, provider_entities):
model = MagicMock()
model.model = "gpt-4"
model.model_properties = {}
@@ -178,7 +183,9 @@ class TestModelConfigManager:
assert updated_config["model"]["mode"] == "completion"
def test_validate_and_set_defaults_missing_completion_params(self, mocker, provider_entities, valid_model_list):
def test_validate_and_set_defaults_missing_completion_params(
self, mocker: MockerFixture, provider_entities, valid_model_list
):
config = {"model": {"provider": "openai/gpt", "name": "gpt-4"}}
self._patch_model_assembly(
mocker,
@@ -189,7 +196,7 @@ class TestModelConfigManager:
with pytest.raises(ValueError, match="completion_params is required"):
ModelConfigManager.validate_and_set_defaults("tenant1", config)
def test_validate_and_set_defaults_provider_without_slash_converted(self, mocker, valid_model_list):
def test_validate_and_set_defaults_provider_without_slash_converted(self, mocker: MockerFixture, valid_model_list):
"""
Covers branch where provider does not contain '/' and
ModelProviderID conversion is triggered (line 64).
@@ -1,6 +1,7 @@
from unittest.mock import MagicMock
import pytest
from pytest_mock import MockerFixture
from core.app.app_config.easy_ui_based_app.prompt_template.manager import (
PromptTemplateConfigManager,
@@ -38,7 +39,7 @@ class TestPromptTemplateConfigManagerConvert:
with pytest.raises(ValueError, match="prompt_type is required"):
PromptTemplateConfigManager.convert({})
def test_convert_simple_prompt(self, mocker):
def test_convert_simple_prompt(self, mocker: MockerFixture):
mock_prompt_entity_cls = MagicMock()
mock_prompt_entity_cls.PromptType = DummyPromptType()
@@ -56,7 +57,7 @@ class TestPromptTemplateConfigManagerConvert:
assert result == "simple_entity"
mock_prompt_entity_cls.assert_called_once_with(prompt_type="simple", simple_prompt_template="hello")
def test_convert_advanced_chat_valid(self, mocker):
def test_convert_advanced_chat_valid(self, mocker: MockerFixture):
mock_prompt_entity_cls = MagicMock()
mock_prompt_entity_cls.PromptType = DummyPromptType()
mock_prompt_entity_cls.return_value = "advanced_entity"
@@ -97,7 +98,7 @@ class TestPromptTemplateConfigManagerConvert:
{"text": "hi", "role": 123},
],
)
def test_convert_advanced_invalid_message_fields(self, mocker, message):
def test_convert_advanced_invalid_message_fields(self, mocker: MockerFixture, message):
mock_prompt_entity_cls = MagicMock()
mock_prompt_entity_cls.PromptType = DummyPromptType()
@@ -114,7 +115,7 @@ class TestPromptTemplateConfigManagerConvert:
with pytest.raises(ValueError):
PromptTemplateConfigManager.convert(config)
def test_convert_advanced_completion_with_roles(self, mocker):
def test_convert_advanced_completion_with_roles(self, mocker: MockerFixture):
mock_prompt_entity_cls = MagicMock()
mock_prompt_entity_cls.PromptType = DummyPromptType()
mock_prompt_entity_cls.return_value = "advanced_entity"
@@ -154,7 +155,7 @@ class TestValidateAndSetDefaults:
def setup_method(self):
self.valid_model = {"mode": "chat"}
def _patch_prompt_type(self, mocker):
def _patch_prompt_type(self, mocker: MockerFixture):
mock_prompt_entity_cls = MagicMock()
mock_prompt_entity_cls.PromptType = DummyPromptType()
mocker.patch(
@@ -163,7 +164,7 @@ class TestValidateAndSetDefaults:
)
return mock_prompt_entity_cls
def test_default_prompt_type_set(self, mocker):
def test_default_prompt_type_set(self, mocker: MockerFixture):
self._patch_prompt_type(mocker)
config = {"model": self.valid_model}
@@ -173,7 +174,7 @@ class TestValidateAndSetDefaults:
assert result["prompt_type"] == "simple"
assert isinstance(keys, list)
def test_invalid_prompt_type_raises(self, mocker):
def test_invalid_prompt_type_raises(self, mocker: MockerFixture):
class InvalidEnum(DummyPromptType):
def __iter__(self):
return iter([DummyEnumValue("valid")])
@@ -191,7 +192,7 @@ class TestValidateAndSetDefaults:
with pytest.raises(ValueError):
PromptTemplateConfigManager.validate_and_set_defaults("chat_app", config)
def test_invalid_chat_prompt_config_type(self, mocker):
def test_invalid_chat_prompt_config_type(self, mocker: MockerFixture):
self._patch_prompt_type(mocker)
config = {
@@ -203,7 +204,7 @@ class TestValidateAndSetDefaults:
with pytest.raises(ValueError):
PromptTemplateConfigManager.validate_and_set_defaults("chat_app", config)
def test_simple_mode_invalid_pre_prompt_type(self, mocker):
def test_simple_mode_invalid_pre_prompt_type(self, mocker: MockerFixture):
self._patch_prompt_type(mocker)
config = {
@@ -215,7 +216,7 @@ class TestValidateAndSetDefaults:
with pytest.raises(ValueError):
PromptTemplateConfigManager.validate_and_set_defaults("chat_app", config)
def test_advanced_requires_one_config(self, mocker):
def test_advanced_requires_one_config(self, mocker: MockerFixture):
self._patch_prompt_type(mocker)
config = {
@@ -228,7 +229,7 @@ class TestValidateAndSetDefaults:
with pytest.raises(ValueError):
PromptTemplateConfigManager.validate_and_set_defaults("chat_app", config)
def test_advanced_invalid_model_mode(self, mocker):
def test_advanced_invalid_model_mode(self, mocker: MockerFixture):
self._patch_prompt_type(mocker)
config = {
@@ -240,7 +241,7 @@ class TestValidateAndSetDefaults:
with pytest.raises(ValueError):
PromptTemplateConfigManager.validate_and_set_defaults("chat_app", config)
def test_advanced_chat_prompt_length_exceeds(self, mocker):
def test_advanced_chat_prompt_length_exceeds(self, mocker: MockerFixture):
self._patch_prompt_type(mocker)
config = {
@@ -252,7 +253,7 @@ class TestValidateAndSetDefaults:
with pytest.raises(ValueError):
PromptTemplateConfigManager.validate_and_set_defaults("chat_app", config)
def test_completion_prefix_defaults_set_when_empty(self, mocker):
def test_completion_prefix_defaults_set_when_empty(self, mocker: MockerFixture):
self._patch_prompt_type(mocker)
config = {
@@ -1,4 +1,5 @@
import pytest
from pytest_mock import MockerFixture
from core.app.app_config.easy_ui_based_app.variables.manager import (
BasicVariablesConfigManager,
@@ -15,7 +16,7 @@ class TestBasicVariablesConfigManagerConvert:
assert variables == []
assert external == []
def test_convert_external_data_tools_enabled_and_disabled(self, mocker):
def test_convert_external_data_tools_enabled_and_disabled(self, mocker: MockerFixture):
config = {
"external_data_tools": [
{"enabled": False},
@@ -232,7 +233,7 @@ class TestValidateExternalDataToolsAndSetDefaults:
with pytest.raises(ValueError):
BasicVariablesConfigManager.validate_external_data_tools_and_set_defaults("tenant", config)
def test_validate_disabled_tool_skipped(self, mocker):
def test_validate_disabled_tool_skipped(self, mocker: MockerFixture):
config = {"external_data_tools": [{"enabled": False}]}
spy = mocker.patch(
@@ -250,7 +251,7 @@ class TestValidateExternalDataToolsAndSetDefaults:
with pytest.raises(ValueError):
BasicVariablesConfigManager.validate_external_data_tools_and_set_defaults("tenant", config)
def test_validate_enabled_tool_calls_factory(self, mocker):
def test_validate_enabled_tool_calls_factory(self, mocker: MockerFixture):
config = {"external_data_tools": [{"enabled": True, "type": "tool", "config": {"a": 1}}]}
spy = mocker.patch(
@@ -263,7 +264,7 @@ class TestValidateExternalDataToolsAndSetDefaults:
class TestValidateAndSetDefaultsIntegration:
def test_validate_and_set_defaults_calls_both(self, mocker):
def test_validate_and_set_defaults_calls_both(self, mocker: MockerFixture):
config = {}
spy_var = mocker.patch.object(
@@ -2,6 +2,7 @@ from collections import UserDict
from unittest.mock import MagicMock
import pytest
from pytest_mock import MockerFixture
from core.app.app_config.base_app_config_manager import BaseAppConfigManager
@@ -12,7 +13,7 @@ class TestBaseAppConfigManager:
return {"key": "value", "another": 123}
@pytest.fixture
def mock_app_additional_features(self, mocker):
def mock_app_additional_features(self, mocker: MockerFixture):
mock_instance = MagicMock()
mocker.patch(
"core.app.app_config.base_app_config_manager.AppAdditionalFeatures",
@@ -21,7 +22,7 @@ class TestBaseAppConfigManager:
return mock_instance
@pytest.fixture
def mock_managers(self, mocker):
def mock_managers(self, mocker: MockerFixture):
retrieval = mocker.patch(
"core.app.app_config.base_app_config_manager.RetrievalResourceConfigManager.convert",
return_value="retrieval_result",
@@ -72,7 +73,7 @@ class TestBaseAppConfigManager:
)
def test_convert_features_all_modes(
self,
mocker,
mocker: MockerFixture,
mock_config_dict,
mock_app_additional_features,
mock_managers,
@@ -107,7 +108,7 @@ class TestBaseAppConfigManager:
mock_managers["speech_to_text"].assert_called_once_with(config=dict(mock_config_dict.items()))
mock_managers["text_to_speech"].assert_called_once_with(config=dict(mock_config_dict.items()))
def test_convert_features_empty_config(self, mocker, mock_app_additional_features, mock_managers):
def test_convert_features_empty_config(self, mocker: MockerFixture, mock_app_additional_features, mock_managers):
# Arrange
empty_config = {}
mock_app_mode = MagicMock()
@@ -143,7 +144,7 @@ class TestBaseAppConfigManager:
with pytest.raises((TypeError, AttributeError)):
BaseAppConfigManager.convert_features(invalid_config, "CHAT")
def test_convert_features_manager_exception_propagates(self, mocker, mock_config_dict):
def test_convert_features_manager_exception_propagates(self, mocker: MockerFixture, mock_config_dict):
# Arrange
mocker.patch(
"core.app.app_config.base_app_config_manager.RetrievalResourceConfigManager.convert",
@@ -154,7 +155,9 @@ class TestBaseAppConfigManager:
with pytest.raises(RuntimeError):
BaseAppConfigManager.convert_features(mock_config_dict, "CHAT")
def test_convert_features_mapping_subclass(self, mocker, mock_app_additional_features, mock_managers):
def test_convert_features_mapping_subclass(
self, mocker: MockerFixture, mock_app_additional_features, mock_managers
):
# Arrange
class CustomMapping(UserDict):
pass
@@ -1,4 +1,5 @@
import pytest
from pytest_mock import MockerFixture
from core.app.app_config.workflow_ui_based_app.variables.manager import (
WorkflowVariablesConfigManager,
@@ -10,19 +11,19 @@ from core.app.app_config.workflow_ui_based_app.variables.manager import (
@pytest.fixture
def mock_workflow(mocker):
def mock_workflow(mocker: MockerFixture):
workflow = mocker.MagicMock()
workflow.graph_dict = {"nodes": []}
return workflow
@pytest.fixture
def mock_variable_entity(mocker):
def mock_variable_entity(mocker: MockerFixture):
return mocker.patch("core.app.app_config.workflow_ui_based_app.variables.manager.VariableEntity")
@pytest.fixture
def mock_rag_entity(mocker):
def mock_rag_entity(mocker: MockerFixture):
return mocker.patch("core.app.app_config.workflow_ui_based_app.variables.manager.RagPipelineVariableEntity")
@@ -111,7 +111,7 @@ class TestAdvancedChatAppGeneratorInternals:
workflow_id="workflow-id",
)
def test_generate_loads_conversation_and_files(self, monkeypatch):
def test_generate_loads_conversation_and_files(self, monkeypatch: pytest.MonkeyPatch):
generator = AdvancedChatAppGenerator()
app_config = self._build_app_config()
@@ -195,7 +195,7 @@ class TestAdvancedChatAppGeneratorInternals:
assert captured["application_generate_entity"].files == built_files
assert build_files_called["called"] is True
def test_resume_delegates_to_generate(self, monkeypatch):
def test_resume_delegates_to_generate(self, monkeypatch: pytest.MonkeyPatch):
generator = AdvancedChatAppGenerator()
application_generate_entity = AdvancedChatAppGenerateEntity.model_construct(
task_id="task",
@@ -235,7 +235,7 @@ class TestAdvancedChatAppGeneratorInternals:
assert result == {"resumed": True}
assert captured["graph_runtime_state"] is not None
def test_single_iteration_generate_builds_debug_task(self, monkeypatch):
def test_single_iteration_generate_builds_debug_task(self, monkeypatch: pytest.MonkeyPatch):
generator = AdvancedChatAppGenerator()
app_config = self._build_app_config()
captured: dict[str, object] = {}
@@ -293,7 +293,7 @@ class TestAdvancedChatAppGeneratorInternals:
assert captured["variable_loader"] is var_loader
assert captured["application_generate_entity"].single_iteration_run.node_id == "node-1"
def test_single_loop_generate_builds_debug_task(self, monkeypatch):
def test_single_loop_generate_builds_debug_task(self, monkeypatch: pytest.MonkeyPatch):
generator = AdvancedChatAppGenerator()
app_config = self._build_app_config()
captured: dict[str, object] = {}
@@ -351,7 +351,7 @@ class TestAdvancedChatAppGeneratorInternals:
assert captured["variable_loader"] is var_loader
assert captured["application_generate_entity"].single_loop_run.node_id == "node-2"
def test_generate_internal_flow_initial_conversation_with_pause_layer(self, monkeypatch):
def test_generate_internal_flow_initial_conversation_with_pause_layer(self, monkeypatch: pytest.MonkeyPatch):
generator = AdvancedChatAppGenerator()
generator._dialogue_count = 0
app_config = self._build_app_config()
@@ -449,7 +449,7 @@ class TestAdvancedChatAppGeneratorInternals:
assert isinstance(captured["conversation"], ConversationSnapshot)
assert isinstance(captured["message"], MessageSnapshot)
def test_generate_internal_flow_with_existing_records_skips_init(self, monkeypatch):
def test_generate_internal_flow_with_existing_records_skips_init(self, monkeypatch: pytest.MonkeyPatch):
generator = AdvancedChatAppGenerator()
generator._dialogue_count = 0
app_config = self._build_app_config()
@@ -535,7 +535,7 @@ class TestAdvancedChatAppGeneratorInternals:
db_session.refresh.assert_not_called()
db_session.close.assert_called_once()
def test_generate_worker_raises_when_workflow_not_found(self, monkeypatch):
def test_generate_worker_raises_when_workflow_not_found(self, monkeypatch: pytest.MonkeyPatch):
generator = AdvancedChatAppGenerator()
generator._dialogue_count = 1
app_config = self._build_app_config()
@@ -594,7 +594,7 @@ class TestAdvancedChatAppGeneratorInternals:
graph_runtime_state=None,
)
def test_generate_worker_raises_when_app_not_found_for_internal_call(self, monkeypatch):
def test_generate_worker_raises_when_app_not_found_for_internal_call(self, monkeypatch: pytest.MonkeyPatch):
generator = AdvancedChatAppGenerator()
generator._dialogue_count = 1
app_config = self._build_app_config()
@@ -658,7 +658,7 @@ class TestAdvancedChatAppGeneratorInternals:
graph_runtime_state=None,
)
def test_generate_worker_handles_stopped_error(self, monkeypatch):
def test_generate_worker_handles_stopped_error(self, monkeypatch: pytest.MonkeyPatch):
generator = AdvancedChatAppGenerator()
generator._dialogue_count = 1
app_config = self._build_app_config()
@@ -732,7 +732,7 @@ class TestAdvancedChatAppGeneratorInternals:
queue_manager.publish_error.assert_not_called()
def test_generate_worker_handles_validation_error(self, monkeypatch):
def test_generate_worker_handles_validation_error(self, monkeypatch: pytest.MonkeyPatch):
generator = AdvancedChatAppGenerator()
generator._dialogue_count = 1
app_config = self._build_app_config()
@@ -816,7 +816,7 @@ class TestAdvancedChatAppGeneratorInternals:
queue_manager.publish_error.assert_called_once()
def test_generate_worker_handles_value_and_unknown_errors(self, monkeypatch):
def test_generate_worker_handles_value_and_unknown_errors(self, monkeypatch: pytest.MonkeyPatch):
app_config = self._build_app_config()
@contextmanager
@@ -897,7 +897,7 @@ class TestAdvancedChatAppGeneratorInternals:
queue_manager.publish_error.assert_called_once()
def test_handle_response_closed_file_raises_stopped(self, monkeypatch):
def test_handle_response_closed_file_raises_stopped(self, monkeypatch: pytest.MonkeyPatch):
generator = AdvancedChatAppGenerator()
generator._dialogue_count = 1
@@ -953,7 +953,7 @@ class TestAdvancedChatAppGeneratorInternals:
stream=False,
)
def test_handle_response_re_raises_value_error(self, monkeypatch):
def test_handle_response_re_raises_value_error(self, monkeypatch: pytest.MonkeyPatch):
generator = AdvancedChatAppGenerator()
generator._dialogue_count = 1
app_config = self._build_app_config()
@@ -1002,7 +1002,7 @@ class TestAdvancedChatAppGeneratorInternals:
logger_exception.assert_called_once()
def test_generate_worker_handles_invoke_auth_error(self, monkeypatch):
def test_generate_worker_handles_invoke_auth_error(self, monkeypatch: pytest.MonkeyPatch):
generator = AdvancedChatAppGenerator()
generator._dialogue_count = 1
@@ -1088,7 +1088,7 @@ class TestAdvancedChatAppGeneratorInternals:
assert queue_manager.publish_error.called
def test_generate_debugger_enables_retrieve_source(self, monkeypatch):
def test_generate_debugger_enables_retrieve_source(self, monkeypatch: pytest.MonkeyPatch):
generator = AdvancedChatAppGenerator()
app_config = WorkflowUIBasedAppConfig(
@@ -1167,7 +1167,7 @@ class TestAdvancedChatAppGeneratorInternals:
assert app_config.additional_features.show_retrieve_source is True
assert captured["application_generate_entity"].query == "hello"
def test_generate_service_api_sets_parent_message_id(self, monkeypatch):
def test_generate_service_api_sets_parent_message_id(self, monkeypatch: pytest.MonkeyPatch):
generator = AdvancedChatAppGenerator()
app_config = WorkflowUIBasedAppConfig(
@@ -224,7 +224,7 @@ class TestAdvancedChatGenerateTaskPipeline:
assert isinstance(responses[0], ValueError)
def test_handle_workflow_started_event_sets_run_id(self, monkeypatch):
def test_handle_workflow_started_event_sets_run_id(self, monkeypatch: pytest.MonkeyPatch):
pipeline = _make_pipeline()
pipeline._graph_runtime_state = GraphRuntimeState(
variable_pool=build_test_variable_pool(variables=build_system_variables(workflow_execution_id="run-id")),
@@ -368,7 +368,7 @@ class TestAdvancedChatGenerateTaskPipeline:
assert list(pipeline._handle_loop_next_event(loop_next)) == ["loop_next"]
assert list(pipeline._handle_loop_completed_event(loop_done)) == ["loop_done"]
def test_workflow_finish_handlers(self, monkeypatch):
def test_workflow_finish_handlers(self, monkeypatch: pytest.MonkeyPatch):
pipeline = _make_pipeline()
pipeline._workflow_run_id = "run-id"
pipeline._graph_runtime_state = GraphRuntimeState(
@@ -593,7 +593,7 @@ class TestAdvancedChatGenerateTaskPipeline:
assert message.answer == "hello"
assert message.message_metadata
def test_handle_stop_event_saves_message_for_moderation(self, monkeypatch):
def test_handle_stop_event_saves_message_for_moderation(self, monkeypatch: pytest.MonkeyPatch):
pipeline = _make_pipeline()
pipeline._message_end_to_stream_response = lambda: "end"
saved: list[str] = []
@@ -614,7 +614,7 @@ class TestAdvancedChatGenerateTaskPipeline:
assert responses == ["end"]
assert saved == ["saved"]
def test_handle_message_end_event_applies_output_moderation(self, monkeypatch):
def test_handle_message_end_event_applies_output_moderation(self, monkeypatch: pytest.MonkeyPatch):
pipeline = _make_pipeline()
pipeline._graph_runtime_state = GraphRuntimeState(
variable_pool=VariablePool(system_variables=build_system_variables(workflow_execution_id="run-id")),
@@ -2,6 +2,7 @@ import uuid
from types import SimpleNamespace
import pytest
from pytest_mock import MockerFixture
from core.app.app_config.entities import EasyUIBasedAppModelConfigFrom
from core.app.apps.agent_chat.app_config_manager import (
@@ -11,7 +12,7 @@ from core.entities.agent_entities import PlanningStrategy
class TestAgentChatAppConfigManagerGetAppConfig:
def test_get_app_config_override_config(self, mocker):
def test_get_app_config_override_config(self, mocker: MockerFixture):
app_model = mocker.MagicMock(id="app1", tenant_id="tenant", mode="agent-chat")
app_model_config = mocker.MagicMock(id="cfg1")
app_model_config.to_dict.return_value = {"ignored": True}
@@ -45,7 +46,7 @@ class TestAgentChatAppConfigManagerGetAppConfig:
assert result.variables == "variables"
assert result.external_data_variables == "external"
def test_get_app_config_conversation_specific(self, mocker):
def test_get_app_config_conversation_specific(self, mocker: MockerFixture):
app_model = mocker.MagicMock(id="app1", tenant_id="tenant", mode="agent-chat")
app_model_config = mocker.MagicMock(id="cfg1")
app_model_config.to_dict.return_value = {"model": {"provider": "p"}}
@@ -76,7 +77,7 @@ class TestAgentChatAppConfigManagerGetAppConfig:
assert result.app_model_config_dict == app_model_config.to_dict.return_value
assert result.app_model_config_from.value == "conversation-specific-config"
def test_get_app_config_latest_config(self, mocker):
def test_get_app_config_latest_config(self, mocker: MockerFixture):
app_model = mocker.MagicMock(id="app1", tenant_id="tenant", mode="agent-chat")
app_model_config = mocker.MagicMock(id="cfg1")
app_model_config.to_dict.return_value = {"model": {"provider": "p"}}
@@ -107,7 +108,7 @@ class TestAgentChatAppConfigManagerGetAppConfig:
class TestAgentChatAppConfigManagerConfigValidate:
def test_config_validate_filters_related_keys(self, mocker):
def test_config_validate_filters_related_keys(self, mocker: MockerFixture):
config = {
"model": {},
"user_input_form": {},
@@ -247,7 +248,7 @@ class TestValidateAgentModeAndSetDefaults:
{"agent_mode": {"enabled": True, "tools": [{"dataset": {"enabled": True, "id": "bad"}}]}},
)
def test_old_tool_dataset_id_not_exists(self, mocker):
def test_old_tool_dataset_id_not_exists(self, mocker: MockerFixture):
mocker.patch(
"core.app.apps.agent_chat.app_config_manager.DatasetConfigManager.is_dataset_exists",
return_value=False,
@@ -275,7 +276,7 @@ class TestValidateAgentModeAndSetDefaults:
"tenant", {"agent_mode": {"enabled": True, "tools": [tool]}}
)
def test_valid_old_and_new_style_tools(self, mocker):
def test_valid_old_and_new_style_tools(self, mocker: MockerFixture):
mocker.patch(
"core.app.apps.agent_chat.app_config_manager.DatasetConfigManager.is_dataset_exists",
return_value=True,
@@ -2,6 +2,7 @@ import contextlib
import pytest
from pydantic import ValidationError
from pytest_mock import MockerFixture
from core.app.apps.agent_chat.app_generator import AgentChatAppGenerator
from core.app.apps.exc import GenerateTaskStoppedError
@@ -16,7 +17,7 @@ class DummyAccount:
@pytest.fixture
def generator(mocker):
def generator(mocker: MockerFixture):
gen = AgentChatAppGenerator()
mocker.patch(
"core.app.apps.agent_chat.app_generator.current_app",
@@ -27,19 +28,19 @@ def generator(mocker):
class TestAgentChatAppGeneratorGenerate:
def test_generate_rejects_blocking_mode(self, generator, mocker):
def test_generate_rejects_blocking_mode(self, generator, mocker: MockerFixture):
app_model = mocker.MagicMock()
user = DummyAccount("user")
with pytest.raises(ValueError):
generator.generate(app_model=app_model, user=user, args={}, invoke_from=mocker.MagicMock(), streaming=False)
def test_generate_requires_query(self, generator, mocker):
def test_generate_requires_query(self, generator, mocker: MockerFixture):
app_model = mocker.MagicMock()
user = DummyAccount("user")
with pytest.raises(ValueError):
generator.generate(app_model=app_model, user=user, args={"inputs": {}}, invoke_from=mocker.MagicMock())
def test_generate_rejects_non_string_query(self, generator, mocker):
def test_generate_rejects_non_string_query(self, generator, mocker: MockerFixture):
app_model = mocker.MagicMock()
user = DummyAccount("user")
with pytest.raises(ValueError):
@@ -50,7 +51,7 @@ class TestAgentChatAppGeneratorGenerate:
invoke_from=mocker.MagicMock(),
)
def test_generate_override_requires_debugger(self, generator, mocker):
def test_generate_override_requires_debugger(self, generator, mocker: MockerFixture):
app_model = mocker.MagicMock()
user = DummyAccount("user")
@@ -62,7 +63,7 @@ class TestAgentChatAppGeneratorGenerate:
invoke_from=InvokeFrom.WEB_APP,
)
def test_generate_success_with_debugger_override(self, generator, mocker):
def test_generate_success_with_debugger_override(self, generator, mocker: MockerFixture):
app_model = mocker.MagicMock(id="app1", tenant_id="tenant", mode="agent-chat")
app_model_config = mocker.MagicMock(id="cfg1")
app_model_config.to_dict.return_value = {"model": {"provider": "p"}}
@@ -142,7 +143,7 @@ class TestAgentChatAppGeneratorGenerate:
assert result == {"result": "ok"}
thread_obj.start.assert_called_once()
def test_generate_without_file_config(self, generator, mocker):
def test_generate_without_file_config(self, generator, mocker: MockerFixture):
app_model = mocker.MagicMock(id="app1", tenant_id="tenant", mode="agent-chat")
app_model_config = mocker.MagicMock(id="cfg1")
app_model_config.to_dict.return_value = {"model": {"provider": "p"}}
@@ -213,14 +214,14 @@ class TestAgentChatAppGeneratorGenerate:
class TestAgentChatAppGeneratorWorker:
@pytest.fixture(autouse=True)
def patch_context(self, mocker):
def patch_context(self, mocker: MockerFixture):
@contextlib.contextmanager
def ctx_manager(*args, **kwargs):
yield
mocker.patch("core.app.apps.agent_chat.app_generator.preserve_flask_contexts", ctx_manager)
def test_generate_worker_handles_generate_task_stopped(self, generator, mocker):
def test_generate_worker_handles_generate_task_stopped(self, generator, mocker: MockerFixture):
queue_manager = mocker.MagicMock()
generator._get_conversation = mocker.MagicMock(return_value=mocker.MagicMock())
generator._get_message = mocker.MagicMock(return_value=mocker.MagicMock())
@@ -250,7 +251,7 @@ class TestAgentChatAppGeneratorWorker:
Exception("bad"),
],
)
def test_generate_worker_publishes_errors(self, generator, mocker, error):
def test_generate_worker_publishes_errors(self, generator, mocker: MockerFixture, error):
queue_manager = mocker.MagicMock()
generator._get_conversation = mocker.MagicMock(return_value=mocker.MagicMock())
generator._get_message = mocker.MagicMock(return_value=mocker.MagicMock())
@@ -271,7 +272,7 @@ class TestAgentChatAppGeneratorWorker:
assert queue_manager.publish_error.called
def test_generate_worker_logs_value_error_when_debug(self, generator, mocker):
def test_generate_worker_logs_value_error_when_debug(self, generator, mocker: MockerFixture):
queue_manager = mocker.MagicMock()
generator._get_conversation = mocker.MagicMock(return_value=mocker.MagicMock())
generator._get_message = mocker.MagicMock(return_value=mocker.MagicMock())
@@ -1,4 +1,5 @@
import pytest
from pytest_mock import MockerFixture
from core.agent.entities import AgentEntity
from core.app.apps.agent_chat.app_runner import AgentChatAppRunner
@@ -13,7 +14,7 @@ def runner():
class TestAgentChatAppRunnerRun:
def test_run_app_not_found(self, runner, mocker):
def test_run_app_not_found(self, runner, mocker: MockerFixture):
app_config = mocker.MagicMock(app_id="app1", tenant_id="tenant", agent=mocker.MagicMock())
generate_entity = mocker.MagicMock(app_config=app_config, inputs={}, query="q", files=[], stream=True)
@@ -22,7 +23,7 @@ class TestAgentChatAppRunnerRun:
with pytest.raises(ValueError):
runner.run(generate_entity, mocker.MagicMock(), mocker.MagicMock(), mocker.MagicMock())
def test_run_moderation_error_direct_output(self, runner, mocker):
def test_run_moderation_error_direct_output(self, runner, mocker: MockerFixture):
app_record = mocker.MagicMock(id="app1", tenant_id="tenant")
app_config = mocker.MagicMock(app_id="app1", tenant_id="tenant", prompt_template=mocker.MagicMock())
app_config.agent = mocker.MagicMock()
@@ -45,7 +46,7 @@ class TestAgentChatAppRunnerRun:
runner.direct_output.assert_called_once()
def test_run_annotation_reply_short_circuits(self, runner, mocker):
def test_run_annotation_reply_short_circuits(self, runner, mocker: MockerFixture):
app_record = mocker.MagicMock(id="app1", tenant_id="tenant")
app_config = mocker.MagicMock(app_id="app1", tenant_id="tenant", prompt_template=mocker.MagicMock())
app_config.agent = mocker.MagicMock()
@@ -74,7 +75,7 @@ class TestAgentChatAppRunnerRun:
queue_manager.publish.assert_called_once()
runner.direct_output.assert_called_once()
def test_run_hosting_moderation_short_circuits(self, runner, mocker):
def test_run_hosting_moderation_short_circuits(self, runner, mocker: MockerFixture):
app_record = mocker.MagicMock(id="app1", tenant_id="tenant")
app_config = mocker.MagicMock(app_id="app1", tenant_id="tenant", prompt_template=mocker.MagicMock())
app_config.agent = mocker.MagicMock()
@@ -98,7 +99,7 @@ class TestAgentChatAppRunnerRun:
runner.run(generate_entity, mocker.MagicMock(), mocker.MagicMock(), mocker.MagicMock())
def test_run_model_schema_missing(self, runner, mocker):
def test_run_model_schema_missing(self, runner, mocker: MockerFixture):
app_record = mocker.MagicMock(id="app1", tenant_id="tenant")
app_config = mocker.MagicMock(app_id="app1", tenant_id="tenant", prompt_template=mocker.MagicMock())
app_config.agent = AgentEntity(provider="p", model="m", strategy=AgentEntity.Strategy.CHAIN_OF_THOUGHT)
@@ -140,7 +141,7 @@ class TestAgentChatAppRunnerRun:
(LLMMode.COMPLETION, "CotCompletionAgentRunner"),
],
)
def test_run_chain_of_thought_modes(self, runner, mocker, mode, expected_runner):
def test_run_chain_of_thought_modes(self, runner, mocker: MockerFixture, mode, expected_runner):
app_record = mocker.MagicMock(id="app1", tenant_id="tenant")
app_config = mocker.MagicMock(app_id="app1", tenant_id="tenant", prompt_template=mocker.MagicMock())
app_config.agent = AgentEntity(provider="p", model="m", strategy=AgentEntity.Strategy.CHAIN_OF_THOUGHT)
@@ -196,7 +197,7 @@ class TestAgentChatAppRunnerRun:
runner_instance.run.assert_called_once()
runner._handle_invoke_result.assert_called_once()
def test_run_invalid_llm_mode_raises(self, runner, mocker):
def test_run_invalid_llm_mode_raises(self, runner, mocker: MockerFixture):
app_record = mocker.MagicMock(id="app1", tenant_id="tenant")
app_config = mocker.MagicMock(app_id="app1", tenant_id="tenant", prompt_template=mocker.MagicMock())
app_config.agent = AgentEntity(provider="p", model="m", strategy=AgentEntity.Strategy.CHAIN_OF_THOUGHT)
@@ -242,7 +243,7 @@ class TestAgentChatAppRunnerRun:
with pytest.raises(ValueError):
runner.run(generate_entity, mocker.MagicMock(), conversation, message)
def test_run_function_calling_strategy_selected_by_features(self, runner, mocker):
def test_run_function_calling_strategy_selected_by_features(self, runner, mocker: MockerFixture):
app_record = mocker.MagicMock(id="app1", tenant_id="tenant")
app_config = mocker.MagicMock(app_id="app1", tenant_id="tenant", prompt_template=mocker.MagicMock())
app_config.agent = AgentEntity(provider="p", model="m", strategy=AgentEntity.Strategy.CHAIN_OF_THOUGHT)
@@ -298,7 +299,7 @@ class TestAgentChatAppRunnerRun:
assert app_config.agent.strategy == AgentEntity.Strategy.FUNCTION_CALLING
runner_instance.run.assert_called_once()
def test_run_conversation_not_found(self, runner, mocker):
def test_run_conversation_not_found(self, runner, mocker: MockerFixture):
app_record = mocker.MagicMock(id="app1", tenant_id="tenant")
app_config = mocker.MagicMock(app_id="app1", tenant_id="tenant", prompt_template=mocker.MagicMock())
app_config.agent = AgentEntity(provider="p", model="m", strategy=AgentEntity.Strategy.FUNCTION_CALLING)
@@ -332,7 +333,7 @@ class TestAgentChatAppRunnerRun:
with pytest.raises(ValueError):
runner.run(generate_entity, mocker.MagicMock(), mocker.MagicMock(id="conv"), mocker.MagicMock(id="msg"))
def test_run_message_not_found(self, runner, mocker):
def test_run_message_not_found(self, runner, mocker: MockerFixture):
app_record = mocker.MagicMock(id="app1", tenant_id="tenant")
app_config = mocker.MagicMock(app_id="app1", tenant_id="tenant", prompt_template=mocker.MagicMock())
app_config.agent = AgentEntity(provider="p", model="m", strategy=AgentEntity.Strategy.FUNCTION_CALLING)
@@ -366,7 +367,7 @@ class TestAgentChatAppRunnerRun:
with pytest.raises(ValueError):
runner.run(generate_entity, mocker.MagicMock(), mocker.MagicMock(id="conv"), mocker.MagicMock(id="msg"))
def test_run_invalid_agent_strategy_raises(self, runner, mocker):
def test_run_invalid_agent_strategy_raises(self, runner, mocker: MockerFixture):
app_record = mocker.MagicMock(id="app1", tenant_id="tenant")
app_config = mocker.MagicMock(app_id="app1", tenant_id="tenant", prompt_template=mocker.MagicMock())
app_config.agent = mocker.MagicMock(strategy="invalid", provider="p", model="m")
@@ -2,6 +2,7 @@ from types import SimpleNamespace
from unittest.mock import MagicMock
import pytest
from pytest_mock import MockerFixture
import core.app.apps.completion.app_runner as module
from core.app.apps.completion.app_runner import CompletionAppRunner
@@ -47,7 +48,7 @@ def _build_generate_entity(app_config, file_upload_config=None):
class TestCompletionAppRunner:
def test_run_app_not_found(self, runner, mocker):
def test_run_app_not_found(self, runner, mocker: MockerFixture):
session = mocker.MagicMock()
session.scalar.return_value = None
mocker.patch.object(module.db, "session", session)
@@ -58,7 +59,7 @@ class TestCompletionAppRunner:
with pytest.raises(ValueError):
runner.run(app_generate_entity, MagicMock(), MagicMock())
def test_run_moderation_error_outputs_direct(self, runner, mocker):
def test_run_moderation_error_outputs_direct(self, runner, mocker: MockerFixture):
app_record = MagicMock(id="app1", tenant_id="tenant")
session = mocker.MagicMock()
@@ -78,7 +79,7 @@ class TestCompletionAppRunner:
runner.direct_output.assert_called_once()
runner._handle_invoke_result.assert_not_called()
def test_run_hosting_moderation_stops(self, runner, mocker):
def test_run_hosting_moderation_stops(self, runner, mocker: MockerFixture):
app_record = MagicMock(id="app1", tenant_id="tenant")
session = mocker.MagicMock()
@@ -97,7 +98,7 @@ class TestCompletionAppRunner:
runner._handle_invoke_result.assert_not_called()
def test_run_dataset_and_external_tools_flow(self, runner, mocker):
def test_run_dataset_and_external_tools_flow(self, runner, mocker: MockerFixture):
app_record = MagicMock(id="app1", tenant_id="tenant")
session = mocker.MagicMock()
@@ -140,7 +141,7 @@ class TestCompletionAppRunner:
assert dataset_retrieval.retrieve.call_args.kwargs["query"] == "query_from_input"
runner._handle_invoke_result.assert_called_once()
def test_run_uses_low_image_detail_default(self, runner, mocker):
def test_run_uses_low_image_detail_default(self, runner, mocker: MockerFixture):
app_record = MagicMock(id="app1", tenant_id="tenant")
session = mocker.MagicMock()
@@ -1,6 +1,8 @@
from types import SimpleNamespace
from unittest.mock import MagicMock
from pytest_mock import MockerFixture
import core.app.apps.completion.app_config_manager as module
from core.app.app_config.entities import EasyUIBasedAppModelConfigFrom
from core.app.apps.completion.app_config_manager import CompletionAppConfigManager
@@ -8,7 +10,7 @@ from models.model import AppMode
class TestCompletionAppConfigManager:
def test_get_app_config_with_override(self, mocker):
def test_get_app_config_with_override(self, mocker: MockerFixture):
app_model = MagicMock(tenant_id="tenant", id="app1", mode=AppMode.COMPLETION.value)
app_model_config = MagicMock(id="cfg1")
app_model_config.to_dict.return_value = {"model": {"provider": "x"}}
@@ -35,8 +37,8 @@ class TestCompletionAppConfigManager:
assert result.external_data_variables == ["ext1"]
assert result.app_mode == AppMode.COMPLETION
def test_get_app_config_without_override_uses_model_config(self, mocker):
app_model = MagicMock(tenant_id="tenant", id="app1", mode=AppMode.COMPLETION.value)
def test_get_app_config_without_override_uses_model_config(self, mocker: MockerFixture):
app_model = MagicMock(tenant_id="tenant", id="app1", mode=AppMode.COMPLETION)
app_model_config = MagicMock(id="cfg1")
app_model_config.to_dict.return_value = {"model": {"provider": "x"}}
@@ -53,7 +55,7 @@ class TestCompletionAppConfigManager:
assert result.app_model_config_from == EasyUIBasedAppModelConfigFrom.APP_LATEST_CONFIG
assert result.app_model_config_dict == {"model": {"provider": "x"}}
def test_config_validate_filters_related_keys(self, mocker):
def test_config_validate_filters_related_keys(self, mocker: MockerFixture):
config = {
"model": {"provider": "x"},
"variables": ["v"],
@@ -4,6 +4,7 @@ from unittest.mock import MagicMock
import pytest
from pydantic import ValidationError
from pytest_mock import MockerFixture
import core.app.apps.completion.app_generator as module
from core.app.apps.completion.app_generator import CompletionAppGenerator
@@ -15,7 +16,7 @@ from services.errors.message import MessageNotExistsError
@pytest.fixture
def generator(mocker):
def generator(mocker: MockerFixture):
gen = CompletionAppGenerator()
mocker.patch.object(module, "copy_current_request_context", side_effect=lambda fn: fn)
@@ -69,7 +70,7 @@ class TestCompletionAppGenerator:
streaming=False,
)
def test_generate_success_no_file_config(self, generator, mocker):
def test_generate_success_no_file_config(self, generator, mocker: MockerFixture):
app_model_config = _build_app_model_config()
mocker.patch.object(generator, "_get_app_model_config", return_value=app_model_config)
mocker.patch.object(module.FileUploadConfigManager, "convert", return_value=None)
@@ -99,7 +100,7 @@ class TestCompletionAppGenerator:
assert result == "converted"
module.file_factory.build_from_mappings.assert_not_called()
def test_generate_success_with_files(self, generator, mocker):
def test_generate_success_with_files(self, generator, mocker: MockerFixture):
app_model_config = _build_app_model_config()
mocker.patch.object(generator, "_get_app_model_config", return_value=app_model_config)
@@ -131,7 +132,7 @@ class TestCompletionAppGenerator:
assert result == "converted"
module.file_factory.build_from_mappings.assert_called_once()
def test_generate_override_model_config_debugger(self, generator, mocker):
def test_generate_override_model_config_debugger(self, generator, mocker: MockerFixture):
app_model_config = _build_app_model_config()
mocker.patch.object(generator, "_get_app_model_config", return_value=app_model_config)
@@ -165,7 +166,7 @@ class TestCompletionAppGenerator:
assert get_app_config.call_args.kwargs["override_config_dict"] == override_config
def test_generate_more_like_this_message_not_found(self, generator, mocker):
def test_generate_more_like_this_message_not_found(self, generator, mocker: MockerFixture):
session = mocker.MagicMock()
session.scalar.return_value = None
mocker.patch.object(module.db, "session", session)
@@ -178,7 +179,7 @@ class TestCompletionAppGenerator:
invoke_from=InvokeFrom.WEB_APP,
)
def test_generate_more_like_this_disabled(self, generator, mocker):
def test_generate_more_like_this_disabled(self, generator, mocker: MockerFixture):
app_model = _build_app_model()
app_model.app_model_config = MagicMock(more_like_this=False, more_like_this_dict={"enabled": False})
@@ -195,7 +196,7 @@ class TestCompletionAppGenerator:
invoke_from=InvokeFrom.WEB_APP,
)
def test_generate_more_like_this_app_model_config_missing(self, generator, mocker):
def test_generate_more_like_this_app_model_config_missing(self, generator, mocker: MockerFixture):
app_model = _build_app_model()
app_model.app_model_config = None
@@ -212,7 +213,7 @@ class TestCompletionAppGenerator:
invoke_from=InvokeFrom.WEB_APP,
)
def test_generate_more_like_this_message_config_none(self, generator, mocker):
def test_generate_more_like_this_message_config_none(self, generator, mocker: MockerFixture):
app_model = _build_app_model()
app_model.app_model_config = MagicMock(more_like_this=True, more_like_this_dict={"enabled": True})
@@ -229,7 +230,7 @@ class TestCompletionAppGenerator:
invoke_from=InvokeFrom.WEB_APP,
)
def test_generate_more_like_this_success(self, generator, mocker):
def test_generate_more_like_this_success(self, generator, mocker: MockerFixture):
app_model = _build_app_model()
app_model.app_model_config = MagicMock(more_like_this=True, more_like_this_dict={"enabled": True})
@@ -297,7 +298,7 @@ class TestCompletionAppGenerator:
(RuntimeError("boom"), True),
],
)
def test_generate_worker_error_handling(self, generator, mocker, error, should_publish):
def test_generate_worker_error_handling(self, generator, mocker: MockerFixture, error, should_publish):
flask_app = MagicMock()
flask_app.app_context.return_value = contextlib.nullcontext()
@@ -1,12 +1,14 @@
from types import SimpleNamespace
from unittest.mock import MagicMock
from pytest_mock import MockerFixture
import core.app.apps.pipeline.pipeline_config_manager as module
from core.app.apps.pipeline.pipeline_config_manager import PipelineConfigManager
from models.model import AppMode
def test_get_pipeline_config(mocker):
def test_get_pipeline_config(mocker: MockerFixture):
pipeline = MagicMock(tenant_id="tenant", id="pipe1")
workflow = MagicMock(id="wf1")
@@ -26,7 +28,7 @@ def test_get_pipeline_config(mocker):
assert result.rag_pipeline_variables == ["var1"]
def test_config_validate_filters_related_keys(mocker):
def test_config_validate_filters_related_keys(mocker: MockerFixture):
config = {
"file_upload": {"enabled": True},
"tts": {"enabled": True},
@@ -3,6 +3,7 @@ from types import SimpleNamespace
from unittest.mock import MagicMock, PropertyMock
import pytest
from pytest_mock import MockerFixture
import core.app.apps.pipeline.pipeline_generator as module
from core.app.apps.exc import GenerateTaskStoppedError
@@ -23,7 +24,7 @@ class FakeRagPipelineGenerateEntity(SimpleNamespace):
@pytest.fixture
def generator(mocker):
def generator(mocker: MockerFixture):
gen = module.PipelineGenerator()
mocker.patch.object(module, "RagPipelineGenerateEntity", FakeRagPipelineGenerateEntity)
@@ -88,7 +89,7 @@ class DummySession:
return False
def test_generate_dataset_missing(generator, mocker):
def test_generate_dataset_missing(generator, mocker: MockerFixture):
pipeline = _build_pipeline()
pipeline.retrieve_dataset.return_value = None
@@ -106,7 +107,7 @@ def test_generate_dataset_missing(generator, mocker):
)
def test_generate_debugger_calls_generate(generator, mocker):
def test_generate_debugger_calls_generate(generator, mocker: MockerFixture):
pipeline = _build_pipeline()
workflow = _build_workflow()
@@ -150,7 +151,7 @@ def test_generate_debugger_calls_generate(generator, mocker):
assert result == {"result": "ok"}
def test_generate_published_pipeline_creates_documents_and_delay(generator, mocker):
def test_generate_published_pipeline_creates_documents_and_delay(generator, mocker: MockerFixture):
pipeline = _build_pipeline()
workflow = _build_workflow()
@@ -228,7 +229,7 @@ def test_generate_published_pipeline_creates_documents_and_delay(generator, mock
task_proxy.delay.assert_called_once()
def test_generate_is_retry_calls_generate(generator, mocker):
def test_generate_is_retry_calls_generate(generator, mocker: MockerFixture):
pipeline = _build_pipeline()
workflow = _build_workflow()
@@ -273,7 +274,7 @@ def test_generate_is_retry_calls_generate(generator, mocker):
assert result == {"result": "ok"}
def test_generate_worker_handles_errors(generator, mocker):
def test_generate_worker_handles_errors(generator, mocker: MockerFixture):
flask_app = MagicMock()
flask_app.app_context.return_value = contextlib.nullcontext()
mocker.patch.object(module, "preserve_flask_contexts", _dummy_preserve)
@@ -308,7 +309,7 @@ def test_generate_worker_handles_errors(generator, mocker):
queue_manager.publish_error.assert_called_once()
def test_generate_worker_sets_system_user_id_for_external_call(generator, mocker):
def test_generate_worker_sets_system_user_id_for_external_call(generator, mocker: MockerFixture):
flask_app = MagicMock()
flask_app.app_context.return_value = contextlib.nullcontext()
mocker.patch.object(module, "preserve_flask_contexts", _dummy_preserve)
@@ -341,7 +342,7 @@ def test_generate_worker_sets_system_user_id_for_external_call(generator, mocker
assert module.PipelineRunner.call_args.kwargs["system_user_id"] == "session"
def test_generate_raises_when_workflow_not_found(generator, mocker):
def test_generate_raises_when_workflow_not_found(generator, mocker: MockerFixture):
flask_app = MagicMock()
mocker.patch.object(module, "preserve_flask_contexts", _dummy_preserve)
@@ -369,7 +370,7 @@ def test_generate_raises_when_workflow_not_found(generator, mocker):
)
def test_generate_success_returns_converted(generator, mocker):
def test_generate_success_returns_converted(generator, mocker: MockerFixture):
flask_app = MagicMock()
mocker.patch.object(module, "preserve_flask_contexts", _dummy_preserve)
@@ -409,7 +410,7 @@ def test_generate_success_returns_converted(generator, mocker):
assert result == "converted"
def test_single_iteration_generate_validates_inputs(generator, mocker):
def test_single_iteration_generate_validates_inputs(generator, mocker: MockerFixture):
with pytest.raises(ValueError):
generator.single_iteration_generate(_build_pipeline(), _build_workflow(), "", _build_user(), {})
@@ -419,7 +420,7 @@ def test_single_iteration_generate_validates_inputs(generator, mocker):
)
def test_single_iteration_generate_dataset_required(generator, mocker):
def test_single_iteration_generate_dataset_required(generator, mocker: MockerFixture):
pipeline = _build_pipeline()
pipeline.retrieve_dataset.return_value = None
@@ -436,7 +437,7 @@ def test_single_iteration_generate_dataset_required(generator, mocker):
)
def test_single_iteration_generate_success(generator, mocker):
def test_single_iteration_generate_success(generator, mocker: MockerFixture):
pipeline = _build_pipeline()
session = DummySession()
@@ -476,7 +477,7 @@ def test_single_iteration_generate_success(generator, mocker):
assert result == {"ok": True}
def test_single_loop_generate_success(generator, mocker):
def test_single_loop_generate_success(generator, mocker: MockerFixture):
pipeline = _build_pipeline()
session = DummySession()
@@ -516,7 +517,7 @@ def test_single_loop_generate_success(generator, mocker):
assert result == {"ok": True}
def test_handle_response_value_error_triggers_generate_task_stopped(generator, mocker):
def test_handle_response_value_error_triggers_generate_task_stopped(generator, mocker: MockerFixture):
pipeline = _build_pipeline()
workflow = _build_workflow()
app_entity = FakeRagPipelineGenerateEntity(task_id="t")
@@ -536,7 +537,7 @@ def test_handle_response_value_error_triggers_generate_task_stopped(generator, m
)
def test_build_document_sets_metadata_for_builtin_fields(generator, mocker):
def test_build_document_sets_metadata_for_builtin_fields(generator, mocker: MockerFixture):
class DummyDocument(SimpleNamespace):
pass
@@ -620,7 +621,7 @@ def test_format_datasource_info_list_missing_node_data(generator):
)
def test_format_datasource_info_list_online_drive_folder(generator, mocker):
def test_format_datasource_info_list_online_drive_folder(generator, mocker: MockerFixture):
workflow = MagicMock(
graph_dict={
"nodes": [
@@ -1,4 +1,5 @@
import pytest
from pytest_mock import MockerFixture
import core.app.apps.pipeline.pipeline_queue_manager as module
from core.app.apps.base_app_queue_manager import PublishFrom
@@ -16,7 +17,7 @@ from core.app.entities.queue_entities import (
from graphon.model_runtime.entities.llm_entities import LLMResult
def test_publish_sets_stop_listen_and_raises_on_stopped(mocker):
def test_publish_sets_stop_listen_and_raises_on_stopped(mocker: MockerFixture):
manager = PipelineQueueManager(task_id="t", user_id="u", invoke_from=InvokeFrom.WEB_APP, app_mode="rag")
manager._q = mocker.MagicMock()
manager.stop_listen = mocker.MagicMock()
@@ -28,7 +29,7 @@ def test_publish_sets_stop_listen_and_raises_on_stopped(mocker):
manager.stop_listen.assert_called_once()
def test_publish_stop_events_trigger_stop_listen(mocker):
def test_publish_stop_events_trigger_stop_listen(mocker: MockerFixture):
manager = PipelineQueueManager(task_id="t", user_id="u", invoke_from=InvokeFrom.WEB_APP, app_mode="rag")
manager._q = mocker.MagicMock()
manager.stop_listen = mocker.MagicMock()
@@ -46,7 +47,7 @@ def test_publish_stop_events_trigger_stop_listen(mocker):
manager.stop_listen.assert_called_once()
def test_publish_non_stop_event_no_stop_listen(mocker):
def test_publish_non_stop_event_no_stop_listen(mocker: MockerFixture):
manager = PipelineQueueManager(task_id="t", user_id="u", invoke_from=InvokeFrom.WEB_APP, app_mode="rag")
manager._q = mocker.MagicMock()
manager.stop_listen = mocker.MagicMock()
@@ -22,6 +22,7 @@ from types import SimpleNamespace
from unittest.mock import MagicMock
import pytest
from pytest_mock import MockerFixture
import core.app.apps.pipeline.pipeline_runner as module
from core.app.apps.pipeline.pipeline_runner import PipelineRunner
@@ -126,7 +127,7 @@ def test_update_document_status_on_failure(mocker, runner):
session.commit.assert_called_once()
def test_run_pipeline_not_found(mocker):
def test_run_pipeline_not_found(mocker: MockerFixture):
app_generate_entity = _build_app_generate_entity()
app_generate_entity.invoke_from = InvokeFrom.WEB_APP
app_generate_entity.single_iteration_run = None
@@ -150,7 +151,7 @@ def test_run_pipeline_not_found(mocker):
runner.run()
def test_run_workflow_not_initialized(mocker):
def test_run_workflow_not_initialized(mocker: MockerFixture):
app_generate_entity = _build_app_generate_entity()
pipeline = MagicMock(id="pipe")
@@ -174,7 +175,7 @@ def test_run_workflow_not_initialized(mocker):
runner.run()
def test_run_single_iteration_path(mocker):
def test_run_single_iteration_path(mocker: MockerFixture):
app_generate_entity = _build_app_generate_entity()
app_generate_entity.single_iteration_run = MagicMock()
@@ -223,7 +224,7 @@ def test_run_single_iteration_path(mocker):
runner._handle_event.assert_called()
def test_run_normal_path_builds_graph(mocker):
def test_run_normal_path_builds_graph(mocker: MockerFixture):
app_generate_entity = _build_app_generate_entity()
pipeline = MagicMock(id="pipe")
@@ -45,7 +45,7 @@ def _make_generate_entity(app_config: WorkflowUIBasedAppConfig) -> AdvancedChatA
@pytest.fixture(autouse=True)
def _mock_db_session(monkeypatch):
def _mock_db_session(monkeypatch: pytest.MonkeyPatch):
session = MagicMock()
def refresh_side_effect(obj):
@@ -108,7 +108,7 @@ def test_init_generate_records_marks_existing_conversation():
assert entity.is_new_conversation is False
def test_message_cycle_manager_uses_new_conversation_flag(monkeypatch):
def test_message_cycle_manager_uses_new_conversation_flag(monkeypatch: pytest.MonkeyPatch):
app_config = _make_app_config()
entity = _make_generate_entity(app_config)
entity.conversation_id = "existing-conversation-id"
@@ -369,7 +369,7 @@ def test_validate_inputs_optional_file_with_empty_string_ignores_default():
class TestBaseAppGeneratorExtras:
def test_prepare_user_inputs_converts_files_and_lists(self, monkeypatch):
def test_prepare_user_inputs_converts_files_and_lists(self, monkeypatch: pytest.MonkeyPatch):
base_app_generator = BaseAppGenerator()
variables = [
@@ -42,7 +42,7 @@ class _QueueRecorder:
class TestAppRunner:
def test_recalc_llm_max_tokens_updates_parameters(self, monkeypatch):
def test_recalc_llm_max_tokens_updates_parameters(self, monkeypatch: pytest.MonkeyPatch):
runner = AppRunner()
model_schema = SimpleNamespace(
@@ -65,7 +65,7 @@ class TestAppRunner:
assert model_config.parameters["max_tokens"] == 20
def test_recalc_llm_max_tokens_returns_minus_one_when_no_context(self, monkeypatch):
def test_recalc_llm_max_tokens_returns_minus_one_when_no_context(self, monkeypatch: pytest.MonkeyPatch):
runner = AppRunner()
model_schema = SimpleNamespace(
@@ -86,7 +86,7 @@ class TestAppRunner:
assert runner.recalc_llm_max_tokens(model_config, prompt_messages=[]) == -1
def test_direct_output_streaming_publishes_chunks_and_end(self, monkeypatch):
def test_direct_output_streaming_publishes_chunks_and_end(self, monkeypatch: pytest.MonkeyPatch):
runner = AppRunner()
queue = _QueueRecorder()
app_generate_entity = SimpleNamespace(model_conf=SimpleNamespace(model="mock"), stream=True)
@@ -133,7 +133,7 @@ class TestAppRunner:
stream=True,
)
def test_organize_prompt_messages_simple_template(self, monkeypatch):
def test_organize_prompt_messages_simple_template(self, monkeypatch: pytest.MonkeyPatch):
runner = AppRunner()
model_config = SimpleNamespace(mode="chat", stop=["STOP"])
prompt_template_entity = PromptTemplateEntity(
@@ -158,7 +158,7 @@ class TestAppRunner:
assert prompt_messages == ["simple-message"]
assert stop == ["simple-stop"]
def test_organize_prompt_messages_advanced_completion_template(self, monkeypatch):
def test_organize_prompt_messages_advanced_completion_template(self, monkeypatch: pytest.MonkeyPatch):
runner = AppRunner()
model_config = SimpleNamespace(mode="completion", stop=["<END>"])
captured: dict[str, object] = {}
@@ -191,7 +191,7 @@ class TestAppRunner:
assert memory_config.role_prefix.user == "U"
assert memory_config.role_prefix.assistant == "A"
def test_organize_prompt_messages_advanced_chat_template(self, monkeypatch):
def test_organize_prompt_messages_advanced_chat_template(self, monkeypatch: pytest.MonkeyPatch):
runner = AppRunner()
model_config = SimpleNamespace(mode="chat", stop=["<END>"])
captured: dict[str, object] = {}
@@ -245,7 +245,7 @@ class TestAppRunner:
files=[],
)
def test_handle_invoke_result_stream_routes_chunks_and_builds_message(self, monkeypatch):
def test_handle_invoke_result_stream_routes_chunks_and_builds_message(self, monkeypatch: pytest.MonkeyPatch):
runner = AppRunner()
queue = _QueueRecorder()
warning_logger = MagicMock()
@@ -284,7 +284,7 @@ class TestAppRunner:
assert queue.events[-1].llm_result.message.content == "abc"
warning_logger.assert_called_once()
def test_handle_invoke_result_stream_agent_mode_handles_multimodal_errors(self, monkeypatch):
def test_handle_invoke_result_stream_agent_mode_handles_multimodal_errors(self, monkeypatch: pytest.MonkeyPatch):
runner = AppRunner()
queue = _QueueRecorder()
exception_logger = MagicMock()
@@ -331,7 +331,7 @@ class TestAppRunner:
assert queue.events[-1].llm_result.usage == usage
exception_logger.assert_called_once()
def test_handle_multimodal_image_content_fallback_return_branch(self, monkeypatch):
def test_handle_multimodal_image_content_fallback_return_branch(self, monkeypatch: pytest.MonkeyPatch):
runner = AppRunner()
class _ToggleBool:
@@ -367,7 +367,7 @@ class TestAppRunner:
db_session.add.assert_not_called()
queue_manager.publish.assert_not_called()
def test_check_hosting_moderation_direct_output_called(self, monkeypatch):
def test_check_hosting_moderation_direct_output_called(self, monkeypatch: pytest.MonkeyPatch):
runner = AppRunner()
queue = _QueueRecorder()
app_generate_entity = SimpleNamespace(stream=False)
@@ -388,7 +388,7 @@ class TestAppRunner:
assert result is True
assert direct_output.called
def test_fill_in_inputs_from_external_data_tools(self, monkeypatch):
def test_fill_in_inputs_from_external_data_tools(self, monkeypatch: pytest.MonkeyPatch):
runner = AppRunner()
monkeypatch.setattr(
"core.app.apps.base_app_runner.ExternalDataFetch.fetch",
@@ -405,7 +405,7 @@ class TestAppRunner:
assert result == {"foo": "bar"}
def test_moderation_for_inputs_returns_result(self, monkeypatch):
def test_moderation_for_inputs_returns_result(self, monkeypatch: pytest.MonkeyPatch):
runner = AppRunner()
monkeypatch.setattr(
"core.app.apps.base_app_runner.InputModeration.check",
@@ -424,7 +424,7 @@ class TestAppRunner:
assert result == (True, {}, "")
def test_query_app_annotations_to_reply(self, monkeypatch):
def test_query_app_annotations_to_reply(self, monkeypatch: pytest.MonkeyPatch):
runner = AppRunner()
monkeypatch.setattr(
"core.app.apps.base_app_runner.AnnotationReplyFeature.query",
@@ -85,7 +85,7 @@ def _make_chat_generate_entity(app_config: EasyUIBasedAppConfig) -> ChatAppGener
@pytest.fixture(autouse=True)
def _mock_db_session(monkeypatch):
def _mock_db_session(monkeypatch: pytest.MonkeyPatch):
session = MagicMock()
def refresh_side_effect(obj):
@@ -130,7 +130,7 @@ def test_init_generate_records_sets_conversation_fields_for_chat_entity():
class TestMessageBasedAppGeneratorExtras:
def test_handle_response_closed_file_raises_stopped(self, monkeypatch):
def test_handle_response_closed_file_raises_stopped(self, monkeypatch: pytest.MonkeyPatch):
generator = MessageBasedAppGenerator()
class _Pipeline:
@@ -155,7 +155,7 @@ class TestMessageBasedAppGeneratorExtras:
stream=False,
)
def test_get_app_model_config_requires_valid_config(self, monkeypatch):
def test_get_app_model_config_requires_valid_config(self, monkeypatch: pytest.MonkeyPatch):
generator = MessageBasedAppGenerator()
app_model = SimpleNamespace(id="app", app_model_config_id=None, app_model_config=None)
@@ -3,6 +3,8 @@ import time
from types import ModuleType, SimpleNamespace
from typing import Any
from pytest_mock import MockerFixture
import graphon.nodes.human_input.entities # noqa: F401
from core.app.apps.advanced_chat import app_generator as adv_app_gen_module
from core.app.apps.workflow import app_generator as wf_app_gen_module
@@ -101,7 +103,7 @@ class _StubToolNode(Node[_StubToolNodeData]):
yield self._convert_node_run_result_to_graph_node_event(result)
def _patch_tool_node(mocker):
def _patch_tool_node(mocker: MockerFixture):
original_resolve_node_class = node_factory_module.resolve_workflow_node_class
def _patched_resolve_node_class(*, node_type: NodeType, node_version: str) -> type[Node]:
@@ -196,7 +198,7 @@ def _node_successes(events: list[GraphEngineEvent]) -> list[str]:
return [evt.node_id for evt in events if isinstance(evt, NodeRunSucceededEvent)]
def test_workflow_app_pause_resume_matches_baseline(mocker):
def test_workflow_app_pause_resume_matches_baseline(mocker: MockerFixture):
_patch_tool_node(mocker)
baseline_state = _build_runtime_state("baseline")
@@ -236,7 +238,7 @@ def test_workflow_app_pause_resume_matches_baseline(mocker):
assert resumed_state.outputs == baseline_outputs
def test_advanced_chat_pause_resume_matches_baseline(mocker):
def test_advanced_chat_pause_resume_matches_baseline(mocker: MockerFixture):
_patch_tool_node(mocker)
baseline_state = _build_runtime_state("adv-baseline")
@@ -54,7 +54,7 @@ class FakeTopic:
return self._state["subscribed"]
def test_retrieve_events_calls_on_subscribe_after_subscription(monkeypatch):
def test_retrieve_events_calls_on_subscribe_after_subscription(monkeypatch: pytest.MonkeyPatch):
topic = FakeTopic()
def fake_get_response_topic(cls, app_mode, workflow_run_id):
@@ -92,7 +92,7 @@ def test_normalize_terminal_events_empty_values():
assert _normalize_terminal_events([]) == set({})
def test_stream_topic_events_emits_ping_and_idle_timeout(monkeypatch):
def test_stream_topic_events_emits_ping_and_idle_timeout(monkeypatch: pytest.MonkeyPatch):
topic = FakeTopic()
times = [1000.0, 1000.0, 1001.0, 1001.0, 1002.0]
@@ -1,6 +1,8 @@
from types import SimpleNamespace
from unittest.mock import MagicMock
from pytest_mock import MockerFixture
from core.app.apps.workflow.app_generator import SKIP_PREPARE_USER_INPUTS_KEY, WorkflowAppGenerator
@@ -22,7 +24,7 @@ def test_should_prepare_user_inputs_keeps_validation_when_flag_false():
assert WorkflowAppGenerator()._should_prepare_user_inputs(args)
def test_resume_delegates_to_generate(mocker):
def test_resume_delegates_to_generate(mocker: MockerFixture):
generator = WorkflowAppGenerator()
mock_generate = mocker.patch.object(generator, "_generate", return_value="ok")
@@ -52,7 +54,7 @@ def test_resume_delegates_to_generate(mocker):
assert kwargs["invoke_from"] == "debugger"
def test_generate_appends_pause_layer_and_forwards_state(mocker):
def test_generate_appends_pause_layer_and_forwards_state(mocker: MockerFixture):
generator = WorkflowAppGenerator()
mock_queue_manager = MagicMock()
@@ -124,7 +126,7 @@ def test_generate_appends_pause_layer_and_forwards_state(mocker):
assert worker_kwargs["kwargs"]["graph_runtime_state"] is graph_runtime_state
def test_resume_path_runs_worker_with_runtime_state(mocker):
def test_resume_path_runs_worker_with_runtime_state(mocker: MockerFixture):
generator = WorkflowAppGenerator()
runtime_state = MagicMock(name="runtime-state")
@@ -90,7 +90,7 @@ class TestWorkflowBasedAppRunner:
with pytest.raises(ValueError, match="Neither single_iteration_run nor single_loop_run"):
runner._prepare_single_node_execution(workflow, None, None, user_id="00000000-0000-0000-0000-000000000001")
def test_get_graph_and_variable_pool_for_single_node_run(self, monkeypatch):
def test_get_graph_and_variable_pool_for_single_node_run(self, monkeypatch: pytest.MonkeyPatch):
runner = WorkflowBasedAppRunner(queue_manager=SimpleNamespace(), app_id="app")
graph_runtime_state = GraphRuntimeState(
variable_pool=VariablePool(system_variables=default_system_variables()),
@@ -142,7 +142,9 @@ class TestWorkflowBasedAppRunner:
assert graph is not None
assert variable_pool is graph_runtime_state.variable_pool
def test_get_graph_and_variable_pool_preloads_constructor_variables_before_graph_init(self, monkeypatch):
def test_get_graph_and_variable_pool_preloads_constructor_variables_before_graph_init(
self, monkeypatch: pytest.MonkeyPatch
):
variable_loader = SimpleNamespace(
load_variables=lambda selectors: (
[
@@ -232,7 +234,7 @@ class TestWorkflowBasedAppRunner:
assert graph is not None
assert variable_pool.get(["sys", "conversation_id"]).value == "conv-1"
def test_handle_graph_run_events_and_pause_notifications(self, monkeypatch):
def test_handle_graph_run_events_and_pause_notifications(self, monkeypatch: pytest.MonkeyPatch):
published: list[object] = []
class _QueueManager:
@@ -67,7 +67,7 @@ class TestWorkflowAppGeneratorValidation:
class TestWorkflowAppGeneratorHandleResponse:
def test_handle_response_closed_file_raises_stopped(self, monkeypatch):
def test_handle_response_closed_file_raises_stopped(self, monkeypatch: pytest.MonkeyPatch):
generator = WorkflowAppGenerator()
app_config = WorkflowUIBasedAppConfig(
@@ -116,7 +116,7 @@ class TestWorkflowAppGeneratorHandleResponse:
class TestWorkflowAppGeneratorGenerate:
def test_generate_skips_prepare_inputs_when_flag_set(self, monkeypatch):
def test_generate_skips_prepare_inputs_when_flag_set(self, monkeypatch: pytest.MonkeyPatch):
generator = WorkflowAppGenerator()
app_config = WorkflowUIBasedAppConfig(
@@ -187,7 +187,7 @@ class TestWorkflowGenerateTaskPipeline:
assert isinstance(responses[0], ValueError)
def test_handle_workflow_started_event_sets_run_id(self, monkeypatch):
def test_handle_workflow_started_event_sets_run_id(self, monkeypatch: pytest.MonkeyPatch):
pipeline = _make_pipeline()
pipeline._graph_runtime_state = GraphRuntimeState(
variable_pool=build_test_variable_pool(variables=build_system_variables(workflow_execution_id="run-id")),
@@ -408,7 +408,7 @@ class TestWorkflowGenerateTaskPipeline:
assert list(pipeline._handle_human_input_form_timeout_event(timeout_event)) == ["timeout"]
assert list(pipeline._handle_agent_log_event(agent_event)) == ["log"]
def test_wrapper_process_stream_response_emits_audio_end(self, monkeypatch):
def test_wrapper_process_stream_response_emits_audio_end(self, monkeypatch: pytest.MonkeyPatch):
pipeline = _make_pipeline()
pipeline._workflow_features_dict = {
"text_to_speech": {"enabled": True, "autoPlay": "enabled", "voice": "v", "language": "en"}
@@ -560,7 +560,7 @@ class TestWorkflowGenerateTaskPipeline:
responses = list(pipeline._wrapper_process_stream_response())
assert responses == [PingStreamResponse(task_id="task")]
def test_wrapper_process_stream_response_final_audio_none_then_finish(self, monkeypatch):
def test_wrapper_process_stream_response_final_audio_none_then_finish(self, monkeypatch: pytest.MonkeyPatch):
pipeline = _make_pipeline()
pipeline._workflow_features_dict = {
"text_to_speech": {"enabled": True, "autoPlay": "enabled", "voice": "v", "language": "en"}
@@ -597,7 +597,7 @@ class TestWorkflowGenerateTaskPipeline:
assert sleep_spy
assert any(isinstance(item, MessageAudioEndStreamResponse) for item in responses)
def test_wrapper_process_stream_response_handles_audio_exception(self, monkeypatch):
def test_wrapper_process_stream_response_handles_audio_exception(self, monkeypatch: pytest.MonkeyPatch):
pipeline = _make_pipeline()
pipeline._workflow_features_dict = {
"text_to_speech": {"enabled": True, "autoPlay": "enabled", "voice": "v", "language": "en"}
@@ -633,7 +633,7 @@ class TestWorkflowGenerateTaskPipeline:
assert logger_exception
assert any(isinstance(item, MessageAudioEndStreamResponse) for item in responses)
def test_database_session_rolls_back_on_error(self, monkeypatch):
def test_database_session_rolls_back_on_error(self, monkeypatch: pytest.MonkeyPatch):
pipeline = _make_pipeline()
calls = {"enter": 0, "exit_exc": None}
@@ -143,7 +143,7 @@ class TestEasyUiBasedGenerateTaskPipeline:
assert pipeline._listen_audio_msg(publisher=None, task_id="task") is None
def test_process_stream_response_handles_chunks_and_end(self, monkeypatch):
def test_process_stream_response_handles_chunks_and_end(self, monkeypatch: pytest.MonkeyPatch):
conversation = SimpleNamespace(id="conv", mode=AppMode.CHAT)
message = SimpleNamespace(id="msg", created_at=datetime.now(UTC))
@@ -245,7 +245,7 @@ class TestEasyUiBasedGenerateTaskPipeline:
assert any(isinstance(event, QueueLLMChunkEvent) for event in events)
assert any(isinstance(event, QueueStopEvent) for event in events)
def test_handle_stop_updates_usage(self, monkeypatch):
def test_handle_stop_updates_usage(self, monkeypatch: pytest.MonkeyPatch):
conversation = SimpleNamespace(id="conv", mode=AppMode.CHAT)
message = SimpleNamespace(id="msg", created_at=datetime.now(UTC))
@@ -313,7 +313,7 @@ class TestEasyUiBasedGenerateTaskPipeline:
assert pipeline._task_state.llm_result.usage.prompt_tokens == 10
assert pipeline._task_state.llm_result.usage.completion_tokens == 5
def test_record_files_builds_file_payloads(self, monkeypatch):
def test_record_files_builds_file_payloads(self, monkeypatch: pytest.MonkeyPatch):
conversation = SimpleNamespace(id="conv", mode=AppMode.CHAT)
message = SimpleNamespace(id="msg", created_at=datetime.now(UTC))
@@ -405,7 +405,7 @@ class TestEasyUiBasedGenerateTaskPipeline:
assert files
assert len(files) == 3
def test_process_stream_response_handles_annotation_and_error(self, monkeypatch):
def test_process_stream_response_handles_annotation_and_error(self, monkeypatch: pytest.MonkeyPatch):
conversation = SimpleNamespace(id="conv", mode=AppMode.CHAT)
message = SimpleNamespace(id="msg", created_at=datetime.now(UTC))
@@ -472,7 +472,7 @@ class TestEasyUiBasedGenerateTaskPipeline:
assert isinstance(responses[-1], ValueError)
assert pipeline._task_state.llm_result.message.content == "annotatedagent"
def test_agent_thought_to_stream_response_returns_payload(self, monkeypatch):
def test_agent_thought_to_stream_response_returns_payload(self, monkeypatch: pytest.MonkeyPatch):
conversation = SimpleNamespace(id="conv", mode=AppMode.CHAT)
message = SimpleNamespace(id="msg", created_at=datetime.now(UTC))
@@ -681,7 +681,7 @@ class TestEasyUiBasedGenerateTaskPipeline:
assert responses == ["payload"]
def test_wrapper_process_stream_response_with_tts_publisher(self, monkeypatch):
def test_wrapper_process_stream_response_with_tts_publisher(self, monkeypatch: pytest.MonkeyPatch):
conversation = SimpleNamespace(id="conv", mode=AppMode.CHAT)
message = SimpleNamespace(id="msg", created_at=datetime.now(UTC))
entity = _make_entity(ChatAppGenerateEntity, AppMode.CHAT)
@@ -715,7 +715,7 @@ class TestEasyUiBasedGenerateTaskPipeline:
assert responses[1] == "payload"
assert isinstance(responses[-1], MessageAudioEndStreamResponse)
def test_wrapper_process_stream_response_timeout_yields_audio_chunk(self, monkeypatch):
def test_wrapper_process_stream_response_timeout_yields_audio_chunk(self, monkeypatch: pytest.MonkeyPatch):
conversation = SimpleNamespace(id="conv", mode=AppMode.CHAT)
message = SimpleNamespace(id="msg", created_at=datetime.now(UTC))
entity = _make_entity(ChatAppGenerateEntity, AppMode.CHAT)
@@ -756,7 +756,7 @@ class TestEasyUiBasedGenerateTaskPipeline:
assert any(isinstance(item, MessageAudioStreamResponse) for item in responses)
assert isinstance(responses[-1], MessageAudioEndStreamResponse)
def test_process_stream_response_handles_stop_event_and_output_replacement(self, monkeypatch):
def test_process_stream_response_handles_stop_event_and_output_replacement(self, monkeypatch: pytest.MonkeyPatch):
conversation = SimpleNamespace(id="conv", mode=AppMode.CHAT)
message = SimpleNamespace(id="msg", created_at=datetime.now(UTC))
pipeline = EasyUIBasedGenerateTaskPipeline(
@@ -896,7 +896,7 @@ class TestEasyUiBasedGenerateTaskPipeline:
assert list(pipeline._process_stream_response(publisher=None)) == []
def test_save_message_persists_fields_and_emits_trace(self, monkeypatch):
def test_save_message_persists_fields_and_emits_trace(self, monkeypatch: pytest.MonkeyPatch):
conversation = SimpleNamespace(id="conv", mode=AppMode.CHAT)
message = SimpleNamespace(id="msg", created_at=datetime.now(UTC))
pipeline = EasyUIBasedGenerateTaskPipeline(
@@ -981,7 +981,7 @@ class TestEasyUiBasedGenerateTaskPipeline:
with pytest.raises(ValueError, match="Conversation conv not found"):
pipeline._save_message(session=session)
def test_message_end_to_stream_response_includes_usage_metadata(self, monkeypatch):
def test_message_end_to_stream_response_includes_usage_metadata(self, monkeypatch: pytest.MonkeyPatch):
conversation = SimpleNamespace(id="conv", mode=AppMode.CHAT)
message = SimpleNamespace(id="msg", created_at=datetime.now(UTC))
pipeline = EasyUIBasedGenerateTaskPipeline(
@@ -1021,7 +1021,7 @@ class TestEasyUiBasedGenerateTaskPipeline:
assert response.id == "msg"
assert response.metadata["usage"]["prompt_tokens"] == 1
def test_record_files_returns_none_when_message_has_no_files(self, monkeypatch):
def test_record_files_returns_none_when_message_has_no_files(self, monkeypatch: pytest.MonkeyPatch):
conversation = SimpleNamespace(id="conv", mode=AppMode.CHAT)
message = SimpleNamespace(id="msg", created_at=datetime.now(UTC))
pipeline = EasyUIBasedGenerateTaskPipeline(
@@ -1059,7 +1059,7 @@ class TestEasyUiBasedGenerateTaskPipeline:
assert response.files is None
def test_record_files_handles_local_fallback_and_tool_url_variants(self, monkeypatch):
def test_record_files_handles_local_fallback_and_tool_url_variants(self, monkeypatch: pytest.MonkeyPatch):
conversation = SimpleNamespace(id="conv", mode=AppMode.CHAT)
message = SimpleNamespace(id="msg", created_at=datetime.now(UTC))
pipeline = EasyUIBasedGenerateTaskPipeline(
@@ -1155,7 +1155,7 @@ class TestEasyUiBasedGenerateTaskPipeline:
assert response.id == "msg"
assert response.answer == "hello"
def test_agent_thought_to_stream_response_returns_none_when_not_found(self, monkeypatch):
def test_agent_thought_to_stream_response_returns_none_when_not_found(self, monkeypatch: pytest.MonkeyPatch):
conversation = SimpleNamespace(id="conv", mode=AppMode.CHAT)
message = SimpleNamespace(id="msg", created_at=datetime.now(UTC))
pipeline = EasyUIBasedGenerateTaskPipeline(
@@ -46,7 +46,7 @@ class TestDifyNodeFactory:
lambda **_kwargs: node_class,
)
def _factory(self, monkeypatch):
def _factory(self, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr("core.workflow.node_factory.dify_config.CODE_MAX_STRING_LENGTH", 10)
monkeypatch.setattr("core.workflow.node_factory.dify_config.CODE_MAX_NUMBER", 10)
monkeypatch.setattr("core.workflow.node_factory.dify_config.CODE_MIN_NUMBER", -10)
@@ -72,20 +72,20 @@ class TestDifyNodeFactory:
graph_runtime_state=SimpleNamespace(),
)
def test_create_node_unknown_type(self, monkeypatch):
def test_create_node_unknown_type(self, monkeypatch: pytest.MonkeyPatch):
factory = self._factory(monkeypatch)
with pytest.raises(ValueError):
factory.create_node({"id": "node-1", "data": {"type": "unknown"}})
def test_create_node_missing_mapping(self, monkeypatch):
def test_create_node_missing_mapping(self, monkeypatch: pytest.MonkeyPatch):
factory = self._factory(monkeypatch)
monkeypatch.setattr("core.workflow.node_factory.get_node_type_classes_mapping", lambda: {})
with pytest.raises(ValueError):
factory.create_node({"id": "node-1", "data": {"type": BuiltinNodeTypes.START}})
def test_create_node_missing_latest_class(self, monkeypatch):
def test_create_node_missing_latest_class(self, monkeypatch: pytest.MonkeyPatch):
factory = self._factory(monkeypatch)
monkeypatch.setattr(
"core.workflow.node_factory.get_node_type_classes_mapping",
@@ -96,7 +96,7 @@ class TestDifyNodeFactory:
with pytest.raises(ValueError):
factory.create_node({"id": "node-1", "data": {"type": BuiltinNodeTypes.START}})
def test_create_node_selects_versioned_class(self, monkeypatch):
def test_create_node_selects_versioned_class(self, monkeypatch: pytest.MonkeyPatch):
factory = self._factory(monkeypatch)
selected_versions: list[tuple[str, str]] = []
@@ -115,7 +115,7 @@ class TestDifyNodeFactory:
assert node.id == "node-1"
assert selected_versions == [("snapshot", "called")]
def test_create_node_code_branch(self, monkeypatch):
def test_create_node_code_branch(self, monkeypatch: pytest.MonkeyPatch):
factory = self._factory(monkeypatch)
self._stub_node_resolution(monkeypatch, DummyCodeNode)
@@ -124,7 +124,7 @@ class TestDifyNodeFactory:
assert isinstance(node, DummyCodeNode)
assert node.id == "node-1"
def test_create_node_template_transform_branch(self, monkeypatch):
def test_create_node_template_transform_branch(self, monkeypatch: pytest.MonkeyPatch):
factory = self._factory(monkeypatch)
self._stub_node_resolution(monkeypatch, DummyTemplateTransformNode)
@@ -133,7 +133,7 @@ class TestDifyNodeFactory:
assert isinstance(node, DummyTemplateTransformNode)
assert "jinja2_template_renderer" in node.kwargs
def test_create_node_http_request_branch(self, monkeypatch):
def test_create_node_http_request_branch(self, monkeypatch: pytest.MonkeyPatch):
factory = self._factory(monkeypatch)
self._stub_node_resolution(monkeypatch, DummyHttpRequestNode)
@@ -142,7 +142,7 @@ class TestDifyNodeFactory:
assert isinstance(node, DummyHttpRequestNode)
assert "http_request_config" in node.kwargs
def test_create_node_knowledge_retrieval_branch(self, monkeypatch):
def test_create_node_knowledge_retrieval_branch(self, monkeypatch: pytest.MonkeyPatch):
factory = self._factory(monkeypatch)
self._stub_node_resolution(monkeypatch, DummyKnowledgeRetrievalNode)
@@ -151,7 +151,7 @@ class TestDifyNodeFactory:
assert isinstance(node, DummyKnowledgeRetrievalNode)
assert node.kwargs == {}
def test_create_node_document_extractor_branch(self, monkeypatch):
def test_create_node_document_extractor_branch(self, monkeypatch: pytest.MonkeyPatch):
factory = self._factory(monkeypatch)
self._stub_node_resolution(monkeypatch, DummyDocumentExtractorNode)
@@ -2,12 +2,14 @@ from __future__ import annotations
from types import SimpleNamespace
import pytest
from core.app.workflow.layers.observability import ObservabilityLayer
from graphon.enums import BuiltinNodeTypes
class TestObservabilityLayerExtras:
def test_init_tracer_enabled_sets_tracer(self, monkeypatch):
def test_init_tracer_enabled_sets_tracer(self, monkeypatch: pytest.MonkeyPatch):
tracer = object()
monkeypatch.setattr("core.app.workflow.layers.observability.dify_config.ENABLE_OTEL", True)
monkeypatch.setattr("core.app.workflow.layers.observability.is_instrument_flag_enabled", lambda: False)
@@ -18,7 +20,7 @@ class TestObservabilityLayerExtras:
assert layer._is_disabled is False
assert layer._tracer is tracer
def test_init_tracer_disables_when_get_tracer_fails(self, monkeypatch, caplog):
def test_init_tracer_disables_when_get_tracer_fails(self, monkeypatch: pytest.MonkeyPatch, caplog):
monkeypatch.setattr("core.app.workflow.layers.observability.dify_config.ENABLE_OTEL", True)
monkeypatch.setattr("core.app.workflow.layers.observability.is_instrument_flag_enabled", lambda: False)
@@ -33,7 +35,7 @@ class TestObservabilityLayerExtras:
assert layer._tracer is None
assert "Failed to get OpenTelemetry tracer" in caplog.text
def test_init_tracer_disables_when_otel_disabled(self, monkeypatch):
def test_init_tracer_disables_when_otel_disabled(self, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr("core.app.workflow.layers.observability.dify_config.ENABLE_OTEL", False)
monkeypatch.setattr("core.app.workflow.layers.observability.is_instrument_flag_enabled", lambda: False)
@@ -143,7 +145,7 @@ class TestObservabilityLayerExtras:
assert layer._node_contexts == {}
def test_on_node_run_end_calls_span_end(self, monkeypatch):
def test_on_node_run_end_calls_span_end(self, monkeypatch: pytest.MonkeyPatch):
layer = ObservabilityLayer()
layer._is_disabled = False
ended: list[str] = []
@@ -164,7 +166,7 @@ class TestObservabilityLayerExtras:
assert ended == ["ended"]
assert "exec" not in layer._node_contexts
def test_on_node_run_end_logs_detach_failure(self, monkeypatch, caplog):
def test_on_node_run_end_logs_detach_failure(self, monkeypatch: pytest.MonkeyPatch, caplog):
layer = ObservabilityLayer()
layer._is_disabled = False
@@ -186,7 +188,7 @@ class TestObservabilityLayerExtras:
assert "Failed to detach OpenTelemetry token" in caplog.text
assert "exec" not in layer._node_contexts
def test_on_node_run_start_and_end_creates_span(self, monkeypatch):
def test_on_node_run_start_and_end_creates_span(self, monkeypatch: pytest.MonkeyPatch):
layer = ObservabilityLayer()
layer._is_disabled = False
@@ -120,7 +120,7 @@ class TestWorkflowPersistenceLayer:
with pytest.raises(ValueError, match="workflow_execution_id must be provided"):
layer._get_execution_id()
def test_prepare_workflow_inputs_excludes_conversation_id(self, monkeypatch):
def test_prepare_workflow_inputs_excludes_conversation_id(self, monkeypatch: pytest.MonkeyPatch):
layer, _, _, _ = _make_layer()
monkeypatch.setattr(
@@ -3,6 +3,7 @@ import queue
from unittest.mock import MagicMock
import pytest
from pytest_mock import MockerFixture
from core.base.tts.app_generator_tts_publisher import (
AppGeneratorTTSPublisher,
@@ -17,7 +18,7 @@ from core.base.tts.app_generator_tts_publisher import (
@pytest.fixture
def mock_model_instance(mocker):
def mock_model_instance(mocker: MockerFixture):
model = mocker.MagicMock()
model.invoke_tts.return_value = [b"audio1", b"audio2"]
model.get_tts_voices.return_value = [{"value": "voice1"}, {"value": "voice2"}]
@@ -33,7 +34,7 @@ def mock_model_manager(mocker, mock_model_instance):
@pytest.fixture(autouse=True)
def patch_threads(mocker):
def patch_threads(mocker: MockerFixture):
"""Prevent real threads from starting during tests"""
mocker.patch("threading.Thread.start", return_value=None)
@@ -114,7 +115,7 @@ class TestProcessFuture:
finish = audio_queue.get()
assert finish.status == "finish"
def test_process_future_exception(self, mocker):
def test_process_future_exception(self, mocker: MockerFixture):
future_queue = queue.Queue()
audio_queue = queue.Queue()
@@ -222,7 +223,7 @@ class TestAppGeneratorTTSPublisher:
publisher.executor.submit.assert_not_called()
def test_runtime_sentence_threshold_triggers_submit(self, mock_model_manager, mocker):
def test_runtime_sentence_threshold_triggers_submit(self, mock_model_manager, mocker: MockerFixture):
publisher = AppGeneratorTTSPublisher("tenant", "voice1")
publisher.executor = MagicMock()
@@ -297,7 +298,7 @@ class TestAppGeneratorTTSPublisher:
publisher.executor.submit.assert_not_called()
def test_runtime_handles_agent_message_event_list_content(self, mock_model_manager, mocker):
def test_runtime_handles_agent_message_event_list_content(self, mock_model_manager, mocker: MockerFixture):
publisher = AppGeneratorTTSPublisher("tenant", "voice1")
publisher.executor = MagicMock()
@@ -332,7 +333,7 @@ class TestAppGeneratorTTSPublisher:
assert publisher.msg_text == "Hello "
def test_runtime_handles_agent_message_event_empty_content(self, mock_model_manager, mocker):
def test_runtime_handles_agent_message_event_empty_content(self, mock_model_manager, mocker: MockerFixture):
publisher = AppGeneratorTTSPublisher("tenant", "voice1")
publisher.executor = MagicMock()
@@ -358,7 +359,7 @@ class TestAppGeneratorTTSPublisher:
assert publisher.msg_text == ""
def test_runtime_resets_msg_text_when_text_tmp_not_str(self, mock_model_manager, mocker):
def test_runtime_resets_msg_text_when_text_tmp_not_str(self, mock_model_manager, mocker: MockerFixture):
publisher = AppGeneratorTTSPublisher("tenant", "voice1")
publisher.executor = MagicMock()
@@ -1,8 +1,10 @@
from unittest.mock import MagicMock
import pytest
from pytest_mock import MockerFixture
import core.callback_handler.agent_tool_callback_handler as module
from core.callback_handler.agent_tool_callback_handler import DifyAgentCallbackHandler
# -----------------------------
# Fixtures
@@ -10,17 +12,17 @@ import core.callback_handler.agent_tool_callback_handler as module
@pytest.fixture
def enable_debug(mocker):
def enable_debug(mocker: MockerFixture):
mocker.patch.object(module.dify_config, "DEBUG", True)
@pytest.fixture
def disable_debug(mocker):
def disable_debug(mocker: MockerFixture):
mocker.patch.object(module.dify_config, "DEBUG", False)
@pytest.fixture
def mock_print(mocker):
def mock_print(mocker: MockerFixture):
return mocker.patch("builtins.print")
@@ -71,7 +73,7 @@ class TestPrintText:
module.print_text("hello")
mock_print.assert_called_once_with("hello", end="", file=None)
def test_print_text_with_color(self, mocker, mock_print):
def test_print_text_with_color(self, mocker: MockerFixture, mock_print):
mock_get_color = mocker.patch(
"core.callback_handler.agent_tool_callback_handler.get_colored_text",
return_value="colored_text",
@@ -82,7 +84,7 @@ class TestPrintText:
mock_get_color.assert_called_once_with("hello", "green")
mock_print.assert_called_once_with("colored_text", end="", file=None)
def test_print_text_with_file_flush(self, mocker):
def test_print_text_with_file_flush(self, mocker: MockerFixture):
mock_file = MagicMock()
mock_print = mocker.patch("builtins.print")
@@ -107,21 +109,25 @@ class TestDifyAgentCallbackHandler:
assert handler.color == "green"
assert handler.current_loop == 1
def test_on_tool_start_debug_enabled(self, handler, enable_debug, mocker):
def test_on_tool_start_debug_enabled(self, handler: DifyAgentCallbackHandler, enable_debug, mocker: MockerFixture):
mock_print_text = mocker.patch("core.callback_handler.agent_tool_callback_handler.print_text")
handler.on_tool_start("tool1", {"a": 1})
mock_print_text.assert_called()
def test_on_tool_start_debug_disabled(self, handler, disable_debug, mocker):
def test_on_tool_start_debug_disabled(
self, handler: DifyAgentCallbackHandler, disable_debug, mocker: MockerFixture
):
mock_print_text = mocker.patch("core.callback_handler.agent_tool_callback_handler.print_text")
handler.on_tool_start("tool1", {"a": 1})
mock_print_text.assert_not_called()
def test_on_tool_end_debug_enabled_and_trace(self, handler, enable_debug, mocker):
def test_on_tool_end_debug_enabled_and_trace(
self, handler: DifyAgentCallbackHandler, enable_debug, mocker: MockerFixture
):
mock_print_text = mocker.patch("core.callback_handler.agent_tool_callback_handler.print_text")
mock_trace_manager = MagicMock()
@@ -137,7 +143,9 @@ class TestDifyAgentCallbackHandler:
assert mock_print_text.call_count >= 1
mock_trace_manager.add_trace_task.assert_called_once()
def test_on_tool_end_without_trace_manager(self, handler, enable_debug, mocker):
def test_on_tool_end_without_trace_manager(
self, handler: DifyAgentCallbackHandler, enable_debug, mocker: MockerFixture
):
mock_print_text = mocker.patch("core.callback_handler.agent_tool_callback_handler.print_text")
handler.on_tool_end(
@@ -148,14 +156,16 @@ class TestDifyAgentCallbackHandler:
assert mock_print_text.call_count >= 1
def test_on_tool_error_debug_enabled(self, handler, enable_debug, mocker):
def test_on_tool_error_debug_enabled(self, handler: DifyAgentCallbackHandler, enable_debug, mocker: MockerFixture):
mock_print_text = mocker.patch("core.callback_handler.agent_tool_callback_handler.print_text")
handler.on_tool_error(Exception("error"))
mock_print_text.assert_called_once()
def test_on_tool_error_debug_disabled(self, handler, disable_debug, mocker):
def test_on_tool_error_debug_disabled(
self, handler: DifyAgentCallbackHandler, disable_debug, mocker: MockerFixture
):
mock_print_text = mocker.patch("core.callback_handler.agent_tool_callback_handler.print_text")
handler.on_tool_error(Exception("error"))
@@ -163,14 +173,16 @@ class TestDifyAgentCallbackHandler:
mock_print_text.assert_not_called()
@pytest.mark.parametrize("thought", ["thinking", ""])
def test_on_agent_start(self, handler, enable_debug, mocker, thought):
def test_on_agent_start(self, handler: DifyAgentCallbackHandler, enable_debug, mocker: MockerFixture, thought):
mock_print_text = mocker.patch("core.callback_handler.agent_tool_callback_handler.print_text")
handler.on_agent_start(thought)
mock_print_text.assert_called()
def test_on_agent_finish_increments_loop(self, handler, enable_debug, mocker):
def test_on_agent_finish_increments_loop(
self, handler: DifyAgentCallbackHandler, enable_debug, mocker: MockerFixture
):
mock_print_text = mocker.patch("core.callback_handler.agent_tool_callback_handler.print_text")
current_loop = handler.current_loop
@@ -179,19 +191,21 @@ class TestDifyAgentCallbackHandler:
assert handler.current_loop == current_loop + 1
mock_print_text.assert_called()
def test_on_datasource_start_debug_enabled(self, handler, enable_debug, mocker):
def test_on_datasource_start_debug_enabled(
self, handler: DifyAgentCallbackHandler, enable_debug, mocker: MockerFixture
):
mock_print_text = mocker.patch("core.callback_handler.agent_tool_callback_handler.print_text")
handler.on_datasource_start("ds1", {"x": 1})
mock_print_text.assert_called_once()
def test_ignore_agent_property(self, disable_debug, handler):
def test_ignore_agent_property(self, disable_debug, handler: DifyAgentCallbackHandler):
assert handler.ignore_agent is True
def test_ignore_chat_model_property(self, disable_debug, handler):
def test_ignore_chat_model_property(self, disable_debug, handler: DifyAgentCallbackHandler):
assert handler.ignore_chat_model is True
def test_ignore_properties_when_debug_enabled(self, enable_debug, handler):
def test_ignore_properties_when_debug_enabled(self, enable_debug, handler: DifyAgentCallbackHandler):
assert handler.ignore_agent is False
assert handler.ignore_chat_model is False
@@ -1,4 +1,5 @@
import pytest
from pytest_mock import MockerFixture
from core.app.entities.app_invoke_entities import InvokeFrom
from core.callback_handler.index_tool_callback_handler import (
@@ -7,12 +8,12 @@ from core.callback_handler.index_tool_callback_handler import (
@pytest.fixture
def mock_queue_manager(mocker):
def mock_queue_manager(mocker: MockerFixture):
return mocker.Mock()
@pytest.fixture
def handler(mock_queue_manager, mocker):
def handler(mock_queue_manager, mocker: MockerFixture):
mocker.patch(
"core.callback_handler.index_tool_callback_handler.db",
)
@@ -34,7 +35,7 @@ class TestOnQuery:
(InvokeFrom.WEB_APP, "end_user"),
],
)
def test_on_query_success_roles(self, mocker, mock_queue_manager, invoke_from, expected_role):
def test_on_query_success_roles(self, mocker: MockerFixture, mock_queue_manager, invoke_from, expected_role):
# Arrange
mock_db = mocker.patch("core.callback_handler.index_tool_callback_handler.db")
@@ -57,7 +58,7 @@ class TestOnQuery:
assert dataset_query.created_by_role == expected_role
mock_db.session.commit.assert_called_once()
def test_on_query_none_values(self, mocker, mock_queue_manager):
def test_on_query_none_values(self, mocker: MockerFixture, mock_queue_manager):
mock_db = mocker.patch("core.callback_handler.index_tool_callback_handler.db")
handler = DatasetIndexToolCallbackHandler(
@@ -75,7 +76,7 @@ class TestOnQuery:
class TestOnToolEnd:
def test_on_tool_end_no_metadata(self, handler, mocker):
def test_on_tool_end_no_metadata(self, handler: DatasetIndexToolCallbackHandler, mocker: MockerFixture):
mock_db = mocker.patch("core.callback_handler.index_tool_callback_handler.db")
document = mocker.Mock()
@@ -85,7 +86,9 @@ class TestOnToolEnd:
mock_db.session.commit.assert_not_called()
def test_on_tool_end_dataset_document_not_found(self, handler, mocker):
def test_on_tool_end_dataset_document_not_found(
self, handler: DatasetIndexToolCallbackHandler, mocker: MockerFixture
):
mock_db = mocker.patch("core.callback_handler.index_tool_callback_handler.db")
mock_db.session.scalar.return_value = None
@@ -96,7 +99,9 @@ class TestOnToolEnd:
mock_db.session.scalar.assert_called_once()
def test_on_tool_end_parent_child_index_with_child(self, handler, mocker):
def test_on_tool_end_parent_child_index_with_child(
self, handler: DatasetIndexToolCallbackHandler, mocker: MockerFixture
):
mock_db = mocker.patch("core.callback_handler.index_tool_callback_handler.db")
mock_dataset_doc = mocker.Mock()
@@ -119,7 +124,7 @@ class TestOnToolEnd:
mock_db.session.execute.assert_called_once()
mock_db.session.commit.assert_called_once()
def test_on_tool_end_non_parent_child_index(self, handler, mocker):
def test_on_tool_end_non_parent_child_index(self, handler: DatasetIndexToolCallbackHandler, mocker: MockerFixture):
mock_db = mocker.patch("core.callback_handler.index_tool_callback_handler.db")
mock_dataset_doc = mocker.Mock()
@@ -139,12 +144,12 @@ class TestOnToolEnd:
mock_db.session.execute.assert_called_once()
mock_db.session.commit.assert_called_once()
def test_on_tool_end_empty_documents(self, handler):
def test_on_tool_end_empty_documents(self, handler: DatasetIndexToolCallbackHandler):
handler.on_tool_end([])
class TestReturnRetrieverResourceInfo:
def test_publish_called(self, handler, mock_queue_manager, mocker):
def test_publish_called(self, handler: DatasetIndexToolCallbackHandler, mock_queue_manager, mocker: MockerFixture):
mock_event = mocker.patch("core.callback_handler.index_tool_callback_handler.QueueRetrieverResourcesEvent")
resources = [mocker.Mock()]
@@ -1,6 +1,7 @@
from unittest.mock import MagicMock, call
import pytest
from pytest_mock import MockerFixture
from core.callback_handler.workflow_tool_callback_handler import (
DifyWorkflowCallbackHandler,
@@ -26,13 +27,13 @@ def handler():
@pytest.fixture
def mock_print_text(mocker):
def mock_print_text(mocker: MockerFixture):
"""Mock print_text to avoid real stdout printing."""
return mocker.patch("core.callback_handler.workflow_tool_callback_handler.print_text")
class TestDifyWorkflowCallbackHandler:
def test_on_tool_execution_single_output_success(self, handler, mock_print_text):
def test_on_tool_execution_single_output_success(self, handler: DifyWorkflowCallbackHandler, mock_print_text):
# Arrange
tool_name = "test_tool"
tool_inputs = {"a": 1}
@@ -62,7 +63,7 @@ class TestDifyWorkflowCallbackHandler:
]
)
def test_on_tool_execution_multiple_outputs(self, handler, mock_print_text):
def test_on_tool_execution_multiple_outputs(self, handler: DifyWorkflowCallbackHandler, mock_print_text):
# Arrange
tool_name = "multi_tool"
outputs = [
@@ -83,7 +84,7 @@ class TestDifyWorkflowCallbackHandler:
assert results == outputs
assert mock_print_text.call_count == 4 * len(outputs)
def test_on_tool_execution_empty_iterable(self, handler, mock_print_text):
def test_on_tool_execution_empty_iterable(self, handler: DifyWorkflowCallbackHandler, mock_print_text):
# Arrange
tool_name = "empty_tool"
@@ -108,7 +109,9 @@ class TestDifyWorkflowCallbackHandler:
("not_iterable", AttributeError),
],
)
def test_on_tool_execution_invalid_outputs_type(self, handler, invalid_outputs, expected_exception):
def test_on_tool_execution_invalid_outputs_type(
self, handler: DifyWorkflowCallbackHandler, invalid_outputs, expected_exception
):
# Arrange
tool_name = "invalid_tool"
@@ -122,7 +125,7 @@ class TestDifyWorkflowCallbackHandler:
)
)
def test_on_tool_execution_long_json_truncation(self, handler, mock_print_text):
def test_on_tool_execution_long_json_truncation(self, handler: DifyWorkflowCallbackHandler, mock_print_text):
# Arrange
tool_name = "long_json_tool"
long_json = "x" * 1500
@@ -144,7 +147,7 @@ class TestDifyWorkflowCallbackHandler:
color="blue",
)
def test_on_tool_execution_model_dump_json_exception(self, handler, mock_print_text):
def test_on_tool_execution_model_dump_json_exception(self, handler: DifyWorkflowCallbackHandler, mock_print_text):
# Arrange
tool_name = "exception_tool"
bad_message = MagicMock()
@@ -163,7 +166,9 @@ class TestDifyWorkflowCallbackHandler:
# Ensure first two prints happened before failure
assert mock_print_text.call_count >= 2
def test_on_tool_execution_none_message_id_and_trace_manager(self, handler, mock_print_text):
def test_on_tool_execution_none_message_id_and_trace_manager(
self, handler: DifyWorkflowCallbackHandler, mock_print_text
):
# Arrange
tool_name = "optional_params_tool"
message = DummyToolInvokeMessage('{"data": "ok"}')
@@ -2,6 +2,7 @@ import types
from collections.abc import Generator
import pytest
from pytest_mock import MockerFixture
from contexts.wrapper import RecyclableContextVar
from core.datasource.datasource_manager import DatasourceManager
@@ -37,7 +38,7 @@ def _invalidate_recyclable_contextvars() -> None:
RecyclableContextVar.increment_thread_recycles()
def test_get_icon_url_calls_runtime(mocker):
def test_get_icon_url_calls_runtime(mocker: MockerFixture):
fake_runtime = mocker.Mock()
fake_runtime.get_icon_url.return_value = "https://icon"
mocker.patch.object(DatasourceManager, "get_datasource_runtime", return_value=fake_runtime)
@@ -52,7 +53,7 @@ def test_get_icon_url_calls_runtime(mocker):
DatasourceManager.get_datasource_runtime.assert_called_once()
def test_get_datasource_runtime_delegates_to_provider_controller(mocker):
def test_get_datasource_runtime_delegates_to_provider_controller(mocker: MockerFixture):
provider_controller = mocker.Mock()
provider_controller.get_datasource.return_value = object()
mocker.patch.object(DatasourceManager, "get_datasource_plugin_provider", return_value=provider_controller)
@@ -114,7 +115,7 @@ def test_get_datasource_plugin_provider_creates_controller_and_caches(mocker, da
assert ctrl_cls.call_count == 1
def test_get_datasource_plugin_provider_raises_when_provider_entity_missing(mocker):
def test_get_datasource_plugin_provider_raises_when_provider_entity_missing(mocker: MockerFixture):
_invalidate_recyclable_contextvars()
mocker.patch(
"core.datasource.datasource_manager.PluginDatasourceManager.fetch_datasource_provider",
@@ -129,7 +130,7 @@ def test_get_datasource_plugin_provider_raises_when_provider_entity_missing(mock
)
def test_get_datasource_plugin_provider_raises_for_unsupported_type(mocker):
def test_get_datasource_plugin_provider_raises_for_unsupported_type(mocker: MockerFixture):
_invalidate_recyclable_contextvars()
provider_entity = types.SimpleNamespace(declaration=object(), plugin_id="plugin", plugin_unique_identifier="uniq")
mocker.patch(
@@ -145,7 +146,7 @@ def test_get_datasource_plugin_provider_raises_for_unsupported_type(mocker):
)
def test_get_datasource_plugin_provider_raises_when_controller_none(mocker):
def test_get_datasource_plugin_provider_raises_when_controller_none(mocker: MockerFixture):
_invalidate_recyclable_contextvars()
provider_entity = types.SimpleNamespace(declaration=object(), plugin_id="plugin", plugin_unique_identifier="uniq")
mocker.patch(
@@ -165,7 +166,7 @@ def test_get_datasource_plugin_provider_raises_when_controller_none(mocker):
)
def test_stream_online_results_yields_messages_online_document(mocker):
def test_stream_online_results_yields_messages_online_document(mocker: MockerFixture):
# stub runtime to yield a text message
def _doc_messages(**_):
yield from _gen_messages_text_only("hello")
@@ -195,7 +196,7 @@ def test_stream_online_results_yields_messages_online_document(mocker):
assert msgs[0].message.text == "hello"
def test_stream_online_results_sets_credentials_and_returns_empty_dict_online_document(mocker):
def test_stream_online_results_sets_credentials_and_returns_empty_dict_online_document(mocker: MockerFixture):
class _Runtime:
def __init__(self) -> None:
self.runtime = types.SimpleNamespace(credentials=None)
@@ -229,7 +230,7 @@ def test_stream_online_results_sets_credentials_and_returns_empty_dict_online_do
assert final_value == {}
def test_stream_online_results_raises_when_missing_params(mocker):
def test_stream_online_results_raises_when_missing_params(mocker: MockerFixture):
class _Runtime:
def __init__(self) -> None:
self.runtime = types.SimpleNamespace(credentials=None)
@@ -279,7 +280,7 @@ def test_stream_online_results_raises_when_missing_params(mocker):
)
def test_stream_online_results_yields_messages_and_returns_empty_dict_online_drive(mocker):
def test_stream_online_results_yields_messages_and_returns_empty_dict_online_drive(mocker: MockerFixture):
class _Runtime:
def __init__(self) -> None:
self.runtime = types.SimpleNamespace(credentials=None)
@@ -313,7 +314,7 @@ def test_stream_online_results_yields_messages_and_returns_empty_dict_online_dri
assert final_value == {}
def test_stream_online_results_raises_for_unsupported_stream_type(mocker):
def test_stream_online_results_raises_for_unsupported_stream_type(mocker: MockerFixture):
mocker.patch.object(DatasourceManager, "get_datasource_runtime", return_value=mocker.Mock())
mocker.patch(
"core.datasource.datasource_manager.DatasourceProviderService.get_datasource_credentials",
@@ -337,7 +338,7 @@ def test_stream_online_results_raises_for_unsupported_stream_type(mocker):
)
def test_stream_node_events_emits_events_online_document(mocker):
def test_stream_node_events_emits_events_online_document(mocker: MockerFixture):
# make manager's low-level stream produce TEXT only
mocker.patch.object(
DatasourceManager,
@@ -370,7 +371,7 @@ def test_stream_node_events_emits_events_online_document(mocker):
assert events[-1].node_run_result.status == WorkflowNodeExecutionStatus.SUCCEEDED
def test_stream_node_events_builds_file_and_variables_from_messages(mocker):
def test_stream_node_events_builds_file_and_variables_from_messages(mocker: MockerFixture):
mocker.patch.object(DatasourceManager, "stream_online_results", return_value=_gen_messages_text_only("ignored"))
def _transformed(**_kwargs):
@@ -478,7 +479,7 @@ def test_stream_node_events_builds_file_and_variables_from_messages(mocker):
assert events[-1].node_run_result.outputs["x"] == 1
def test_stream_node_events_raises_when_toolfile_missing(mocker):
def test_stream_node_events_raises_when_toolfile_missing(mocker: MockerFixture):
mocker.patch.object(DatasourceManager, "stream_online_results", return_value=_gen_messages_text_only("ignored"))
def _transformed(**_kwargs):
@@ -526,7 +527,7 @@ def test_stream_node_events_raises_when_toolfile_missing(mocker):
)
def test_stream_node_events_online_drive_sets_variable_pool_file_and_outputs(mocker):
def test_stream_node_events_online_drive_sets_variable_pool_file_and_outputs(mocker: MockerFixture):
mocker.patch.object(DatasourceManager, "stream_online_results", return_value=_gen_messages_text_only("ignored"))
file_in = File(
@@ -580,7 +581,7 @@ def test_stream_node_events_online_drive_sets_variable_pool_file_and_outputs(moc
assert completed.node_run_result.outputs["datasource_type"] == DatasourceProviderType.ONLINE_DRIVE
def test_stream_node_events_skips_file_build_for_non_online_types(mocker):
def test_stream_node_events_skips_file_build_for_non_online_types(mocker: MockerFixture):
mocker.patch.object(DatasourceManager, "stream_online_results", return_value=_gen_messages_text_only("ignored"))
def _transformed(**_kwargs):
@@ -620,7 +621,7 @@ def test_stream_node_events_skips_file_build_for_non_online_types(mocker):
assert events[-1].node_run_result.outputs["file"] is None
def test_get_upload_file_by_id_builds_file(mocker):
def test_get_upload_file_by_id_builds_file(mocker: MockerFixture):
# fake UploadFile row
fake_row = types.SimpleNamespace(
id="fid",
@@ -654,7 +655,7 @@ def test_get_upload_file_by_id_builds_file(mocker):
assert f.storage_key == "k"
def test_get_upload_file_by_id_raises_when_missing(mocker):
def test_get_upload_file_by_id_raises_when_missing(mocker: MockerFixture):
class _S:
def __enter__(self):
return self
@@ -1,11 +1,12 @@
import httpx
import pytest
from pytest_mock import MockerFixture
from core.extension.api_based_extension_requestor import APIBasedExtensionRequestor
from models.api_based_extension import APIBasedExtensionPoint
def test_request_success(mocker):
def test_request_success(mocker: MockerFixture):
# Mock httpx.Client and its context manager
mock_client = mocker.MagicMock()
mock_client_instance = mock_client.__enter__.return_value
@@ -28,7 +29,7 @@ def test_request_success(mocker):
)
def test_request_with_ssrf_proxy(mocker):
def test_request_with_ssrf_proxy(mocker: MockerFixture):
# Mock dify_config
mocker.patch("configs.dify_config.SSRF_PROXY_HTTP_URL", "http://proxy:8080")
mocker.patch("configs.dify_config.SSRF_PROXY_HTTPS_URL", "https://proxy:8081")
@@ -59,7 +60,7 @@ def test_request_with_ssrf_proxy(mocker):
assert mock_transport.call_count == 2
def test_request_with_only_one_proxy_config(mocker):
def test_request_with_only_one_proxy_config(mocker: MockerFixture):
# Mock dify_config with only one proxy
mocker.patch("configs.dify_config.SSRF_PROXY_HTTP_URL", "http://proxy:8080")
mocker.patch("configs.dify_config.SSRF_PROXY_HTTPS_URL", None)
@@ -84,7 +85,7 @@ def test_request_with_only_one_proxy_config(mocker):
assert kwargs.get("mounts") is None
def test_request_timeout(mocker):
def test_request_timeout(mocker: MockerFixture):
mock_client = mocker.MagicMock()
mock_client_instance = mock_client.__enter__.return_value
mocker.patch("httpx.Client", return_value=mock_client)
@@ -95,7 +96,7 @@ def test_request_timeout(mocker):
requestor.request(APIBasedExtensionPoint.PING, {})
def test_request_connection_error(mocker):
def test_request_connection_error(mocker: MockerFixture):
mock_client = mocker.MagicMock()
mock_client_instance = mock_client.__enter__.return_value
mocker.patch("httpx.Client", return_value=mock_client)
@@ -106,7 +107,7 @@ def test_request_connection_error(mocker):
requestor.request(APIBasedExtensionPoint.PING, {})
def test_request_error_status_code(mocker):
def test_request_error_status_code(mocker: MockerFixture):
mock_client = mocker.MagicMock()
mock_client_instance = mock_client.__enter__.return_value
mocker.patch("httpx.Client", return_value=mock_client)
@@ -121,7 +122,7 @@ def test_request_error_status_code(mocker):
requestor.request(APIBasedExtensionPoint.PING, {})
def test_request_error_status_code_long_content(mocker):
def test_request_error_status_code_long_content(mocker: MockerFixture):
mock_client = mocker.MagicMock()
mock_client_instance = mock_client.__enter__.return_value
mocker.patch("httpx.Client", return_value=mock_client)
@@ -8,7 +8,7 @@ from yarl import URL
@pytest.fixture(autouse=True)
def _patch_creators_url(monkeypatch):
def _patch_creators_url(monkeypatch: pytest.MonkeyPatch):
"""Patch the module-level creators_platform_api_url for all tests."""
monkeypatch.setattr(
"core.helper.creators.creators_platform_api_url",
@@ -18,7 +18,7 @@ class ConcreteTraceInstance(BaseTraceInstance):
@pytest.fixture
def mock_db_session(monkeypatch):
def mock_db_session(monkeypatch: pytest.MonkeyPatch):
mock_session = MagicMock(spec=Session)
mock_session.__enter__.return_value = mock_session
mock_session.__exit__.return_value = None
@@ -203,7 +203,7 @@ class DummySessionContext:
@pytest.fixture(autouse=True)
def patch_provider_map(monkeypatch):
def patch_provider_map(monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr(
"core.ops.ops_trace_manager.provider_config_map", FakeProviderMap({"dummy": FAKE_PROVIDER_ENTRY})
)
@@ -212,7 +212,7 @@ def patch_provider_map(monkeypatch):
@pytest.fixture(autouse=True)
def patch_timer_and_current_app(monkeypatch):
def patch_timer_and_current_app(monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr("core.ops.ops_trace_manager.threading.Timer", DummyTimer)
monkeypatch.setattr("core.ops.ops_trace_manager.trace_manager_queue", queue.Queue())
monkeypatch.setattr("core.ops.ops_trace_manager.trace_manager_timer", None)
@@ -227,12 +227,12 @@ def patch_timer_and_current_app(monkeypatch):
@pytest.fixture(autouse=True)
def patch_sqlalchemy_session(monkeypatch):
def patch_sqlalchemy_session(monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr("core.ops.ops_trace_manager.Session", DummySessionContext)
@pytest.fixture
def encryption_mocks(monkeypatch):
def encryption_mocks(monkeypatch: pytest.MonkeyPatch):
encrypt_mock = MagicMock(side_effect=lambda tenant, value: f"enc-{value}")
batch_decrypt_mock = MagicMock(side_effect=lambda tenant, values: [f"dec-{value}" for value in values])
obfuscate_mock = MagicMock(side_effect=lambda value: f"ob-{value}")
@@ -243,7 +243,7 @@ def encryption_mocks(monkeypatch):
@pytest.fixture
def mock_db(monkeypatch):
def mock_db(monkeypatch: pytest.MonkeyPatch):
session = MagicMock()
session.scalars.return_value.all.return_value = ["chat"]
db_mock = MagicMock()
@@ -254,7 +254,7 @@ def mock_db(monkeypatch):
@pytest.fixture
def workflow_repo_fixture(monkeypatch):
def workflow_repo_fixture(monkeypatch: pytest.MonkeyPatch):
repo = MagicMock()
repo.get_workflow_run_by_id_without_tenant.return_value = make_workflow_run()
monkeypatch.setattr(TraceTask, "_get_workflow_run_repo", classmethod(lambda cls: repo))
@@ -340,13 +340,13 @@ def test_get_ops_trace_instance_handles_none_app(mock_db):
assert OpsTraceManager.get_ops_trace_instance("app-id") is None
def test_get_ops_trace_instance_returns_none_when_disabled(mock_db, monkeypatch):
def test_get_ops_trace_instance_returns_none_when_disabled(mock_db, monkeypatch: pytest.MonkeyPatch):
app = SimpleNamespace(id="app-id", tracing=json.dumps({"enabled": False}))
mock_db.get.return_value = app
assert OpsTraceManager.get_ops_trace_instance("app-id") is None
def test_get_ops_trace_instance_invalid_provider(mock_db, monkeypatch):
def test_get_ops_trace_instance_invalid_provider(mock_db, monkeypatch: pytest.MonkeyPatch):
app = SimpleNamespace(id="app-id", tracing=json.dumps({"enabled": True, "tracing_provider": "missing"}))
mock_db.get.return_value = app
monkeypatch.setattr("core.ops.ops_trace_manager.provider_config_map", FakeProviderMap({}))
@@ -388,7 +388,7 @@ def test_get_app_config_through_message_id_app_model_config(mock_db):
assert result.id == "cfg"
def test_update_app_tracing_config_invalid_provider(mock_db, monkeypatch):
def test_update_app_tracing_config_invalid_provider(mock_db, monkeypatch: pytest.MonkeyPatch):
mock_db.get.return_value = None
with pytest.raises(ValueError, match="Invalid tracing provider"):
OpsTraceManager.update_app_tracing_config("app", True, "bad")
@@ -421,7 +421,7 @@ def test_get_app_tracing_config_returns_payload(mock_db):
assert OpsTraceManager.get_app_tracing_config("app-id", mock_db) == payload
def test_check_and_project_helpers(monkeypatch):
def test_check_and_project_helpers(monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr(
"core.ops.ops_trace_manager.provider_config_map",
FakeProviderMap(
@@ -449,7 +449,7 @@ def test_check_and_project_helpers(monkeypatch):
assert OpsTraceManager.get_trace_config_project_url({}, "dummy") == "url"
def test_trace_task_conversation_and_extract(monkeypatch):
def test_trace_task_conversation_and_extract(monkeypatch: pytest.MonkeyPatch):
task = TraceTask(trace_type=TraceTaskName.CONVERSATION_TRACE, message_id="msg")
assert task.conversation_trace(foo="bar") == {"foo": "bar"}
assert task._extract_streaming_metrics(make_message_data(message_metadata="not json")) == {}
@@ -525,7 +525,7 @@ def test_extract_streaming_metrics_invalid_json():
assert task._extract_streaming_metrics(fake_message) == {}
def test_trace_queue_manager_add_and_collect(monkeypatch):
def test_trace_queue_manager_add_and_collect(monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr(
"core.ops.ops_trace_manager.OpsTraceManager.get_ops_trace_instance", classmethod(lambda cls, aid: True)
)
@@ -536,7 +536,7 @@ def test_trace_queue_manager_add_and_collect(monkeypatch):
assert tasks == [task]
def test_trace_queue_manager_run_invokes_send(monkeypatch):
def test_trace_queue_manager_run_invokes_send(monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr(
"core.ops.ops_trace_manager.OpsTraceManager.get_ops_trace_instance", classmethod(lambda cls, aid: True)
)
@@ -556,7 +556,7 @@ def test_trace_queue_manager_run_invokes_send(monkeypatch):
assert called["tasks"] == [task]
def test_trace_queue_manager_send_to_celery(monkeypatch):
def test_trace_queue_manager_send_to_celery(monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr(
"core.ops.ops_trace_manager.OpsTraceManager.get_ops_trace_instance", classmethod(lambda cls, aid: True)
)
@@ -19,7 +19,7 @@ import pytest
@pytest.fixture
def trace_queue_manager_and_task(monkeypatch):
def trace_queue_manager_and_task(monkeypatch: pytest.MonkeyPatch):
"""Fixture to provide TraceQueueManager and TraceTask with delayed imports."""
module_name = "core.ops.ops_trace_manager"
if module_name not in sys.modules:
@@ -1,5 +1,7 @@
from types import SimpleNamespace
from pytest_mock import MockerFixture
from core.plugin.entities.request import PluginInvokeContext
from core.plugin.impl.agent import PluginAgentClient
@@ -15,7 +17,7 @@ def _agent_provider(name: str = "agent") -> SimpleNamespace:
class TestPluginAgentClient:
def test_fetch_agent_strategy_providers(self, mocker):
def test_fetch_agent_strategy_providers(self, mocker: MockerFixture):
client = PluginAgentClient()
provider = _agent_provider("remote")
@@ -43,7 +45,7 @@ class TestPluginAgentClient:
assert result[0].declaration.identity.name == "org/plugin/remote"
assert result[0].declaration.strategies[0].identity.provider == "org/plugin/remote"
def test_fetch_agent_strategy_provider(self, mocker):
def test_fetch_agent_strategy_provider(self, mocker: MockerFixture):
client = PluginAgentClient()
provider = _agent_provider("provider")
@@ -63,7 +65,7 @@ class TestPluginAgentClient:
assert result.declaration.identity.name == "org/plugin/provider"
assert result.declaration.strategies[0].identity.provider == "org/plugin/provider"
def test_invoke_merges_chunks_and_passes_context(self, mocker):
def test_invoke_merges_chunks_and_passes_context(self, mocker: MockerFixture):
client = PluginAgentClient()
stream_mock = mocker.patch.object(
client, "_request_with_plugin_daemon_response_stream", return_value=iter(["raw"])
@@ -1,12 +1,13 @@
from unittest.mock import MagicMock
import pytest
from pytest_mock import MockerFixture
from core.plugin.impl.asset import PluginAssetManager
class TestPluginAssetManager:
def test_fetch_asset_success(self, mocker):
def test_fetch_asset_success(self, mocker: MockerFixture):
manager = PluginAssetManager()
response = MagicMock(status_code=200, content=b"asset-bytes")
request_mock = mocker.patch.object(manager, "_request", return_value=response)
@@ -16,14 +17,14 @@ class TestPluginAssetManager:
assert result == b"asset-bytes"
request_mock.assert_called_once_with(method="GET", path="plugin/tenant-1/asset/asset-1")
def test_fetch_asset_not_found_raises(self, mocker):
def test_fetch_asset_not_found_raises(self, mocker: MockerFixture):
manager = PluginAssetManager()
mocker.patch.object(manager, "_request", return_value=MagicMock(status_code=404, content=b""))
with pytest.raises(ValueError, match="can not found asset asset-1"):
manager.fetch_asset("tenant-1", "asset-1")
def test_extract_asset_success(self, mocker):
def test_extract_asset_success(self, mocker: MockerFixture):
manager = PluginAssetManager()
response = MagicMock(status_code=200, content=b"file-content")
request_mock = mocker.patch.object(manager, "_request", return_value=response)
@@ -37,7 +38,7 @@ class TestPluginAssetManager:
params={"plugin_unique_identifier": "org/plugin:1", "file_path": "README.md"},
)
def test_extract_asset_not_found_raises(self, mocker):
def test_extract_asset_not_found_raises(self, mocker: MockerFixture):
manager = PluginAssetManager()
mocker.patch.object(manager, "_request", return_value=MagicMock(status_code=404, content=b""))
@@ -1,6 +1,7 @@
import json
import pytest
from pytest_mock import MockerFixture
from core.plugin.endpoint.exc import EndpointSetupFailedError
from core.plugin.entities.plugin_daemon import PluginDaemonInnerError
@@ -39,7 +40,7 @@ class _StreamContext:
class TestBasePluginClientImpl:
def test_inject_trace_headers(self, mocker):
def test_inject_trace_headers(self, mocker: MockerFixture):
client = BasePluginClient()
mocker.patch("core.plugin.impl.base.dify_config.ENABLE_OTEL", True)
trace_header = "00-abc-xyz-01"
@@ -54,7 +55,7 @@ class TestBasePluginClientImpl:
client._inject_trace_headers(headers_with_existing)
assert headers_with_existing["TraceParent"] == "exists"
def test_stream_request_handles_data_lines_and_dict_payload(self, mocker):
def test_stream_request_handles_data_lines_and_dict_payload(self, mocker: MockerFixture):
client = BasePluginClient()
stream_mock = mocker.patch(
"httpx.Client.stream",
@@ -66,14 +67,14 @@ class TestBasePluginClientImpl:
assert result == ["hello", "world"]
assert stream_mock.call_args.kwargs["data"] == {"k": "v"}
def test_request_with_plugin_daemon_response_handles_request_exception(self, mocker):
def test_request_with_plugin_daemon_response_handles_request_exception(self, mocker: MockerFixture):
client = BasePluginClient()
mocker.patch.object(client, "_request", side_effect=RuntimeError("boom"))
with pytest.raises(ValueError, match="Failed to request plugin daemon"):
client._request_with_plugin_daemon_response("GET", "plugin/tenant/path", bool)
def test_request_with_plugin_daemon_response_applies_transformer(self, mocker):
def test_request_with_plugin_daemon_response_applies_transformer(self, mocker: MockerFixture):
client = BasePluginClient()
mocker.patch.object(client, "_request", return_value=_ResponseStub({"code": 0, "message": "", "data": True}))
@@ -88,14 +89,14 @@ class TestBasePluginClientImpl:
assert result is True
assert transformed == {"code": 0, "message": "", "data": True}
def test_request_with_plugin_daemon_response_stream_malformed_json_error(self, mocker):
def test_request_with_plugin_daemon_response_stream_malformed_json_error(self, mocker: MockerFixture):
client = BasePluginClient()
mocker.patch.object(client, "_stream_request", return_value=iter(['{"error":"bad-line"}']))
with pytest.raises(ValueError, match="bad-line"):
list(client._request_with_plugin_daemon_response_stream("GET", "p", bool))
def test_request_with_plugin_daemon_response_stream_plugin_daemon_inner_error(self, mocker):
def test_request_with_plugin_daemon_response_stream_plugin_daemon_inner_error(self, mocker: MockerFixture):
client = BasePluginClient()
mocker.patch.object(
client, "_stream_request", return_value=iter(['{"code":-500,"message":"not-json","data":null}'])
@@ -105,14 +106,14 @@ class TestBasePluginClientImpl:
list(client._request_with_plugin_daemon_response_stream("GET", "p", bool))
assert exc_info.value.message == "not-json"
def test_request_with_plugin_daemon_response_stream_plugin_daemon_error(self, mocker):
def test_request_with_plugin_daemon_response_stream_plugin_daemon_error(self, mocker: MockerFixture):
client = BasePluginClient()
mocker.patch.object(client, "_stream_request", return_value=iter(['{"code":-1,"message":"err","data":null}']))
with pytest.raises(ValueError, match="plugin daemon: err, code: -1"):
list(client._request_with_plugin_daemon_response_stream("GET", "p", bool))
def test_request_with_plugin_daemon_response_stream_empty_data_error(self, mocker):
def test_request_with_plugin_daemon_response_stream_empty_data_error(self, mocker: MockerFixture):
client = BasePluginClient()
mocker.patch.object(client, "_stream_request", return_value=iter(['{"code":0,"message":"","data":null}']))
@@ -1,5 +1,7 @@
from types import SimpleNamespace
from pytest_mock import MockerFixture
from core.datasource.entities.datasource_entities import (
GetOnlineDocumentPageContentRequest,
OnlineDriveBrowseFilesRequest,
@@ -19,7 +21,7 @@ def _datasource_provider(name: str = "provider") -> SimpleNamespace:
class TestPluginDatasourceManager:
def test_fetch_datasource_providers(self, mocker):
def test_fetch_datasource_providers(self, mocker: MockerFixture):
manager = PluginDatasourceManager()
provider = _datasource_provider("remote")
repack = mocker.patch("core.plugin.impl.datasource.ToolTransformService.repack_provider")
@@ -52,7 +54,7 @@ class TestPluginDatasourceManager:
assert result[1].declaration.datasources[0].identity.provider == "org/plugin/remote"
repack.assert_called_once_with(tenant_id="tenant-1", provider=provider)
def test_fetch_installed_datasource_providers(self, mocker):
def test_fetch_installed_datasource_providers(self, mocker: MockerFixture):
manager = PluginDatasourceManager()
provider = _datasource_provider("remote")
repack = mocker.patch("core.plugin.impl.datasource.ToolTransformService.repack_provider")
@@ -83,7 +85,7 @@ class TestPluginDatasourceManager:
assert result[0].declaration.datasources[0].identity.provider == "org/plugin/remote"
repack.assert_called_once_with(tenant_id="tenant-1", provider=provider)
def test_fetch_datasource_provider_local_and_remote(self, mocker):
def test_fetch_datasource_provider_local_and_remote(self, mocker: MockerFixture):
manager = PluginDatasourceManager()
local = manager.fetch_datasource_provider("tenant-1", "langgenius/file/file")
@@ -113,7 +115,7 @@ class TestPluginDatasourceManager:
assert result.declaration.identity.name == "org/plugin/provider"
assert result.declaration.datasources[0].identity.provider == "org/plugin/provider"
def test_get_website_crawl_streaming(self, mocker):
def test_get_website_crawl_streaming(self, mocker: MockerFixture):
manager = PluginDatasourceManager()
stream_mock = mocker.patch.object(manager, "_request_with_plugin_daemon_response_stream")
stream_mock.return_value = iter(["crawl"])
@@ -132,7 +134,7 @@ class TestPluginDatasourceManager:
assert stream_mock.call_count == 1
def test_get_online_document_pages_streaming(self, mocker):
def test_get_online_document_pages_streaming(self, mocker: MockerFixture):
manager = PluginDatasourceManager()
stream_mock = mocker.patch.object(manager, "_request_with_plugin_daemon_response_stream")
stream_mock.return_value = iter(["pages"])
@@ -151,7 +153,7 @@ class TestPluginDatasourceManager:
assert stream_mock.call_count == 1
def test_get_online_document_page_content_streaming(self, mocker):
def test_get_online_document_page_content_streaming(self, mocker: MockerFixture):
manager = PluginDatasourceManager()
stream_mock = mocker.patch.object(manager, "_request_with_plugin_daemon_response_stream")
stream_mock.return_value = iter(["content"])
@@ -170,7 +172,7 @@ class TestPluginDatasourceManager:
assert stream_mock.call_count == 1
def test_online_drive_browse_files_streaming(self, mocker):
def test_online_drive_browse_files_streaming(self, mocker: MockerFixture):
manager = PluginDatasourceManager()
stream_mock = mocker.patch.object(manager, "_request_with_plugin_daemon_response_stream")
stream_mock.return_value = iter(["browse"])
@@ -189,7 +191,7 @@ class TestPluginDatasourceManager:
assert stream_mock.call_count == 1
def test_online_drive_download_file_streaming(self, mocker):
def test_online_drive_download_file_streaming(self, mocker: MockerFixture):
manager = PluginDatasourceManager()
stream_mock = mocker.patch.object(manager, "_request_with_plugin_daemon_response_stream")
stream_mock.return_value = iter(["download"])
@@ -208,14 +210,14 @@ class TestPluginDatasourceManager:
assert stream_mock.call_count == 1
def test_validate_provider_credentials_returns_true_when_stream_yields_result(self, mocker):
def test_validate_provider_credentials_returns_true_when_stream_yields_result(self, mocker: MockerFixture):
manager = PluginDatasourceManager()
stream_mock = mocker.patch.object(manager, "_request_with_plugin_daemon_response_stream")
stream_mock.return_value = iter([SimpleNamespace(result=True)])
assert manager.validate_provider_credentials("tenant-1", "user-1", "provider", "org/plugin", {"k": "v"}) is True
def test_validate_provider_credentials_returns_false_when_stream_empty(self, mocker):
def test_validate_provider_credentials_returns_false_when_stream_empty(self, mocker: MockerFixture):
manager = PluginDatasourceManager()
stream_mock = mocker.patch.object(manager, "_request_with_plugin_daemon_response_stream")
stream_mock.return_value = iter([])
@@ -1,10 +1,12 @@
from types import SimpleNamespace
from pytest_mock import MockerFixture
from core.plugin.impl.debugging import PluginDebuggingClient
class TestPluginDebuggingClient:
def test_get_debugging_key(self, mocker):
def test_get_debugging_key(self, mocker: MockerFixture):
client = PluginDebuggingClient()
request_mock = mocker.patch.object(
client,
@@ -1,11 +1,12 @@
import pytest
from pytest_mock import MockerFixture
from core.plugin.impl.endpoint import PluginEndpointClient
from core.plugin.impl.exc import PluginDaemonInternalServerError
class TestPluginEndpointClientImpl:
def test_create_endpoint(self, mocker):
def test_create_endpoint(self, mocker: MockerFixture):
client = PluginEndpointClient()
request_mock = mocker.patch.object(client, "_request_with_plugin_daemon_response", return_value=True)
@@ -18,7 +19,7 @@ class TestPluginEndpointClientImpl:
assert args[:3] == ("POST", "plugin/tenant-1/endpoint/setup", bool)
assert kwargs["data"]["plugin_unique_identifier"] == "org/plugin:1"
def test_list_endpoints(self, mocker):
def test_list_endpoints(self, mocker: MockerFixture):
client = PluginEndpointClient()
request_mock = mocker.patch.object(client, "_request_with_plugin_daemon_response", return_value=["endpoint"])
@@ -28,7 +29,7 @@ class TestPluginEndpointClientImpl:
assert request_mock.call_args.args[1] == "plugin/tenant-1/endpoint/list"
assert request_mock.call_args.kwargs["params"] == {"page": 2, "page_size": 20}
def test_list_endpoints_for_single_plugin(self, mocker):
def test_list_endpoints_for_single_plugin(self, mocker: MockerFixture):
client = PluginEndpointClient()
request_mock = mocker.patch.object(client, "_request_with_plugin_daemon_response", return_value=["endpoint"])
@@ -38,7 +39,7 @@ class TestPluginEndpointClientImpl:
assert request_mock.call_args.args[1] == "plugin/tenant-1/endpoint/list/plugin"
assert request_mock.call_args.kwargs["params"] == {"plugin_id": "org/plugin", "page": 1, "page_size": 10}
def test_update_endpoint(self, mocker):
def test_update_endpoint(self, mocker: MockerFixture):
client = PluginEndpointClient()
request_mock = mocker.patch.object(client, "_request_with_plugin_daemon_response", return_value=True)
@@ -47,7 +48,7 @@ class TestPluginEndpointClientImpl:
assert result is True
assert request_mock.call_args.args[:3] == ("POST", "plugin/tenant-1/endpoint/update", bool)
def test_enable_and_disable_endpoint(self, mocker):
def test_enable_and_disable_endpoint(self, mocker: MockerFixture):
client = PluginEndpointClient()
request_mock = mocker.patch.object(client, "_request_with_plugin_daemon_response", return_value=True)
@@ -58,7 +59,7 @@ class TestPluginEndpointClientImpl:
assert calls[0].args[1] == "plugin/tenant-1/endpoint/enable"
assert calls[1].args[1] == "plugin/tenant-1/endpoint/disable"
def test_delete_endpoint_idempotent_and_re_raise(self, mocker):
def test_delete_endpoint_idempotent_and_re_raise(self, mocker: MockerFixture):
client = PluginEndpointClient()
request_mock = mocker.patch.object(client, "_request_with_plugin_daemon_response")
@@ -1,11 +1,13 @@
import json
from pytest_mock import MockerFixture
from core.plugin.impl import exc as exc_module
from core.plugin.impl.exc import PluginDaemonError, PluginInvokeError
class TestPluginImplExceptions:
def test_plugin_daemon_error_str_contains_request_id(self, mocker):
def test_plugin_daemon_error_str_contains_request_id(self, mocker: MockerFixture):
mocker.patch("core.plugin.impl.exc.get_request_id", return_value="req-123")
error = PluginDaemonError("bad")
@@ -21,7 +23,7 @@ class TestPluginImplExceptions:
assert "RateLimit" in friendly
assert "too many" in friendly
def test_plugin_invoke_error_invalid_json_and_fallback(self, mocker):
def test_plugin_invoke_error_invalid_json_and_fallback(self, mocker: MockerFixture):
err = PluginInvokeError("plain text")
assert err._get_error_object() == {}
@@ -32,7 +34,7 @@ class TestPluginImplExceptions:
err2 = PluginInvokeError("plain text")
assert err2.get_error_message() == "plain text"
def test_plugin_invoke_error_get_error_object_handles_adapter_exception(self, mocker):
def test_plugin_invoke_error_get_error_object_handles_adapter_exception(self, mocker: MockerFixture):
adapter = mocker.patch.object(exc_module, "TypeAdapter")
adapter.return_value.validate_json.side_effect = RuntimeError("invalid")
@@ -4,13 +4,14 @@ import io
from types import SimpleNamespace
import pytest
from pytest_mock import MockerFixture
from core.plugin.entities.plugin_daemon import PluginDaemonInnerError
from core.plugin.impl.model import PluginModelClient
class TestPluginModelClient:
def test_fetch_model_providers(self, mocker):
def test_fetch_model_providers(self, mocker: MockerFixture):
client = PluginModelClient()
request_mock = mocker.patch.object(client, "_request_with_plugin_daemon_response", return_value=["provider-a"])
@@ -23,7 +24,7 @@ class TestPluginModelClient:
)
assert request_mock.call_args.kwargs["params"] == {"page": 1, "page_size": 256}
def test_get_model_schema(self, mocker):
def test_get_model_schema(self, mocker: MockerFixture):
client = PluginModelClient()
schema = SimpleNamespace(name="schema")
stream_mock = mocker.patch.object(
@@ -45,7 +46,7 @@ class TestPluginModelClient:
assert result is schema
assert stream_mock.call_args.args[:2] == ("POST", "plugin/tenant-1/dispatch/model/schema")
def test_get_model_schema_empty_stream_returns_none(self, mocker):
def test_get_model_schema_empty_stream_returns_none(self, mocker: MockerFixture):
client = PluginModelClient()
mocker.patch.object(client, "_request_with_plugin_daemon_response_stream", return_value=iter([]))
@@ -53,7 +54,7 @@ class TestPluginModelClient:
assert result is None
def test_validate_provider_credentials(self, mocker):
def test_validate_provider_credentials(self, mocker: MockerFixture):
client = PluginModelClient()
stream_mock = mocker.patch.object(
client,
@@ -77,7 +78,7 @@ class TestPluginModelClient:
"plugin/tenant-1/dispatch/model/validate_provider_credentials",
)
def test_validate_provider_credentials_without_dict_update(self, mocker):
def test_validate_provider_credentials_without_dict_update(self, mocker: MockerFixture):
client = PluginModelClient()
mocker.patch.object(
client,
@@ -91,13 +92,13 @@ class TestPluginModelClient:
assert result is False
assert credentials == {"api_key": "same"}
def test_validate_provider_credentials_empty_returns_false(self, mocker):
def test_validate_provider_credentials_empty_returns_false(self, mocker: MockerFixture):
client = PluginModelClient()
mocker.patch.object(client, "_request_with_plugin_daemon_response_stream", return_value=iter([]))
assert client.validate_provider_credentials("tenant-1", "user-1", "org/plugin:1", "provider-a", {}) is False
def test_validate_model_credentials(self, mocker):
def test_validate_model_credentials(self, mocker: MockerFixture):
client = PluginModelClient()
stream_mock = mocker.patch.object(
client,
@@ -123,7 +124,7 @@ class TestPluginModelClient:
"plugin/tenant-1/dispatch/model/validate_model_credentials",
)
def test_validate_model_credentials_empty_returns_false(self, mocker):
def test_validate_model_credentials_empty_returns_false(self, mocker: MockerFixture):
client = PluginModelClient()
mocker.patch.object(client, "_request_with_plugin_daemon_response_stream", return_value=iter([]))
@@ -132,7 +133,7 @@ class TestPluginModelClient:
is False
)
def test_invoke_llm(self, mocker):
def test_invoke_llm(self, mocker: MockerFixture):
client = PluginModelClient()
stream_mock = mocker.patch.object(
client, "_request_with_plugin_daemon_response_stream", return_value=iter(["chunk-1"])
@@ -160,7 +161,7 @@ class TestPluginModelClient:
assert call_kwargs["data"]["data"]["stream"] is False
assert call_kwargs["data"]["data"]["model_parameters"] == {"temperature": 0.1}
def test_invoke_llm_wraps_plugin_daemon_inner_error(self, mocker):
def test_invoke_llm_wraps_plugin_daemon_inner_error(self, mocker: MockerFixture):
client = PluginModelClient()
def _boom():
@@ -182,7 +183,7 @@ class TestPluginModelClient:
)
)
def test_get_llm_num_tokens(self, mocker):
def test_get_llm_num_tokens(self, mocker: MockerFixture):
client = PluginModelClient()
mocker.patch.object(
client,
@@ -204,7 +205,7 @@ class TestPluginModelClient:
assert result == 42
def test_get_llm_num_tokens_empty_returns_zero(self, mocker):
def test_get_llm_num_tokens_empty_returns_zero(self, mocker: MockerFixture):
client = PluginModelClient()
mocker.patch.object(client, "_request_with_plugin_daemon_response_stream", return_value=iter([]))
@@ -213,7 +214,7 @@ class TestPluginModelClient:
== 0
)
def test_invoke_text_embedding(self, mocker):
def test_invoke_text_embedding(self, mocker: MockerFixture):
client = PluginModelClient()
embedding_result = SimpleNamespace(data=[[0.1, 0.2]])
mocker.patch.object(
@@ -233,7 +234,7 @@ class TestPluginModelClient:
assert result is embedding_result
def test_invoke_text_embedding_empty_raises(self, mocker):
def test_invoke_text_embedding_empty_raises(self, mocker: MockerFixture):
client = PluginModelClient()
mocker.patch.object(client, "_request_with_plugin_daemon_response_stream", return_value=iter([]))
@@ -242,7 +243,7 @@ class TestPluginModelClient:
"tenant-1", "user-1", "org/plugin:1", "provider-a", "embedding-a", {}, ["hello"], "x"
)
def test_invoke_multimodal_embedding(self, mocker):
def test_invoke_multimodal_embedding(self, mocker: MockerFixture):
client = PluginModelClient()
embedding_result = SimpleNamespace(data=[[0.3, 0.4]])
mocker.patch.object(
@@ -262,7 +263,7 @@ class TestPluginModelClient:
assert result is embedding_result
def test_invoke_multimodal_embedding_empty_raises(self, mocker):
def test_invoke_multimodal_embedding_empty_raises(self, mocker: MockerFixture):
client = PluginModelClient()
mocker.patch.object(client, "_request_with_plugin_daemon_response_stream", return_value=iter([]))
@@ -271,7 +272,7 @@ class TestPluginModelClient:
"tenant-1", "user-1", "org/plugin:1", "provider-a", "embedding-a", {}, [{"type": "image"}], "x"
)
def test_get_text_embedding_num_tokens(self, mocker):
def test_get_text_embedding_num_tokens(self, mocker: MockerFixture):
client = PluginModelClient()
mocker.patch.object(
client,
@@ -287,7 +288,7 @@ class TestPluginModelClient:
3,
]
def test_get_text_embedding_num_tokens_empty_returns_list(self, mocker):
def test_get_text_embedding_num_tokens_empty_returns_list(self, mocker: MockerFixture):
client = PluginModelClient()
mocker.patch.object(client, "_request_with_plugin_daemon_response_stream", return_value=iter([]))
@@ -298,7 +299,7 @@ class TestPluginModelClient:
== []
)
def test_invoke_rerank(self, mocker):
def test_invoke_rerank(self, mocker: MockerFixture):
client = PluginModelClient()
rerank_result = SimpleNamespace(scores=[0.9])
mocker.patch.object(client, "_request_with_plugin_daemon_response_stream", return_value=iter([rerank_result]))
@@ -318,14 +319,14 @@ class TestPluginModelClient:
assert result is rerank_result
def test_invoke_rerank_empty_raises(self, mocker):
def test_invoke_rerank_empty_raises(self, mocker: MockerFixture):
client = PluginModelClient()
mocker.patch.object(client, "_request_with_plugin_daemon_response_stream", return_value=iter([]))
with pytest.raises(ValueError, match="Failed to invoke rerank"):
client.invoke_rerank("tenant-1", "user-1", "org/plugin:1", "provider-a", "rerank-a", {}, "q", ["doc-1"])
def test_invoke_multimodal_rerank(self, mocker):
def test_invoke_multimodal_rerank(self, mocker: MockerFixture):
client = PluginModelClient()
rerank_result = SimpleNamespace(scores=[0.8])
mocker.patch.object(client, "_request_with_plugin_daemon_response_stream", return_value=iter([rerank_result]))
@@ -345,7 +346,7 @@ class TestPluginModelClient:
assert result is rerank_result
def test_invoke_multimodal_rerank_empty_raises(self, mocker):
def test_invoke_multimodal_rerank_empty_raises(self, mocker: MockerFixture):
client = PluginModelClient()
mocker.patch.object(client, "_request_with_plugin_daemon_response_stream", return_value=iter([]))
@@ -361,7 +362,7 @@ class TestPluginModelClient:
[{"type": "image"}],
)
def test_invoke_tts(self, mocker):
def test_invoke_tts(self, mocker: MockerFixture):
client = PluginModelClient()
mocker.patch.object(
client,
@@ -384,7 +385,7 @@ class TestPluginModelClient:
assert result == [b"hello", b"!"]
def test_invoke_tts_wraps_plugin_daemon_inner_error(self, mocker):
def test_invoke_tts_wraps_plugin_daemon_inner_error(self, mocker: MockerFixture):
client = PluginModelClient()
def _boom():
@@ -396,7 +397,7 @@ class TestPluginModelClient:
with pytest.raises(ValueError, match="tts error-400"):
list(client.invoke_tts("tenant-1", "user-1", "org/plugin:1", "provider-a", "tts-a", {}, "hello", "alloy"))
def test_get_tts_model_voices(self, mocker):
def test_get_tts_model_voices(self, mocker: MockerFixture):
client = PluginModelClient()
mocker.patch.object(
client,
@@ -425,13 +426,13 @@ class TestPluginModelClient:
assert result == [{"name": "Alloy", "value": "alloy"}, {"name": "Echo", "value": "echo"}]
def test_get_tts_model_voices_empty_returns_list(self, mocker):
def test_get_tts_model_voices_empty_returns_list(self, mocker: MockerFixture):
client = PluginModelClient()
mocker.patch.object(client, "_request_with_plugin_daemon_response_stream", return_value=iter([]))
assert client.get_tts_model_voices("tenant-1", "user-1", "org/plugin:1", "provider-a", "tts-a", {}) == []
def test_invoke_speech_to_text(self, mocker):
def test_invoke_speech_to_text(self, mocker: MockerFixture):
client = PluginModelClient()
stream_mock = mocker.patch.object(
client,
@@ -452,7 +453,7 @@ class TestPluginModelClient:
assert result == "transcribed text"
assert stream_mock.call_args.kwargs["data"]["data"]["file"] == "616263"
def test_invoke_speech_to_text_empty_raises(self, mocker):
def test_invoke_speech_to_text_empty_raises(self, mocker: MockerFixture):
client = PluginModelClient()
mocker.patch.object(client, "_request_with_plugin_daemon_response_stream", return_value=iter([]))
@@ -461,7 +462,7 @@ class TestPluginModelClient:
"tenant-1", "user-1", "org/plugin:1", "provider-a", "stt-a", {}, io.BytesIO(b"abc")
)
def test_invoke_moderation(self, mocker):
def test_invoke_moderation(self, mocker: MockerFixture):
client = PluginModelClient()
stream_mock = mocker.patch.object(
client,
@@ -482,7 +483,7 @@ class TestPluginModelClient:
assert result is True
assert stream_mock.call_args.kwargs["path"] == "plugin/tenant-1/dispatch/moderation/invoke"
def test_invoke_moderation_empty_raises(self, mocker):
def test_invoke_moderation_empty_raises(self, mocker: MockerFixture):
client = PluginModelClient()
mocker.patch.object(client, "_request_with_plugin_daemon_response_stream", return_value=iter([]))

Some files were not shown because too many files have changed in this diff Show More