refactor(api): migrate installed app workflow and task control (#41912)

This commit is contained in:
非法操作
2026-09-23 08:25:37 +00:00
committed by GitHub
parent 14d8eecea6
commit f0078b7935
11 changed files with 990 additions and 333 deletions
@@ -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,
+36 -48
View File
@@ -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/<uuid:installed_app_id>/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/<uuid:installed_app_id>/workflows/tasks/<string:task_id>/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")
+7 -4
View File
@@ -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:
+4 -3
View File
@@ -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:
+32 -19
View File
@@ -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,
)
@@ -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,
*,
@@ -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/<uuid:installed_app_id>/completion-messages/<string:task_id>/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)
@@ -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/<uuid:installed_app_id>/workflows/run",
)
harness.api.add_resource(
workflow_module.InstalledAppWorkflowTaskStopApi,
"/installed-apps/<uuid:installed_app_id>/workflows/tasks/<string:task_id>/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()
@@ -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:
@@ -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)
@@ -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()