mirror of
https://github.com/langgenius/dify.git
synced 2026-09-24 23:22:26 +08:00
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:
co-authored by
autofix-ci[bot]
parent
e03eb3a76c
commit
140ad6ba4e
+4
-4
@@ -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."""
|
||||
|
||||
|
||||
+3
-2
@@ -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(
|
||||
|
||||
+3
-1
@@ -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,
|
||||
|
||||
+1
-1
@@ -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()
|
||||
|
||||
+43
-28
@@ -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
|
||||
|
||||
+4
-3
@@ -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)
|
||||
|
||||
+5
-5
@@ -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",
|
||||
|
||||
+2
-1
@@ -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(
|
||||
|
||||
+3
-3
@@ -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()
|
||||
|
||||
|
||||
+2
-1
@@ -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()
|
||||
|
||||
+3
-2
@@ -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)
|
||||
|
||||
+2
-2
@@ -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",
|
||||
|
||||
+2
-2
@@ -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
|
||||
|
||||
+6
-5
@@ -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(
|
||||
|
||||
+11
-8
@@ -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()
|
||||
|
||||
+16
-9
@@ -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).
|
||||
|
||||
+14
-13
@@ -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 = {
|
||||
|
||||
+5
-4
@@ -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
|
||||
|
||||
+4
-3
@@ -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(
|
||||
|
||||
+4
-4
@@ -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"],
|
||||
|
||||
+11
-10
@@ -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}
|
||||
|
||||
|
||||
+13
-13
@@ -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
Reference in New Issue
Block a user