From f0078b793589bc8002dfcb3bf6dc08d01b786262 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E9=9D=9E=E6=B3=95=E6=93=8D=E4=BD=9C?= Date: Wed, 23 Sep 2026 08:25:37 +0000 Subject: [PATCH] refactor(api): migrate installed app workflow and task control (#41912) --- api/controllers/console/explore/completion.py | 5 +- api/controllers/console/explore/workflow.py | 84 ++- api/core/app/apps/base_app_queue_manager.py | 11 +- api/core/app/apps/execution_coordinator.py | 7 +- api/services/app_task_service.py | 51 +- .../installed_app_generation_service.py | 20 +- .../explore/test_installed_app_admission.py | 140 +++-- .../console/explore/test_workflow.py | 544 +++++++++++++----- .../test_ext_application_services.py | 26 + .../services/test_app_task_service.py | 326 +++++++++-- .../test_installed_app_generation_adapters.py | 109 +++- 11 files changed, 990 insertions(+), 333 deletions(-) diff --git a/api/controllers/console/explore/completion.py b/api/controllers/console/explore/completion.py index 491342f8c36..733db9ea5f2 100644 --- a/api/controllers/console/explore/completion.py +++ b/api/controllers/console/explore/completion.py @@ -36,7 +36,6 @@ from machinery.context import RequestContext from models.model import AppMode from services.account_errors import AccountNotFoundError from services.app_definition_query_service import AppDefinitionUnavailableError -from services.app_task_service import AppTaskService from services.errors.llm import InvokeRateLimitError from services.installed_app_access_service import InstalledAppNotFoundError, InstalledAppRef from services.installed_app_generation_service import InstalledAppNotChatError, InstalledAppNotCompletionError @@ -155,7 +154,7 @@ class CompletionStopApi(Resource): if app_mode != AppMode.COMPLETION: raise NotCompletionAppError() - AppTaskService.stop_task( + application_services().app_tasks.stop_task( task_id=task_id, invoke_from=InvokeFrom.EXPLORE, user_id=request_context.account_id, @@ -241,7 +240,7 @@ class ChatStopApi(Resource): if app_mode not in {AppMode.CHAT, AppMode.AGENT_CHAT, AppMode.ADVANCED_CHAT}: raise NotChatAppError() - AppTaskService.stop_task( + application_services().app_tasks.stop_task( task_id=task_id, invoke_from=InvokeFrom.EXPLORE, user_id=request_context.account_id, diff --git a/api/controllers/console/explore/workflow.py b/api/controllers/console/explore/workflow.py index 29e76d81498..4200ae798d7 100644 --- a/api/controllers/console/explore/workflow.py +++ b/api/controllers/console/explore/workflow.py @@ -1,7 +1,8 @@ import logging -from sqlalchemy.orm import Session -from werkzeug.exceptions import InternalServerError +from flask import Response +from flask_restx import Resource +from werkzeug.exceptions import InternalServerError, Unauthorized from controllers.common.controller_schemas import WorkflowRunPayload from controllers.common.fields import SimpleResultResponse @@ -12,26 +13,25 @@ from controllers.console.app.error import ( ProviderNotInitializeError, ProviderQuotaExceededError, ) -from controllers.console.app.wraps import with_session from controllers.console.explore.error import NotWorkflowAppError -from controllers.console.explore.wraps import InstalledAppResource -from controllers.console.wraps import model_validate, with_current_user +from controllers.console.explore.installed_app_admission import get_installed_app +from controllers.console.flask_admission import console_account_admission +from controllers.console.wraps import model_validate from controllers.web.error import InvokeRateLimitError as InvokeRateLimitHttpError -from core.app.apps.base_app_queue_manager import AppQueueManager -from core.app.entities.app_invoke_entities import InvokeFrom from core.errors.error import ( ModelCurrentlyNotSupportError, ProviderTokenNotInitError, QuotaExceededError, ) -from extensions.ext_redis import redis_client -from graphon.graph_engine.manager import GraphEngineManager +from extensions.ext_application_services import application_services from graphon.model_runtime.errors.invoke import InvokeError from libs import helper -from models import Account -from models.model import AppMode, InstalledApp -from services.app_generate_service import AppGenerateService +from machinery.context import RequestContext +from services.account_errors import AccountNotFoundError +from services.app_definition_query_service import AppDefinitionUnavailableError from services.errors.llm import InvokeRateLimitError +from services.installed_app_access_service import InstalledAppRef +from services.installed_app_generation_service import InstalledAppNotWorkflowError from .. import console_ns @@ -42,42 +42,34 @@ register_response_schema_models(console_ns, SimpleResultResponse) @console_ns.route("/installed-apps//workflows/run") -class InstalledAppWorkflowRunApi(InstalledAppResource): +class InstalledAppWorkflowRunApi(Resource): @console_ns.expect(console_ns.models[WorkflowRunPayload.__name__]) @console_ns.response(200, "Success") - @with_current_user - @with_session + @console_account_admission() + @get_installed_app @model_validate(WorkflowRunPayload) def post( self, req_data: WorkflowRunPayload, - session: Session, - current_user: Account, - installed_app: InstalledApp, - ): + request_context: RequestContext, + installed_app: InstalledAppRef, + ) -> Response: """ Run workflow """ - app_model = installed_app.app_with_session(session=session) - if not app_model: - raise NotWorkflowAppError() - app_mode = AppMode.value_of(app_model.mode) - if app_mode != AppMode.WORKFLOW: - raise NotWorkflowAppError() - - args = req_data.model_dump(exclude_none=True) try: - response = AppGenerateService.generate( - session=session, - app_model=app_model, - user=current_user, - args=args, - invoke_from=InvokeFrom.EXPLORE, - streaming=True, + response = application_services().installed_app_generation.generate_workflow( + installed_app=installed_app, + account_id=request_context.account_id, + args=req_data.model_dump(exclude_none=True), ) # response-contract:ignore compact_generate_response return helper.compact_generate_response(response) + except (AppDefinitionUnavailableError, InstalledAppNotWorkflowError): + raise NotWorkflowAppError() from None + except AccountNotFoundError: + raise Unauthorized("Account no longer exists.") from None except ProviderTokenNotInitError as ex: raise ProviderNotInitializeError(ex.description) except QuotaExceededError: @@ -96,25 +88,21 @@ class InstalledAppWorkflowRunApi(InstalledAppResource): @console_ns.route("/installed-apps//workflows/tasks//stop") -class InstalledAppWorkflowTaskStopApi(InstalledAppResource): +class InstalledAppWorkflowTaskStopApi(Resource): @console_ns.response(200, "Success", console_ns.models[SimpleResultResponse.__name__]) - @with_session(write=False) - def post(self, session: Session, installed_app: InstalledApp, task_id: str): + @console_account_admission() + @get_installed_app + def post(self, request_context: RequestContext, installed_app: InstalledAppRef, task_id: str) -> dict[str, object]: """ Stop workflow task """ - app_model = installed_app.app_with_session(session=session) - if not app_model: - raise NotWorkflowAppError() - app_mode = AppMode.value_of(app_model.mode) - if app_mode != AppMode.WORKFLOW: + try: + app_mode = application_services().app_definitions.get_mode(installed_app.app_id) + except AppDefinitionUnavailableError: + raise NotWorkflowAppError() from None + if app_mode != "workflow": raise NotWorkflowAppError() - # Stop using both mechanisms for backward compatibility - # Legacy stop flag mechanism (without user check) - AppQueueManager.set_stop_flag_no_user_check(task_id) - - # New graph engine command channel mechanism - GraphEngineManager(redis_client).send_stop_command(task_id) + application_services().app_tasks.stop_workflow_task_no_user_check(task_id=task_id) return SimpleResultResponse(result="success").model_dump(mode="json") diff --git a/api/core/app/apps/base_app_queue_manager.py b/api/core/app/apps/base_app_queue_manager.py index 651529ee81d..ea8d50b74ae 100644 --- a/api/core/app/apps/base_app_queue_manager.py +++ b/api/core/app/apps/base_app_queue_manager.py @@ -25,7 +25,7 @@ from core.app.entities.queue_entities import ( QueueStopEvent, WorkflowQueueMessage, ) -from extensions.ext_redis import redis_client +from extensions.ext_redis import RedisClientWrapper, redis_client from graphon.runtime import GraphRuntimeState logger = logging.getLogger(__name__) @@ -186,12 +186,15 @@ class AppQueueManager(ABC): raise NotImplementedError @classmethod - def set_stop_flag(cls, task_id: str, invoke_from: InvokeFrom, user_id: str): + def set_stop_flag( + cls, task_id: str, invoke_from: InvokeFrom, user_id: str, *, redis: RedisClientWrapper | None = None + ) -> None: """ Set task stop flag :return: """ - result: Any | None = redis_client.get(cls._generate_task_belong_cache_key(task_id)) + client = redis if redis is not None else redis_client + result: Any | None = client.get(cls._generate_task_belong_cache_key(task_id)) if result is None: return @@ -200,7 +203,7 @@ class AppQueueManager(ABC): return stopped_cache_key = cls._generate_stopped_cache_key(task_id) - redis_client.setex(stopped_cache_key, 600, 1) + client.setex(stopped_cache_key, 600, 1) @classmethod def set_stop_flag_no_user_check(cls, task_id: str) -> None: diff --git a/api/core/app/apps/execution_coordinator.py b/api/core/app/apps/execution_coordinator.py index e0fb9d36c65..f8ef6347f83 100644 --- a/api/core/app/apps/execution_coordinator.py +++ b/api/core/app/apps/execution_coordinator.py @@ -7,7 +7,7 @@ from collections.abc import Callable from enum import Enum, auto from configs import dify_config -from extensions.ext_redis import redis_client +from extensions.ext_redis import RedisClientWrapper, redis_client from graphon.graph_engine.command_channels import RedisChannel from graphon.graph_engine.manager import GraphEngineManager @@ -31,11 +31,12 @@ def app_task_stop_flag_key(task_id: str) -> str: return f"generate_task_stopped:{task_id}" -def set_app_task_stop_flag(task_id: str) -> None: +def set_app_task_stop_flag(task_id: str, *, redis: RedisClientWrapper | None = None) -> None: if not task_id: return - redis_client.setex(app_task_stop_flag_key(task_id), 600, 1) + client = redis if redis is not None else redis_client + client.setex(app_task_stop_flag_key(task_id), 600, 1) def is_app_task_stop_flag_set(task_id: str) -> bool: diff --git a/api/services/app_task_service.py b/api/services/app_task_service.py index d6e82226ee9..453ce4d9bc0 100644 --- a/api/services/app_task_service.py +++ b/api/services/app_task_service.py @@ -6,7 +6,7 @@ new GraphEngine command channel mechanism. """ from core.app.apps.base_app_queue_manager import AppQueueManager -from core.app.apps.execution_coordinator import app_task_stop_flag_key +from core.app.apps.execution_coordinator import set_app_task_stop_flag from core.app.entities.app_invoke_entities import InvokeFrom from extensions.ext_redis import RedisClientWrapper, redis_client from graphon.graph_engine.manager import GraphEngineManager @@ -14,27 +14,13 @@ from models.model import AppMode class AppTaskControlService: - """Injected task control for entry points that already admit app access.""" + """Service for managing application task operations.""" def __init__(self, *, redis_client: RedisClientWrapper) -> None: self._redis_client: RedisClientWrapper = redis_client - def stop_workflow_task_no_user_check(self, *, task_id: str) -> None: - """Send both cancellation signals after the caller has admitted the app. - - Preserve the legacy stop flag even if the GraphEngine command fails. - Trial workflow stops do not consult the task's user ownership cache. - """ - if task_id: - self._redis_client.setex(app_task_stop_flag_key(task_id), 600, 1) - GraphEngineManager(self._redis_client).send_stop_command(task_id) - - -class AppTaskService: - """Service for managing application task operations.""" - - @staticmethod def stop_task( + self, task_id: str, invoke_from: InvokeFrom, user_id: str, @@ -56,9 +42,36 @@ class AppTaskService: None """ # Legacy mechanism: Set stop flag in Redis - AppQueueManager.set_stop_flag(task_id, invoke_from, user_id) + AppQueueManager.set_stop_flag(task_id, invoke_from, user_id, redis=self._redis_client) # New mechanism: Send stop command via GraphEngine for workflow-based apps # This ensures proper workflow status recording in the persistence layer if app_mode in (AppMode.ADVANCED_CHAT, AppMode.WORKFLOW): - GraphEngineManager(redis_client).send_stop_command(task_id) + GraphEngineManager(self._redis_client).send_stop_command(task_id) + + def stop_workflow_task_no_user_check(self, *, task_id: str) -> None: + """Stop a workflow after app admission, without consulting the user ownership cache. + + Keep the legacy stop flag before the GraphEngine command, including when + the latter fails. Callers must authorize access to the workflow app first. + """ + set_app_task_stop_flag(task_id, redis=self._redis_client) + GraphEngineManager(self._redis_client).send_stop_command(task_id) + + +class AppTaskService: + """Compatibility entry point for callers outside ApplicationServices.""" + + @staticmethod + def stop_task( + task_id: str, + invoke_from: InvokeFrom, + user_id: str, + app_mode: AppMode, + ) -> None: + AppTaskControlService(redis_client=redis_client).stop_task( + task_id=task_id, + invoke_from=invoke_from, + user_id=user_id, + app_mode=app_mode, + ) diff --git a/api/services/installed_app_generation_service.py b/api/services/installed_app_generation_service.py index e1bf446a9af..381a4a57228 100644 --- a/api/services/installed_app_generation_service.py +++ b/api/services/installed_app_generation_service.py @@ -1,4 +1,4 @@ -"""Generate completion and chat responses for an admitted workspace installation.""" +"""Generate completion, chat, and workflow responses for an admitted installation.""" from collections.abc import Iterator, Mapping from datetime import datetime @@ -28,6 +28,10 @@ class InstalledAppNotChatError(ValueError): """The installed app does not support the chat endpoint.""" +class InstalledAppNotWorkflowError(ValueError): + """The installed app does not support the workflow endpoint.""" + + class InstalledAppUsageRecorder(Protocol): def record(self, *, installed_app: InstalledAppRef, used_at: datetime) -> None: ... @@ -76,6 +80,20 @@ class InstalledAppGenerationService: return self._generate(installed_app=installed_app, account_id=account_id, args=args, streaming=True) + def generate_workflow( + self, *, installed_app: InstalledAppRef, account_id: str, args: Mapping[str, object] + ) -> GenerationResponse: + if self._app_definitions.get_mode(installed_app.app_id) != "workflow": + raise InstalledAppNotWorkflowError(f"App {installed_app.app_id} is not a workflow app") + + # Workflow runs do not update installation usage or add chat naming options. + return self._runtime.generate( + app_id=installed_app.app_id, + account_id=account_id, + args=args, + streaming=True, + ) + def _generate( self, *, diff --git a/api/tests/unit_tests/controllers/console/explore/test_installed_app_admission.py b/api/tests/unit_tests/controllers/console/explore/test_installed_app_admission.py index be2aefeaff7..ac38975761e 100644 --- a/api/tests/unit_tests/controllers/console/explore/test_installed_app_admission.py +++ b/api/tests/unit_tests/controllers/console/explore/test_installed_app_admission.py @@ -1,6 +1,5 @@ import json from collections.abc import Callable, Generator -from contextlib import contextmanager from dataclasses import dataclass, field from operator import itemgetter from uuid import UUID, uuid4 @@ -21,6 +20,7 @@ import controllers.console.explore.saved_message as saved_message_module import controllers.console.flask_admission as console_admission import controllers.console.wraps as console_wraps import core.app.apps.base_app_queue_manager as app_queue_module +import core.app.apps.execution_coordinator as coordinator_module import libs.login as login_module import services.app_task_service as app_task_module from controllers.console.explore.installed_app_admission import get_installed_app @@ -34,9 +34,11 @@ from models.account import AccountStatus from repositories.app_definition_query_repository import AppDefinitionQueryRepository from repositories.installed_app_repository import SQLAlchemyInstalledAppRepository from services.app_definition_query_service import AppDefinitionQueryService +from services.app_task_service import AppTaskControlService from services.installed_app_access_service import InstalledAppAccessService, InstalledAppRef from services.saved_message_service import SavedMessageActor, SavedMessagePage, SavedMessageRecord, SavedMessageService from services.webapp_access_query_service import WebAppAccessUnavailableError +from tests.unit_tests.services.test_app_task_service import _StopRedis @dataclass @@ -508,70 +510,47 @@ def test_migrated_saved_message_and_parameter_handlers_dispatch_through_full_adm assert harness.state.permission_calls == [(harness.account.id, harness.target_app.id)] * 5 -@dataclass -class _StopRedis: - values: dict[str, bytes] = field(default_factory=dict) - commands: dict[str, list[str]] = field(default_factory=dict) - expirations: dict[str, int] = field(default_factory=dict) - reads: list[str] = field(default_factory=list) - read_error: Exception | None = None - - def get(self, key: str) -> bytes | None: - self.reads.append(key) - if self.read_error is not None: - raise self.read_error - return self.values.get(key) - - def setex(self, key: str, ttl: int, value: int) -> None: - self.values[key] = str(value).encode() - self.expirations[key] = ttl - - @contextmanager - def pipeline(self) -> Generator["_StopRedis"]: - yield self - - def rpush(self, key: str, value: str) -> int: - values = self.commands.setdefault(key, []) - values.append(value) - return len(values) - - def expire(self, key: str, ttl: int) -> bool: - self.expirations[key] = ttl - return True - - def set(self, key: str, value: str, *, ex: int) -> bool: - self.values[key] = value.encode() - self.expirations[key] = ex - return True - - def execute(self) -> list[object]: - return [] - - @dataclass(frozen=True) class _StopServices: app_definitions: AppDefinitionQueryService + app_tasks: AppTaskControlService _TASK_ID = "task-with-non-uuid-id" +@pytest.fixture +def _stop_global_redis(monkeypatch: pytest.MonkeyPatch) -> Generator[_StopRedis]: + redis = _StopRedis( + read_error=AssertionError("Must use the injected Redis for ownership reads"), + flag_error=AssertionError("Must use the injected Redis for stop flags"), + command_error=AssertionError("Must use the injected Redis for GraphEngine commands"), + ) + monkeypatch.setattr(app_queue_module, "redis_client", redis) + monkeypatch.setattr(coordinator_module, "redis_client", redis) + monkeypatch.setattr(app_task_module, "redis_client", redis) + yield redis + # GraphEngine catches Redis failures, so inspect the trap even after HTTP success. + assert redis.reads == [] + assert redis.operations == [] + + @pytest.fixture def stop_redis( harness: _Harness, monkeypatch: pytest.MonkeyPatch, sqlite_session_factory: sessionmaker[Session], + _stop_global_redis: _StopRedis, ) -> _StopRedis: redis = _StopRedis(values={f"generate_task_belong:{_TASK_ID}": f"account-{harness.account.id}".encode()}) services = _StopServices( app_definitions=AppDefinitionQueryService( definitions=AppDefinitionQueryRepository(session_factory=sqlite_session_factory), builtin_icon_url_prefix="/tools/icons", - ) + ), + app_tasks=AppTaskControlService(redis_client=redis), ) monkeypatch.setattr(completion_module, "application_services", lambda: services) - monkeypatch.setattr(app_queue_module, "redis_client", redis) - monkeypatch.setattr(app_task_module, "redis_client", redis) harness.api.add_resource( completion_module.CompletionStopApi, "/installed-apps//completion-messages//stop", @@ -629,6 +608,42 @@ def test_stop_handlers_preserve_mode_specific_commands_and_response( assert harness.state.permission_calls == [(harness.account.id, harness.target_app.id)] +@pytest.mark.parametrize( + ("message_kind", "mode"), [("completion", AppMode.COMPLETION), ("chat", AppMode.ADVANCED_CHAT)] +) +@pytest.mark.parametrize("ownership", ["missing", "different-account", "end-user"]) +def test_stop_handlers_preserve_mode_specific_behavior_when_task_ownership_does_not_match( + harness: _Harness, + stop_redis: _StopRedis, + sqlite_session_factory: sessionmaker[Session], + message_kind: str, + mode: AppMode, + ownership: str, +) -> None: + _set_app_mode(harness, sqlite_session_factory, mode) + owner_key = f"generate_task_belong:{_TASK_ID}" + if ownership == "missing": + stop_redis.values.pop(owner_key) + elif ownership == "different-account": + stop_redis.values[owner_key] = b"account-someone-else" + else: + stop_redis.values[owner_key] = f"end-user-{harness.account.id}".encode() + + response = harness.app.test_client().post(_stop_url(harness, message_kind)) + + _assert_json_response(response, status=200, body={"result": "success"}) + assert stop_redis.reads == [owner_key] + assert f"generate_task_stopped:{_TASK_ID}" not in stop_redis.values + if mode == AppMode.ADVANCED_CHAT: + assert stop_redis.operations == ["graph_command"] + assert [json.loads(command) for command in stop_redis.commands[f"workflow:{_TASK_ID}:commands"]] == [ + {"command_type": "abort", "payload": None, "reason": "User requested stop"} + ] + else: + assert stop_redis.operations == [] + assert stop_redis.commands == {} + + @pytest.mark.parametrize( ("message_kind", "mode", "code", "message"), [ @@ -706,17 +721,29 @@ def test_stop_handlers_enforce_admission_before_sending_commands( assert stop_redis.commands == {} -@pytest.mark.parametrize(("message_kind", "mode"), [("completion", AppMode.COMPLETION), ("chat", AppMode.CHAT)]) +@pytest.mark.parametrize( + ("message_kind", "mode", "failure_stage"), + [ + ("completion", AppMode.COMPLETION, "read"), + ("chat", AppMode.CHAT, "read"), + ("completion", AppMode.COMPLETION, "flag"), + ("chat", AppMode.ADVANCED_CHAT, "flag"), + ], +) def test_stop_handlers_propagate_redis_failure_to_existing_http_error_handler( harness: _Harness, stop_redis: _StopRedis, sqlite_session_factory: sessionmaker[Session], message_kind: str, mode: AppMode, + failure_stage: str, ) -> None: _set_app_mode(harness, sqlite_session_factory, mode) failure = RedisConnectionError("Redis unavailable") - stop_redis.read_error = failure + if failure_stage == "read": + stop_redis.read_error = failure + else: + stop_redis.flag_error = failure exceptions: list[Exception] = [] def capture_exception(_sender: Flask, exception: Exception) -> None: @@ -733,3 +760,24 @@ def test_stop_handlers_propagate_redis_failure_to_existing_http_error_handler( assert any(exception is failure for exception in exceptions) assert f"generate_task_stopped:{_TASK_ID}" not in stop_redis.values assert stop_redis.commands == {} + assert stop_redis.operations == (["legacy_flag"] if failure_stage == "flag" else []) + + +def test_chat_stop_preserves_success_and_legacy_flag_when_graph_redis_fails( + harness: _Harness, + stop_redis: _StopRedis, + sqlite_session_factory: sessionmaker[Session], + caplog: pytest.LogCaptureFixture, +) -> None: + _set_app_mode(harness, sqlite_session_factory, AppMode.ADVANCED_CHAT) + failure = RedisConnectionError("Graph channel unavailable") + stop_redis.command_error = failure + + response = harness.app.test_client().post(_stop_url(harness, "chat")) + + _assert_json_response(response, status=200, body={"result": "success"}) + assert stop_redis.operations == ["legacy_flag", "graph_command"] + assert stop_redis.values[f"generate_task_stopped:{_TASK_ID}"] == b"1" + assert stop_redis.expirations[f"generate_task_stopped:{_TASK_ID}"] == 600 + assert stop_redis.commands == {} + assert any(record.exc_info is not None and record.exc_info[1] is failure for record in caplog.records) diff --git a/api/tests/unit_tests/controllers/console/explore/test_workflow.py b/api/tests/unit_tests/controllers/console/explore/test_workflow.py index 835b316dab7..303f925ec2f 100644 --- a/api/tests/unit_tests/controllers/console/explore/test_workflow.py +++ b/api/tests/unit_tests/controllers/console/explore/test_workflow.py @@ -1,150 +1,434 @@ -from inspect import unwrap -from unittest.mock import MagicMock, patch +import json +from collections.abc import Generator, Mapping +from dataclasses import dataclass, field +from datetime import datetime import pytest -from flask import Flask -from sqlalchemy.orm import Session -from werkzeug.exceptions import InternalServerError +from flask import Flask, got_request_exception +from redis.exceptions import ConnectionError as RedisConnectionError +from sqlalchemy.orm import Session, sessionmaker -from controllers.common.controller_schemas import WorkflowRunPayload -from controllers.console.explore.error import NotWorkflowAppError -from controllers.console.explore.workflow import ( - InstalledAppWorkflowRunApi, - InstalledAppWorkflowTaskStopApi, -) -from controllers.web.error import InvokeRateLimitError as InvokeRateLimitHttpError -from models import Account -from models.model import App, AppMode, InstalledApp +import controllers.console.explore.workflow as workflow_module +from core.errors.error import ModelCurrentlyNotSupportError, ProviderTokenNotInitError, QuotaExceededError +from graphon.model_runtime.errors.invoke import InvokeError +from models import App, AppMode, InstalledApp, Tenant +from repositories.app_definition_query_repository import AppDefinitionQueryRepository +from repositories.installed_app_repository import SQLAlchemyInstalledAppRepository +from services.account_errors import AccountNotFoundError +from services.app_definition_query_service import AppDefinitionQueryService, AppDefinitionUnavailableError +from services.app_task_service import AppTaskControlService from services.errors.llm import InvokeRateLimitError +from services.installed_app_generation_service import GenerationResponse, InstalledAppGenerationService +from tests.unit_tests.controllers.console.explore.test_installed_app_admission import ( + _assert_json_response, + _Harness, + _set_app_mode, + harness, +) +from tests.unit_tests.services.test_app_task_service import _StopRedis + +__all__ = ["harness"] + +_LAST_USED_AT = datetime(2026, 9, 1, 12, 0, 0) +_TASK_ID = "workflow-task-with-non-uuid-id" +_STOP_KEY = f"generate_task_stopped:{_TASK_ID}" +_COMMAND_KEY = f"workflow:{_TASK_ID}:commands" + + +@dataclass(frozen=True) +class _RuntimeCall: + app_id: str + account_id: str + args: Mapping[str, object] + streaming: bool + + +@dataclass +class _Runtime: + session_factory: sessionmaker[Session] + installed_app_id: str + response: GenerationResponse = field( + default_factory=lambda: {"workflow_run_id": "run-1", "data": {"total_tokens": 0}} + ) + error: Exception | None = None + last_used_at: datetime | None = _LAST_USED_AT + calls: list[_RuntimeCall] = field(default_factory=list) + + def generate( + self, *, app_id: str, account_id: str, args: Mapping[str, object], streaming: bool + ) -> GenerationResponse: + # Workflow run never updates usage, including before a failed generation. + self.assert_usage_unchanged() + self.calls.append(_RuntimeCall(app_id, account_id, dict(args), streaming)) + if self.error is not None: + raise self.error + return self.response + + def assert_usage_unchanged(self) -> None: + with self.session_factory() as session: + installation = session.get(InstalledApp, self.installed_app_id) + assert installation is not None + assert installation.last_used_at == self.last_used_at + + +@dataclass(frozen=True) +class _Services: + installed_app_generation: InstalledAppGenerationService + app_definitions: AppDefinitionQueryService + app_tasks: AppTaskControlService @pytest.fixture -def app(): - app = Flask(__name__) - app.config["TESTING"] = True - return app - - -def make_installed_app(session: Session, *, mode: AppMode) -> InstalledApp: - app = App( - tenant_id="owner-tenant", - name="Explore App", - mode=mode, - enable_site=True, - enable_api=False, +def runtime( + harness: _Harness, + monkeypatch: pytest.MonkeyPatch, + sqlite_session_factory: sessionmaker[Session], + stop_redis: _StopRedis, +) -> _Runtime: + _set_app_mode(harness, sqlite_session_factory, AppMode.WORKFLOW) + with sqlite_session_factory.begin() as session: + installation = session.get(InstalledApp, harness.installed_app.id) + assert installation is not None + installation.last_used_at = _LAST_USED_AT + runtime = _Runtime(sqlite_session_factory, harness.installed_app.id) + definitions = AppDefinitionQueryService( + definitions=AppDefinitionQueryRepository(session_factory=sqlite_session_factory), + builtin_icon_url_prefix="/tools/icons", ) - session.add(app) - session.flush() - installed_app = InstalledApp( - tenant_id="viewer-tenant", - app_id=app.id, - app_owner_tenant_id=app.tenant_id, - position=0, - is_pinned=False, - last_used_at=None, + services = _Services( + installed_app_generation=InstalledAppGenerationService( + app_definitions=definitions, + usage=SQLAlchemyInstalledAppRepository(session_factory=sqlite_session_factory), + runtime=runtime, + ), + app_definitions=definitions, + app_tasks=AppTaskControlService(redis_client=stop_redis), ) - session.add(installed_app) - session.commit() - return installed_app + monkeypatch.setattr(workflow_module, "application_services", lambda: services) + harness.api.add_resource( + workflow_module.InstalledAppWorkflowRunApi, + "/installed-apps//workflows/run", + ) + harness.api.add_resource( + workflow_module.InstalledAppWorkflowTaskStopApi, + "/installed-apps//workflows/tasks//stop", + ) + return runtime @pytest.fixture -def payload(): - return {"inputs": {"a": 1}} +def stop_redis() -> _StopRedis: + return _StopRedis( + values={f"generate_task_belong:{_TASK_ID}": b"account-someone-else"}, + read_error=AssertionError("Workflow stop must not inspect task ownership"), + ) -class TestInstalledAppWorkflowRunApi: - def test_not_workflow_app(self, app: Flask, sqlite_session: Session): - api = InstalledAppWorkflowRunApi() - method = unwrap(api.post) - installed_app = make_installed_app(sqlite_session, mode=AppMode.CHAT) - user = Account(name="User", email="user@example.com") - - with app.test_request_context("/"): - with pytest.raises(NotWorkflowAppError): - method( - api, - WorkflowRunPayload.model_validate({"inputs": {}}), - sqlite_session, - user, - installed_app, - ) - - def test_success(self, app: Flask, sqlite_session: Session, payload): - api = InstalledAppWorkflowRunApi() - method = unwrap(api.post) - req_data = WorkflowRunPayload.model_validate(payload) - installed_app = make_installed_app(sqlite_session, mode=AppMode.WORKFLOW) - user = Account(name="User", email="user@example.com") - - with ( - app.test_request_context("/", json=payload), - patch( - "controllers.console.explore.workflow.AppGenerateService.generate", - return_value=MagicMock(), - ) as generate_mock, - ): - result = method(api, req_data, sqlite_session, user, installed_app) - - generate_mock.assert_called_once() - assert generate_mock.call_args.kwargs["user"] is user - assert result is not None - - def test_rate_limit_error(self, app: Flask, sqlite_session: Session, payload): - api = InstalledAppWorkflowRunApi() - method = unwrap(api.post) - req_data = WorkflowRunPayload.model_validate(payload) - installed_app = make_installed_app(sqlite_session, mode=AppMode.WORKFLOW) - user = Account(name="User", email="user@example.com") - - with ( - app.test_request_context("/", json=payload), - patch( - "controllers.console.explore.workflow.AppGenerateService.generate", - side_effect=InvokeRateLimitError("rate limit"), - ), - ): - with pytest.raises(InvokeRateLimitHttpError): - method(api, req_data, sqlite_session, user, installed_app) - - def test_unexpected_exception(self, app: Flask, sqlite_session: Session, payload): - api = InstalledAppWorkflowRunApi() - method = unwrap(api.post) - req_data = WorkflowRunPayload.model_validate(payload) - installed_app = make_installed_app(sqlite_session, mode=AppMode.WORKFLOW) - user = Account(name="User", email="user@example.com") - - with ( - app.test_request_context("/", json=payload), - patch( - "controllers.console.explore.workflow.AppGenerateService.generate", - side_effect=Exception("boom"), - ), - ): - with pytest.raises(InternalServerError): - method(api, req_data, sqlite_session, user, installed_app) +def _url(harness: _Harness, action: str = "run") -> str: + suffix = "run" if action == "run" else f"tasks/{_TASK_ID}/stop" + return f"/installed-apps/{harness.installed_app.id}/workflows/{suffix}" -class TestInstalledAppWorkflowTaskStopApi: - def test_not_workflow_app(self, sqlite_session: Session): - api = InstalledAppWorkflowTaskStopApi() - method = unwrap(api.post) - installed_app = make_installed_app(sqlite_session, mode=AppMode.CHAT) +@pytest.mark.parametrize("last_used_at", [None, _LAST_USED_AT]) +@pytest.mark.parametrize( + ("payload", "expected_args"), + [ + ( + {"inputs": {"zero": 0, "enabled": False, "empty": [], "nullable": None}}, + {"inputs": {"zero": 0, "enabled": False, "empty": [], "nullable": None}}, + ), + ({"inputs": {}, "files": None}, {"inputs": {}}), + ( + {"inputs": {}, "files": [], "response_mode": "blocking", "auto_generate_name": True}, + {"inputs": {}, "files": []}, + ), + ( + { + "inputs": {}, + "files": [{"type": "image", "transfer_method": "remote_url", "url": "https://example.com/i.png"}], + }, + { + "inputs": {}, + "files": [{"type": "image", "transfer_method": "remote_url", "url": "https://example.com/i.png"}], + }, + ), + ], +) +def test_workflow_run_preserves_args_and_response_without_updating_usage( + harness: _Harness, + runtime: _Runtime, + sqlite_session_factory: sessionmaker[Session], + payload: dict[str, object], + expected_args: dict[str, object], + last_used_at: datetime | None, +) -> None: + runtime.last_used_at = last_used_at + with sqlite_session_factory.begin() as session: + installation = session.get(InstalledApp, harness.installed_app.id) + assert installation is not None + installation.last_used_at = last_used_at - with pytest.raises(NotWorkflowAppError): - method(api, sqlite_session, installed_app, "task-1") + response = harness.app.test_client().post(_url(harness), json=payload) - def test_success(self, sqlite_session: Session): - api = InstalledAppWorkflowTaskStopApi() - method = unwrap(api.post) - installed_app = make_installed_app(sqlite_session, mode=AppMode.WORKFLOW) + assert response.status_code == 200 + assert response.get_json() == {"workflow_run_id": "run-1", "data": {"total_tokens": 0}} + assert dict(response.headers) == { + "Content-Type": "application/json; charset=utf-8", + "Content-Length": str(len(response.data)), + } + assert runtime.calls == [_RuntimeCall(harness.target_app.id, harness.account.id, expected_args, True)] + runtime.assert_usage_unchanged() - with ( - patch("controllers.console.explore.workflow.AppQueueManager.set_stop_flag_no_user_check") as stop_flag, - patch("controllers.console.explore.workflow.GraphEngineManager.send_stop_command") as send_stop, - ): - result = method(api, sqlite_session, installed_app, "task-1") - stop_flag.assert_called_once_with("task-1") - send_stop.assert_called_once_with("task-1") - assert result == {"result": "success"} +@pytest.mark.parametrize("consume_all", [False, True]) +def test_workflow_run_preserves_stream_bytes_headers_and_closes_on_completion_or_disconnect( + harness: _Harness, runtime: _Runtime, consume_all: bool +) -> None: + closed: list[bool] = [] + + def chunks() -> Generator[str]: + try: + yield 'data: {"event":"workflow_started"}\n\n' + yield 'data: {"event":"workflow_finished","answer":"你好"}\n\n' + finally: + closed.append(True) + + runtime.response = chunks() + response = harness.app.test_client().post(_url(harness), json={"inputs": {}}, buffered=False) + + assert response.status_code == 200 + assert dict(response.headers) == {"Content-Type": "text/event-stream; charset=utf-8"} + if consume_all: + assert ( + response.data + == ( + 'data: {"event":"workflow_started"}\n\ndata: {"event":"workflow_finished","answer":"你好"}\n\n' + ).encode() + ) + else: + assert next(iter(response.response)) == b'data: {"event":"workflow_started"}\n\n' + assert closed == [] + response.close() + assert closed == [True] + assert runtime.calls == [_RuntimeCall(harness.target_app.id, harness.account.id, {"inputs": {}}, True)] + runtime.assert_usage_unchanged() + + +@pytest.mark.parametrize( + ("failure", "status", "code", "message"), + [ + (ProviderTokenNotInitError("Missing credentials"), 400, "provider_not_initialize", "Missing credentials"), + ( + QuotaExceededError(), + 400, + "provider_quota_exceeded", + "Your quota for Dify Hosted Model Provider has been exhausted. " + "Please go to Settings -> Model Provider to complete your own provider credentials.", + ), + ( + ModelCurrentlyNotSupportError(), + 400, + "model_currently_not_support", + "Dify Hosted OpenAI trial currently not support the GPT-4 model.", + ), + (InvokeError("Provider rejected input"), 400, "completion_request_error", "Provider rejected input"), + (InvokeRateLimitError("Too many requests"), 429, "rate_limit_error", "Too many requests"), + (ValueError("Invalid runtime arguments"), 400, "invalid_param", "Invalid runtime arguments"), + (AccountNotFoundError(), 401, "unauthorized", "Account no longer exists."), + (AppDefinitionUnavailableError("App removed"), 400, "not_workflow_app", "Only support workflow app."), + ( + RuntimeError("Unexpected generation failure"), + 500, + "internal_server_error", + "The server encountered an internal error and was unable to complete your request. " + "Either the server is overloaded or there is an error in the application.", + ), + ], +) +def test_workflow_run_preserves_error_contract_without_updating_usage( + harness: _Harness, runtime: _Runtime, failure: Exception, status: int, code: str, message: str +) -> None: + runtime.error = failure + + response = harness.app.test_client().post(_url(harness), json={"inputs": {}}) + + _assert_json_response(response, status=status, body={"code": code, "message": message, "status": status}) + assert len(runtime.calls) == 1 + assert runtime.calls[0].streaming is True + runtime.assert_usage_unchanged() + if status == 401: + assert response.headers["WWW-Authenticate"] == 'Bearer realm="api"' + + +@pytest.mark.parametrize( + ("payload", "error"), + [ + ({}, {"type": "missing", "loc": ["inputs"], "msg": "Field required"}), + ({"inputs": None}, {"type": "dict_type", "loc": ["inputs"], "msg": "Input should be a valid dictionary"}), + ( + {"inputs": {}, "files": {}}, + {"type": "list_type", "loc": ["files"], "msg": "Input should be a valid list"}, + ), + ], +) +def test_workflow_run_validates_payload_before_mode( + harness: _Harness, + runtime: _Runtime, + sqlite_session_factory: sessionmaker[Session], + payload: dict[str, object], + error: dict[str, object], +) -> None: + _set_app_mode(harness, sqlite_session_factory, AppMode.CHAT) + + response = harness.app.test_client().post(_url(harness), json=payload) + + _assert_json_response( + response, + status=422, + body={"code": "unprocessable_entity", "message": json.dumps([error]), "status": 422}, + ) + assert runtime.calls == [] + runtime.assert_usage_unchanged() + + +@pytest.mark.parametrize("action", ["run", "stop"]) +@pytest.mark.parametrize("mode", [AppMode.COMPLETION, AppMode.CHAT, AppMode.ADVANCED_CHAT, AppMode.AGENT_CHAT]) +def test_workflow_handlers_reject_other_modes_without_side_effects( + harness: _Harness, + runtime: _Runtime, + stop_redis: _StopRedis, + sqlite_session_factory: sessionmaker[Session], + action: str, + mode: AppMode, +) -> None: + _set_app_mode(harness, sqlite_session_factory, mode) + + response = harness.app.test_client().post(_url(harness, action), json={"inputs": {}}) + + _assert_json_response( + response, status=400, body={"code": "not_workflow_app", "message": "Only support workflow app.", "status": 400} + ) + assert runtime.calls == [] + assert stop_redis.operations == [] + assert stop_redis.reads == [] + runtime.assert_usage_unchanged() + + +@pytest.mark.parametrize("action", ["run", "stop"]) +@pytest.mark.parametrize("rejection", ["permission", "tenant", "missing"]) +def test_workflow_handlers_require_admission_before_payload_or_task_actions( + harness: _Harness, + runtime: _Runtime, + stop_redis: _StopRedis, + sqlite_session_factory: sessionmaker[Session], + action: str, + rejection: str, +) -> None: + if rejection == "permission": + harness.state.allowed = False + elif rejection == "tenant": + harness.account._current_tenant = Tenant(name="Other workspace") + else: + with sqlite_session_factory.begin() as session: + installation = session.get(InstalledApp, harness.installed_app.id) + assert installation is not None + session.delete(installation) + + response = harness.app.test_client().post(_url(harness, action), json={}) + + if rejection == "permission": + _assert_json_response( + response, status=403, body={"code": "access_denied", "message": "App access denied.", "status": 403} + ) + else: + _assert_json_response( + response, status=404, body={"code": "not_found", "message": "Installed app not found", "status": 404} + ) + assert runtime.calls == [] + assert stop_redis.operations == [] + assert stop_redis.reads == [] + + +@pytest.mark.parametrize("action", ["run", "stop"]) +def test_workflow_handlers_preserve_not_workflow_error_when_app_disappears_after_admission( + harness: _Harness, + runtime: _Runtime, + stop_redis: _StopRedis, + sqlite_session_factory: sessionmaker[Session], + action: str, +) -> None: + def remove_app() -> None: + with sqlite_session_factory.begin() as session: + app = session.get(App, harness.target_app.id) + assert app is not None + session.delete(app) + + harness.state.permission_action = remove_app + + response = harness.app.test_client().post(_url(harness, action), json={"inputs": {}}) + + _assert_json_response( + response, status=400, body={"code": "not_workflow_app", "message": "Only support workflow app.", "status": 400} + ) + assert runtime.calls == [] + assert stop_redis.operations == [] + runtime.assert_usage_unchanged() + + +def test_workflow_stop_sets_both_signals_in_order_without_reading_task_ownership( + harness: _Harness, runtime: _Runtime, stop_redis: _StopRedis +) -> None: + response = harness.app.test_client().post(_url(harness, "stop")) + + _assert_json_response(response, status=200, body={"result": "success"}) + assert stop_redis.operations == ["legacy_flag", "graph_command"] + assert stop_redis.reads == [] + assert stop_redis.values[_STOP_KEY] == b"1" + assert stop_redis.expirations[_STOP_KEY] == 600 + assert set(stop_redis.commands) == {_COMMAND_KEY} + assert [json.loads(command) for command in stop_redis.commands[_COMMAND_KEY]] == [ + {"command_type": "abort", "payload": None, "reason": "User requested stop"} + ] + assert stop_redis.expirations[_COMMAND_KEY] == 3600 + assert runtime.calls == [] + runtime.assert_usage_unchanged() + + +def test_workflow_stop_preserves_legacy_redis_failure_without_attempting_graph_command( + harness: _Harness, runtime: _Runtime, stop_redis: _StopRedis +) -> None: + failure = RedisConnectionError("Legacy Redis unavailable") + stop_redis.flag_error = failure + exceptions: list[Exception] = [] + + def capture_exception(_sender: Flask, exception: Exception) -> None: + exceptions.append(exception) + + with got_request_exception.connected_to(capture_exception): + response = harness.app.test_client().post(_url(harness, "stop")) + + _assert_json_response( + response, status=500, body={"code": "unknown", "message": "Internal Server Error", "status": 500} + ) + assert any(exception is failure for exception in exceptions) + assert stop_redis.operations == ["legacy_flag"] + assert stop_redis.reads == [] + assert _STOP_KEY not in stop_redis.values + assert stop_redis.commands == {} + runtime.assert_usage_unchanged() + + +def test_workflow_stop_keeps_legacy_flag_and_success_when_graph_command_fails( + harness: _Harness, runtime: _Runtime, stop_redis: _StopRedis, caplog: pytest.LogCaptureFixture +) -> None: + failure = RedisConnectionError("Graph channel unavailable") + stop_redis.command_error = failure + + response = harness.app.test_client().post(_url(harness, "stop")) + + _assert_json_response(response, status=200, body={"result": "success"}) + assert stop_redis.operations == ["legacy_flag", "graph_command"] + assert stop_redis.reads == [] + assert stop_redis.values[_STOP_KEY] == b"1" + assert any(record.exc_info is not None and record.exc_info[1] is failure for record in caplog.records) + runtime.assert_usage_unchanged() diff --git a/api/tests/unit_tests/extensions/test_ext_application_services.py b/api/tests/unit_tests/extensions/test_ext_application_services.py index 03c57bf0d29..6f8d8b9f207 100644 --- a/api/tests/unit_tests/extensions/test_ext_application_services.py +++ b/api/tests/unit_tests/extensions/test_ext_application_services.py @@ -98,6 +98,7 @@ from services.workflow_app_log_query_service import WorkflowAppLogQueryService from services.workflow_run_service import WorkflowRunService from services.workflow_statistic_query_service import WorkflowStatisticQueryService from tests.unit_tests.config_override import apply_config_overrides +from tests.unit_tests.services.test_app_task_service import _StopRedis @pytest.mark.parametrize( @@ -690,6 +691,31 @@ def test_build_application_services_wires_data_source_api_key_auth( assert isinstance(services.data_source_api_key_auth, DataSourceApiKeyAuthService) +def test_build_application_services_uses_supplied_redis_for_both_workflow_stop_signals( + sqlite_session_factory: sessionmaker[Session], +) -> None: + redis = _StopRedis(read_error=AssertionError("Workflow stop must not inspect task ownership")) + script = MagicMock(side_effect=AssertionError("Workflow stop must not execute Redis scripts")) + with patch.object(redis, "register_script", return_value=script): + services = ext_application_services.build_application_services( + database_client=sqlite_session_factory, + deployment_edition=DeploymentEdition.COMMUNITY, + initialization_password="", + redis=redis, + ) + + services.app_tasks.stop_workflow_task_no_user_check(task_id="workflow-task") + + assert redis.reads == [] + assert redis.operations == ["legacy_flag", "graph_command"] + assert redis.values["generate_task_stopped:workflow-task"] == b"1" + assert redis.expirations["generate_task_stopped:workflow-task"] == 600 + assert [json.loads(command) for command in redis.commands["workflow:workflow-task:commands"]] == [ + {"command_type": "abort", "payload": None, "reason": "User requested stop"} + ] + assert redis.expirations["workflow:workflow-task:commands"] == 3600 + + def test_build_application_services_wires_trial_app_usage( sqlite_session_factory: sessionmaker[Session], ) -> None: diff --git a/api/tests/unit_tests/services/test_app_task_service.py b/api/tests/unit_tests/services/test_app_task_service.py index 06caa8170de..603dcc9509b 100644 --- a/api/tests/unit_tests/services/test_app_task_service.py +++ b/api/tests/unit_tests/services/test_app_task_service.py @@ -1,78 +1,282 @@ -from unittest.mock import patch +import json +from collections.abc import Generator +from contextlib import contextmanager +from dataclasses import dataclass, field +from datetime import timedelta +from functools import partial +from typing import override import pytest +from redis.exceptions import ConnectionError as RedisConnectionError +import core.app.apps.base_app_queue_manager as queue_module +import core.app.apps.execution_coordinator as coordinator_module +import services.app_task_service as task_module from core.app.entities.app_invoke_entities import InvokeFrom +from extensions.ext_redis import RedisClientWrapper from models.model import AppMode -from services.app_task_service import AppTaskService +from services.app_task_service import AppTaskControlService, AppTaskService + +_TASK_ID = "task-with-non-uuid-id" +_USER_ID = "user-1" +_OWNER_KEY = f"generate_task_belong:{_TASK_ID}" +_STOP_KEY = f"generate_task_stopped:{_TASK_ID}" +_COMMAND_KEY = f"workflow:{_TASK_ID}:commands" -class TestAppTaskService: - """Test suite for AppTaskService.stop_task method.""" +@dataclass +class _StopRedis(RedisClientWrapper): + """In-memory Redis boundary; task policy and GraphEngine serialization stay real.""" - @pytest.mark.parametrize( - ("app_mode", "should_call_graph_engine"), - [ - (AppMode.CHAT, False), - (AppMode.COMPLETION, False), - (AppMode.AGENT_CHAT, False), - (AppMode.AGENT, False), - (AppMode.CHANNEL, False), - (AppMode.RAG_PIPELINE, False), - (AppMode.ADVANCED_CHAT, True), - (AppMode.WORKFLOW, True), - ], + values: dict[str, bytes] = field(default_factory=dict) + commands: dict[str, list[str]] = field(default_factory=dict) + expirations: dict[str, int] = field(default_factory=dict) + reads: list[str] = field(default_factory=list) + operations: list[str] = field(default_factory=list) + read_error: Exception | None = None + flag_error: Exception | None = None + command_error: Exception | None = None + + @override + def get(self, name: str | bytes) -> bytes | None: + key = name.decode() if isinstance(name, bytes) else name + self.reads.append(key) + if self.read_error is not None: + raise self.read_error + return self.values.get(key) + + @override + def setex(self, name: str | bytes, time: int | timedelta, value: object) -> None: + self.operations.append("legacy_flag") + if self.flag_error is not None: + raise self.flag_error + key = name.decode() if isinstance(name, bytes) else name + self.values[key] = str(value).encode() + self.expirations[key] = int(time.total_seconds()) if isinstance(time, timedelta) else time + + @override + @contextmanager + def pipeline(self, transaction: bool = True, shard_hint: str | None = None) -> Generator["_StopPipeline"]: + yield _StopPipeline(self) + + +@dataclass +class _StopPipeline: + redis: _StopRedis + values: dict[str, bytes] = field(default_factory=dict) + commands: dict[str, list[str]] = field(default_factory=dict) + expirations: dict[str, int] = field(default_factory=dict) + + def rpush(self, name: str, value: str) -> None: + self.commands.setdefault(name, []).append(value) + + def expire(self, name: str, time: int) -> None: + self.expirations[name] = time + + def set(self, name: str, value: str, *, ex: int) -> None: + self.values[name] = value.encode() + self.expirations[name] = ex + + def execute(self) -> list[object]: + self.redis.operations.append("graph_command") + if self.redis.command_error is not None: + raise self.redis.command_error + self.redis.values.update(self.values) + self.redis.expirations.update(self.expirations) + for name, values in self.commands.items(): + self.redis.commands.setdefault(name, []).extend(values) + return [] + + +@pytest.fixture(autouse=True) +def global_redis(monkeypatch: pytest.MonkeyPatch) -> Generator[_StopRedis]: + redis = _StopRedis( + read_error=AssertionError("Must use the injected Redis for ownership reads"), + flag_error=AssertionError("Must use the injected Redis for stop flags"), + command_error=AssertionError("Must use the injected Redis for GraphEngine commands"), ) - @patch("services.app_task_service.AppQueueManager") - @patch("services.app_task_service.GraphEngineManager") - def test_stop_task_with_different_app_modes( - self, mock_graph_engine_manager, mock_app_queue_manager, app_mode, should_call_graph_engine - ): - """Test stop_task behavior with different app modes. + monkeypatch.setattr(queue_module, "redis_client", redis) + monkeypatch.setattr(coordinator_module, "redis_client", redis) + monkeypatch.setattr(task_module, "redis_client", redis) + yield redis + # GraphEngine catches Redis failures, so a trap exception alone would not fail the test. + assert redis.reads == [] + assert redis.operations == [] - Verifies that: - - Legacy Redis flag is always set via AppQueueManager - - GraphEngine stop command is only sent for ADVANCED_CHAT and WORKFLOW modes - """ - # Arrange - task_id = "task-123" - invoke_from = InvokeFrom.WEB_APP - user_id = "user-456" - # Act - AppTaskService.stop_task(task_id, invoke_from, user_id, app_mode) +def _assert_stop_flag(redis: _StopRedis) -> None: + assert redis.values[_STOP_KEY] == b"1" + assert redis.expirations[_STOP_KEY] == 600 - # Assert - mock_app_queue_manager.set_stop_flag.assert_called_once_with(task_id, invoke_from, user_id) - if should_call_graph_engine: - mock_graph_engine_manager.assert_called_once() - mock_graph_engine_manager.return_value.send_stop_command.assert_called_once_with(task_id) - else: - mock_graph_engine_manager.assert_not_called() - @patch("services.app_task_service.GraphEngineManager") - @patch("services.app_task_service.AppQueueManager") - def test_stop_task_legacy_mechanism_called_even_if_graph_engine_fails( - self, mock_app_queue_manager, mock_graph_engine_manager - ): - """Test that legacy Redis flag is set even if GraphEngine fails. +def _assert_graph_command(redis: _StopRedis) -> None: + assert set(redis.commands) == {_COMMAND_KEY} + assert [json.loads(command) for command in redis.commands[_COMMAND_KEY]] == [ + {"command_type": "abort", "payload": None, "reason": "User requested stop"} + ] + assert redis.expirations[_COMMAND_KEY] == 3600 + assert redis.values[f"{_COMMAND_KEY}:pending"] == b"1" + assert redis.expirations[f"{_COMMAND_KEY}:pending"] == 3600 - This ensures backward compatibility: the legacy mechanism should complete - before attempting the GraphEngine command, so the stop flag is set - regardless of GraphEngine success. - """ - # Arrange - task_id = "task-123" - invoke_from = InvokeFrom.WEB_APP - user_id = "user-456" - app_mode = AppMode.ADVANCED_CHAT - # Simulate GraphEngine failure - mock_graph_engine_manager.return_value.send_stop_command.side_effect = Exception("GraphEngine error") +@pytest.mark.parametrize("app_mode", list(AppMode)) +def test_stop_task_sends_graph_command_only_for_workflow_modes(app_mode: AppMode) -> None: + redis = _StopRedis(values={_OWNER_KEY: f"end-user-{_USER_ID}".encode()}) - # Act & Assert - should raise the exception since it's not caught - with pytest.raises(Exception, match="GraphEngine error"): - AppTaskService.stop_task(task_id, invoke_from, user_id, app_mode) + AppTaskControlService(redis_client=redis).stop_task(_TASK_ID, InvokeFrom.WEB_APP, _USER_ID, app_mode) - # Verify legacy mechanism was still called before the exception - mock_app_queue_manager.set_stop_flag.assert_called_once_with(task_id, invoke_from, user_id) + assert redis.reads == [_OWNER_KEY] + _assert_stop_flag(redis) + if app_mode in (AppMode.ADVANCED_CHAT, AppMode.WORKFLOW): + assert redis.operations == ["legacy_flag", "graph_command"] + _assert_graph_command(redis) + else: + assert redis.operations == ["legacy_flag"] + assert redis.commands == {} + + +@pytest.mark.parametrize( + ("invoke_from", "owner", "should_set_flag"), + [ + (InvokeFrom.EXPLORE, b"account-user-1", True), + (InvokeFrom.DEBUGGER, b"account-user-1", True), + (InvokeFrom.WEB_APP, b"end-user-user-1", True), + (InvokeFrom.SERVICE_API, b"end-user-user-1", True), + (InvokeFrom.EXPLORE, b"end-user-user-1", False), + (InvokeFrom.DEBUGGER, b"end-user-user-1", False), + (InvokeFrom.WEB_APP, b"account-user-1", False), + (InvokeFrom.SERVICE_API, b"account-user-1", False), + (InvokeFrom.EXPLORE, b"account-another-user", False), + (InvokeFrom.WEB_APP, b"end-user-another-user", False), + (InvokeFrom.EXPLORE, None, False), + (InvokeFrom.WEB_APP, None, False), + ], +) +def test_task_ownership_controls_only_the_legacy_flag( + invoke_from: InvokeFrom, owner: bytes | None, should_set_flag: bool +) -> None: + redis = _StopRedis(values={_OWNER_KEY: owner} if owner is not None else {}) + + AppTaskControlService(redis_client=redis).stop_task(_TASK_ID, invoke_from, _USER_ID, AppMode.WORKFLOW) + + assert redis.reads == [_OWNER_KEY] + if should_set_flag: + _assert_stop_flag(redis) + assert redis.operations == ["legacy_flag", "graph_command"] + else: + assert _STOP_KEY not in redis.values + assert redis.operations == ["graph_command"] + # Preserve the existing behavior even when the legacy ownership check does not match. + _assert_graph_command(redis) + + +@pytest.mark.parametrize("owner", [None, b"account-another-user"]) +def test_unchecked_workflow_stop_skips_ownership_read(owner: bytes | None) -> None: + redis = _StopRedis( + values={_OWNER_KEY: owner} if owner is not None else {}, + read_error=AssertionError("Unchecked workflow stop must not read ownership"), + ) + + AppTaskControlService(redis_client=redis).stop_workflow_task_no_user_check(task_id=_TASK_ID) + + assert redis.reads == [] + assert redis.operations == ["legacy_flag", "graph_command"] + _assert_stop_flag(redis) + _assert_graph_command(redis) + + +def test_unchecked_workflow_stop_with_empty_task_id_is_noop() -> None: + redis = _StopRedis() + + AppTaskControlService(redis_client=redis).stop_workflow_task_no_user_check(task_id="") + + assert redis.reads == [] + assert redis.operations == [] + assert redis.values == {} + assert redis.commands == {} + + +@pytest.mark.parametrize("unchecked", [False, True]) +def test_flag_write_failure_propagates_before_graph_command(unchecked: bool) -> None: + error = RedisConnectionError("stop flag write failed") + redis = _StopRedis(values={_OWNER_KEY: b"account-user-1"}, flag_error=error) + service = AppTaskControlService(redis_client=redis) + stop = ( + partial(service.stop_workflow_task_no_user_check, task_id=_TASK_ID) + if unchecked + else partial(service.stop_task, _TASK_ID, InvokeFrom.EXPLORE, _USER_ID, AppMode.WORKFLOW) + ) + + with pytest.raises(RedisConnectionError, match="stop flag write failed") as caught: + stop() + + assert caught.value is error + assert redis.operations == ["legacy_flag"] + assert _STOP_KEY not in redis.values + assert redis.commands == {} + + +@pytest.mark.parametrize("unchecked", [False, True]) +def test_graph_redis_failure_is_swallowed_after_legacy_flag(unchecked: bool, caplog: pytest.LogCaptureFixture) -> None: + redis = _StopRedis( + values={_OWNER_KEY: b"account-user-1"}, command_error=RedisConnectionError("command write failed") + ) + service = AppTaskControlService(redis_client=redis) + + if unchecked: + service.stop_workflow_task_no_user_check(task_id=_TASK_ID) + else: + service.stop_task(_TASK_ID, InvokeFrom.EXPLORE, _USER_ID, AppMode.WORKFLOW) + + _assert_stop_flag(redis) + assert redis.operations == ["legacy_flag", "graph_command"] + assert redis.commands == {} + assert "Failed to send graph engine command AbortCommand" in caplog.text + + +def test_ownership_read_failure_propagates_without_either_stop_signal() -> None: + error = RedisConnectionError("ownership read failed") + redis = _StopRedis(read_error=error) + + with pytest.raises(RedisConnectionError, match="ownership read failed") as caught: + AppTaskControlService(redis_client=redis).stop_task(_TASK_ID, InvokeFrom.EXPLORE, _USER_ID, AppMode.WORKFLOW) + + assert caught.value is error + assert redis.reads == [_OWNER_KEY] + assert redis.operations == [] + assert redis.values == {} + assert redis.commands == {} + + +def test_service_instances_keep_their_redis_dependencies_separate() -> None: + first = _StopRedis(values={_OWNER_KEY: b"account-user-1"}) + second = _StopRedis(values={_OWNER_KEY: b"account-another-user"}) + first_service = AppTaskControlService(redis_client=first) + second_service = AppTaskControlService(redis_client=second) + + first_service.stop_task(_TASK_ID, InvokeFrom.EXPLORE, _USER_ID, AppMode.CHAT) + second_service.stop_workflow_task_no_user_check(task_id=_TASK_ID) + + assert first.reads == [_OWNER_KEY] + assert first.operations == ["legacy_flag"] + assert first.commands == {} + assert second.reads == [] + assert second.operations == ["legacy_flag", "graph_command"] + _assert_stop_flag(first) + _assert_stop_flag(second) + _assert_graph_command(second) + + +def test_legacy_static_entry_point_passes_global_client_through_the_same_implementation( + monkeypatch: pytest.MonkeyPatch, +) -> None: + redis = _StopRedis(values={_OWNER_KEY: b"end-user-user-1"}) + # Only the legacy composition point gets a global client; lower-level globals remain traps. + monkeypatch.setattr(task_module, "redis_client", redis) + + AppTaskService.stop_task(_TASK_ID, InvokeFrom.SERVICE_API, _USER_ID, AppMode.ADVANCED_CHAT) + + assert redis.reads == [_OWNER_KEY] + assert redis.operations == ["legacy_flag", "graph_command"] + _assert_stop_flag(redis) + _assert_graph_command(redis) diff --git a/api/tests/unit_tests/services/test_installed_app_generation_adapters.py b/api/tests/unit_tests/services/test_installed_app_generation_adapters.py index 08d3b346b12..451c7dfe118 100644 --- a/api/tests/unit_tests/services/test_installed_app_generation_adapters.py +++ b/api/tests/unit_tests/services/test_installed_app_generation_adapters.py @@ -27,6 +27,7 @@ from services.installed_app_generation_adapters import AppGenerateServiceRuntime from services.installed_app_generation_service import GenerationResponse _ARGS: dict[str, object] = {"inputs": {"count": 0}, "query": "hello", "auto_generate_name": False} +_WORKFLOW_ARGS: dict[str, object] = {"inputs": {"count": 0}, "files": []} @dataclass @@ -455,28 +456,67 @@ def test_visible_conversation_uses_shared_chat_dispatch_after_preflight_session_ assert result.closed is True -def test_advanced_chat_dispatch_starts_task_after_subscription_with_runtime_session_closed( +def _generate_workflow( + harness: _RuntimeHarness, + session_factory: sessionmaker[Session], + args: Mapping[str, object], +) -> GenerationResponse: + # WorkflowService also creates its repository factory from Flask's db.engine. + # Bind it to the same real database without replacing workflow lookup or dispatch. + runtime_app = Flask(__name__) + runtime_app.config["SQLALCHEMY_DATABASE_URI"] = str(session_factory.kw["bind"].url) + db.init_app(runtime_app) + with runtime_app.app_context(): + try: + return harness.runtime.generate( + app_id=harness.app_id, account_id=harness.account_id, args=args, streaming=True + ) + finally: + db.engine.dispose() + + +@pytest.mark.parametrize( + ("mode", "trigger"), + [ + pytest.param(AppMode.ADVANCED_CHAT, False, id="advanced-chat"), + pytest.param(AppMode.WORKFLOW, False, id="workflow"), + pytest.param(AppMode.WORKFLOW, True, id="explore-trigger-workflow"), + ], +) +def test_workflow_dispatch_starts_task_after_subscription_with_runtime_session_closed( harness: _RuntimeHarness, monkeypatch: pytest.MonkeyPatch, sqlite_session_factory: sessionmaker[Session], config_overrides: Callable[..., None], + mode: AppMode, + trigger: bool, ) -> None: config_overrides( DEPLOYMENT_EDITION=DeploymentEdition.COMMUNITY, ENABLE_OTEL=False, + APP_DEFAULT_ACTIVE_REQUESTS=0, APP_MAX_ACTIVE_REQUESTS=0, PUBSUB_REDIS_CHANNEL_TYPE="streams", ) - conversation_id = _seed_chat_conversation(harness, sqlite_session_factory, mode=AppMode.ADVANCED_CHAT) + args = _WORKFLOW_ARGS + if mode == AppMode.ADVANCED_CHAT: + conversation_id = _seed_chat_conversation(harness, sqlite_session_factory, mode=mode) + args = {**_ARGS, "conversation_id": conversation_id} with sqlite_session_factory.begin() as session: app = session.get(App, harness.app_id) assert app is not None + app.mode = mode workflow = Workflow( tenant_id=app.tenant_id, app_id=app.id, - type=WorkflowType.CHAT, + type=WorkflowType.WORKFLOW if mode == AppMode.WORKFLOW else WorkflowType.CHAT, version="2026-09-07 00:00:00", - graph='{"nodes": [], "edges": []}', + graph=json.dumps( + { + "nodes": [{"id": "schedule", "data": {"type": "trigger-schedule"}}] if trigger else [], + "edges": [], + } + ), _features="{}", created_by=harness.account_id, ) @@ -503,7 +543,6 @@ def test_advanced_chat_dispatch_starts_task_after_subscription_with_runtime_sess subscription.__enter__.side_effect = activate_subscription subscription.receive.return_value = b'{"event":"workflow_finished"}' monkeypatch.setattr(message_based_app_generator, "get_pubsub_broadcast_channel", lambda: channel) - args = {**_ARGS, "conversation_id": conversation_id} submitted: list[generation_module.AppExecutionParams] = [] def enqueue(payload_json: str) -> None: @@ -515,21 +554,11 @@ def test_advanced_chat_dispatch_starts_task_after_subscription_with_runtime_sess monkeypatch.setattr(generation_module.workflow_based_app_execution_task, "delay", enqueue) - # WorkflowService also creates its repository factory from Flask's db.engine. - # Bind it to the same real database without replacing workflow lookup or dispatch. - runtime_app = Flask(__name__) - runtime_app.config["SQLALCHEMY_DATABASE_URI"] = str(sqlite_session_factory.kw["bind"].url) - db.init_app(runtime_app) - with runtime_app.app_context(): - try: - result = harness.runtime.generate( - app_id=harness.app_id, account_id=harness.account_id, args=args, streaming=True - ) - finally: - db.engine.dispose() + result = _generate_workflow(harness, sqlite_session_factory, args) assert isinstance(result, RateLimitGenerator) assert len(harness.closed_sessions) == 2 + assert harness.committed_sessions == [harness.closed_sessions[1]] assert submitted == [] subscriber.prepare_subscription.assert_called_once_with() subscription.__enter__.assert_not_called() @@ -541,10 +570,54 @@ def test_advanced_chat_dispatch_starts_task_after_subscription_with_runtime_sess assert len(submitted) == 1 payload = submitted[0] assert (payload.app_id, payload.workflow_id, payload.tenant_id) == (harness.app_id, workflow_id, tenant_id) - assert payload.app_mode == AppMode.ADVANCED_CHAT + assert payload.app_mode == mode assert payload.user.model_dump(mode="json") == {"TYPE": "account", "user_id": harness.account_id} assert payload.args == args assert payload.invoke_from == InvokeFrom.EXPLORE assert payload.streaming is True subscription.__exit__.assert_called_once() assert result.closed is True + + +def test_unpublished_workflow_raises_before_subscription_or_task_creation( + harness: _RuntimeHarness, + monkeypatch: pytest.MonkeyPatch, + sqlite_session_factory: sessionmaker[Session], + config_overrides: Callable[..., None], +) -> None: + config_overrides( + DEPLOYMENT_EDITION=DeploymentEdition.COMMUNITY, + ENABLE_OTEL=False, + APP_DEFAULT_ACTIVE_REQUESTS=0, + APP_MAX_ACTIVE_REQUESTS=0, + ) + with sqlite_session_factory.begin() as session: + app = session.get(App, harness.app_id) + assert app is not None + app.mode = AppMode.WORKFLOW + session.add( + Workflow( + tenant_id=app.tenant_id, + app_id=app.id, + type=WorkflowType.WORKFLOW, + version="draft", + graph='{"nodes": [], "edges": []}', + _features="{}", + created_by=harness.account_id, + ) + ) + + channel = MagicMock(spec=BroadcastChannel) + monkeypatch.setattr(message_based_app_generator, "get_pubsub_broadcast_channel", lambda: channel) + enqueue = MagicMock() + monkeypatch.setattr(generation_module.workflow_based_app_execution_task, "delay", enqueue) + + with pytest.raises(ValueError, match="^Workflow not published$") as raised: + _generate_workflow(harness, sqlite_session_factory, _WORKFLOW_ARGS) + + assert type(raised.value) is ValueError + assert len(harness.closed_sessions) == 2 + assert all(not session.in_transaction() for session in harness.closed_sessions) + assert harness.committed_sessions == [] + channel.topic.assert_not_called() + enqueue.assert_not_called()