refactor(api): extract console app site command boundary (#41216)

This commit is contained in:
非法操作
2026-08-26 06:31:06 +00:00
committed by GitHub
parent d5a8ace1c4
commit c4ace15afd
9 changed files with 608 additions and 102 deletions
+15
View File
@@ -245,6 +245,21 @@ forbidden_modules =
sqlalchemy
werkzeug
[importlinter:contract:app-site-service-boundary]
name = App site application service is framework and persistence neutral
type = forbidden
source_modules =
services.app_site_service
forbidden_modules =
configs
controllers
extensions
flask
models
repositories
sqlalchemy
werkzeug
[importlinter:contract:webapp-access-query-service-boundary]
name = Web app access query application service is framework and persistence neutral
type = forbidden
+60 -74
View File
@@ -1,34 +1,45 @@
from typing import Literal
from uuid import UUID
from flask_restx import Resource
from pydantic import BaseModel, Field, field_validator
from sqlalchemy import select
from sqlalchemy.orm import Session
from werkzeug.exceptions import NotFound
from constants.languages import supported_language
from controllers.common.schema import register_schema_models
from controllers.common.session import with_session
from controllers.console import console_ns
from controllers.console.app.wraps import agent_manage_required_for_agent_app, get_app_model
from controllers.console.app.error import AppNotFoundError
from controllers.console.app.wraps import agent_manage_required_for_agent_app
from controllers.console.flask_admission import console_account_admission
from controllers.console.wraps import (
RBACPermission,
RBACResourceScope,
account_initialization_required,
edit_permission_required,
is_admin_or_owner_required,
model_validate,
rbac_permission_required,
setup_required,
with_current_user,
)
from extensions.ext_application_services import application_services
from fields.base import ResponseModel
from libs.datetime_utils import naive_utc_now
from libs.helper import dump_response
from libs.login import login_required
from models import Site
from models.account import Account
from models.model import App
from machinery.context import RequestContext
from models.account import TenantAccountRole
from services.app_site_service import (
AppSiteAppNotFoundError,
AppSiteChanges,
AppSiteNotFoundError,
AppSiteTokenStrategy,
)
_APP_SITE_EDIT_ROLES = frozenset(
{
TenantAccountRole.OWNER,
TenantAccountRole.ADMIN,
TenantAccountRole.EDITOR,
}
)
_APP_SITE_TOKEN_RESET_ROLES = frozenset(
{
TenantAccountRole.OWNER,
TenantAccountRole.ADMIN,
}
)
class AppSiteUpdatePayload(BaseModel):
@@ -45,7 +56,7 @@ class AppSiteUpdatePayload(BaseModel):
privacy_policy: str | None = Field(default=None)
input_placeholder: str | None = Field(default=None)
custom_disclaimer: str | None = Field(default=None)
customize_token_strategy: Literal["must", "allow", "not_allow"] | None = Field(default=None)
customize_token_strategy: AppSiteTokenStrategy | None = Field(default=None)
prompt_public: bool | None = Field(default=None)
show_workflow_steps: bool | None = Field(default=None)
use_icon_as_answer_icon: bool | None = Field(default=None)
@@ -57,6 +68,9 @@ class AppSiteUpdatePayload(BaseModel):
return value
return supported_language(value)
def to_changes(self) -> AppSiteChanges:
return AppSiteChanges(**self.model_dump())
class AppSiteResponse(ResponseModel):
app_id: str
@@ -90,47 +104,25 @@ class AppSite(Resource):
@console_ns.response(200, "Site configuration updated successfully", console_ns.models[AppSiteResponse.__name__])
@console_ns.response(403, "Insufficient permissions")
@console_ns.response(404, "App not found")
@setup_required
@login_required
@edit_permission_required
@rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_RELEASE_AND_VERSION)
@console_account_admission(
allowed_roles=_APP_SITE_EDIT_ROLES,
rbac_resource_scope=RBACResourceScope.APP,
rbac_permission=RBACPermission.APP_RELEASE_AND_VERSION,
)
@agent_manage_required_for_agent_app
@account_initialization_required
@with_current_user
@with_session
@get_app_model
@model_validate(AppSiteUpdatePayload)
def post(self, req_data: AppSiteUpdatePayload, session: Session, current_user: Account, app_model: App):
site = session.scalar(select(Site).where(Site.app_id == app_model.id).limit(1))
if not site:
raise NotFound
for attr_name in [
"title",
"icon_type",
"icon",
"icon_background",
"description",
"default_language",
"chat_color_theme",
"chat_color_theme_inverted",
"customize_domain",
"copyright",
"privacy_policy",
"input_placeholder",
"custom_disclaimer",
"customize_token_strategy",
"prompt_public",
"show_workflow_steps",
"use_icon_as_answer_icon",
]:
value = getattr(req_data, attr_name)
if value is not None:
setattr(site, attr_name, value)
site.updated_by = current_user.id
site.updated_at = naive_utc_now()
session.flush()
def post(
self,
req_data: AppSiteUpdatePayload,
request_context: RequestContext,
app_id: UUID,
):
try:
site = application_services().app_sites.update(request_context, str(app_id), req_data.to_changes())
except AppSiteAppNotFoundError as error:
raise AppNotFoundError() from error
except AppSiteNotFoundError as error:
raise NotFound from error
return dump_response(AppSiteResponse, site)
@@ -143,24 +135,18 @@ class AppSiteAccessTokenReset(Resource):
@console_ns.response(200, "Access token reset successfully", console_ns.models[AppSiteResponse.__name__])
@console_ns.response(403, "Insufficient permissions (admin/owner required)")
@console_ns.response(404, "App or site not found")
@setup_required
@login_required
@is_admin_or_owner_required
@rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_RELEASE_AND_VERSION)
@console_account_admission(
allowed_roles=_APP_SITE_TOKEN_RESET_ROLES,
rbac_resource_scope=RBACResourceScope.APP,
rbac_permission=RBACPermission.APP_RELEASE_AND_VERSION,
)
@agent_manage_required_for_agent_app
@account_initialization_required
@with_current_user
@with_session
@get_app_model
def post(self, session: Session, current_user: Account, app_model: App):
site = session.scalar(select(Site).where(Site.app_id == app_model.id).limit(1))
if not site:
raise NotFound
site.code = Site.generate_code(16, session=session)
site.updated_by = current_user.id
site.updated_at = naive_utc_now()
session.flush()
def post(self, request_context: RequestContext, app_id: UUID):
try:
site = application_services().app_sites.reset_access_token(request_context, str(app_id))
except AppSiteAppNotFoundError as error:
raise AppNotFoundError() from error
except AppSiteNotFoundError as error:
raise NotFound from error
return dump_response(AppSiteResponse, site)
@@ -21,6 +21,7 @@ from repositories.account_activation_repository import SQLAlchemyAccountActivati
from repositories.account_integration_repository import SQLAlchemyAccountIntegrationRepository
from repositories.account_repository import SQLAlchemyAccountRepository
from repositories.app_definition_query_repository import AppDefinitionQueryRepository
from repositories.app_site_command_repository import AppSiteCommandRepository
from repositories.data_source_api_key_auth_repository import SQLAlchemyDataSourceApiKeyAuthBindingRepository
from repositories.explore_banner_query_repository import ExploreBannerQueryRepository
from repositories.installation_state_repository import InstallationStateRepository
@@ -68,6 +69,7 @@ from services.account_password_hasher import LegacyAccountPasswordHasher
from services.account_password_service import AccountPasswordService
from services.account_profile_service import AccountProfileService
from services.app_definition_query_service import AppDefinitionQueryService
from services.app_site_service import AppSiteService
from services.auth.data_source_api_key_auth_gateways import (
ProviderApiKeyAuthCredentialValidator,
TenantApiKeyAuthCredentialEncryptor,
@@ -144,6 +146,7 @@ class ApplicationServices:
accounts: AccountServices
account_activation: AccountActivationService
app_definitions: AppDefinitionQueryService
app_sites: AppSiteService
billing_portal: BillingPortalService
data_source_api_key_auth: DataSourceApiKeyAuthService
webapp_access: WebAppAccessQueryService
@@ -275,6 +278,9 @@ def build_application_services(
dify_config.CONSOLE_API_URL + "/console/api/workspaces/current/tool-provider/builtin/"
),
),
app_sites=AppSiteService(
sites=AppSiteCommandRepository(session_factory=database_client),
),
billing_portal=BillingPortalService(
accounts=accounts,
get_subscription=BillingService.get_subscription,
@@ -0,0 +1,108 @@
"""SQLAlchemy persistence adapter for Console app site management."""
from dataclasses import asdict
from typing import override
from sqlalchemy import select
from sqlalchemy.orm import Session, sessionmaker
from libs.datetime_utils import naive_utc_now
from models.enums import AppStatus
from models.model import App, Site
from services.app_site_service import (
AppSiteAppNotFoundError,
AppSiteChanges,
AppSiteCommandResult,
AppSiteNotFoundError,
AppSiteStore,
)
class AppSiteCommandRepository(AppSiteStore):
def __init__(self, *, session_factory: sessionmaker[Session]) -> None:
self._session_factory = session_factory
@override
def update_site(
self,
*,
workspace_id: str,
app_id: str,
actor_id: str,
changes: AppSiteChanges,
) -> AppSiteCommandResult:
with self._session_factory.begin() as session:
site = self._get_site(session, workspace_id, app_id)
for field_name, value in asdict(changes).items():
if value is not None:
setattr(site, field_name, value)
site.updated_by = actor_id
site.updated_at = naive_utc_now()
session.flush()
return self._to_command_result(site)
@override
def reset_access_token(
self,
*,
workspace_id: str,
app_id: str,
actor_id: str,
) -> AppSiteCommandResult:
with self._session_factory.begin() as session:
site = self._get_site(session, workspace_id, app_id)
site.code = Site.generate_code(16, session=session)
site.updated_by = actor_id
site.updated_at = naive_utc_now()
session.flush()
return self._to_command_result(site)
@staticmethod
def _get_site(session: Session, workspace_id: str, app_id: str) -> Site:
site = session.scalar(
select(Site)
.join(App, App.id == Site.app_id)
.where(
App.id == app_id,
App.tenant_id == workspace_id,
App.status == AppStatus.NORMAL,
)
.limit(1)
)
if site is not None:
return site
app_exists = session.scalar(
select(App.id)
.where(
App.id == app_id,
App.tenant_id == workspace_id,
App.status == AppStatus.NORMAL,
)
.limit(1)
)
if app_exists is None:
raise AppSiteAppNotFoundError
raise AppSiteNotFoundError
@staticmethod
def _to_command_result(site: Site) -> AppSiteCommandResult:
return AppSiteCommandResult(
app_id=site.app_id,
code=site.code,
title=site.title,
icon=site.icon,
icon_background=site.icon_background,
description=site.description,
default_language=site.default_language,
customize_domain=site.customize_domain,
copyright=site.copyright,
privacy_policy=site.privacy_policy,
input_placeholder=site.input_placeholder,
custom_disclaimer=site.custom_disclaimer,
customize_token_strategy=str(site.customize_token_strategy),
prompt_public=site.prompt_public,
show_workflow_steps=site.show_workflow_steps,
use_icon_as_answer_icon=site.use_icon_as_answer_icon,
)
+107
View File
@@ -0,0 +1,107 @@
"""Application boundary for Console app site management."""
from dataclasses import dataclass
from typing import Literal, NamedTuple, Protocol
from machinery.context import RequestContext
AppSiteTokenStrategy = Literal["must", "allow", "not_allow"]
@dataclass(frozen=True, slots=True)
class AppSiteChanges:
title: str | None = None
icon_type: str | None = None
icon: str | None = None
icon_background: str | None = None
description: str | None = None
default_language: str | None = None
chat_color_theme: str | None = None
chat_color_theme_inverted: bool | None = None
customize_domain: str | None = None
copyright: str | None = None
privacy_policy: str | None = None
input_placeholder: str | None = None
custom_disclaimer: str | None = None
customize_token_strategy: AppSiteTokenStrategy | None = None
prompt_public: bool | None = None
show_workflow_steps: bool | None = None
use_icon_as_answer_icon: bool | None = None
class AppSiteCommandResult(NamedTuple):
app_id: str
code: str | None
title: str
icon: str | None
icon_background: str | None
description: str | None
default_language: str
customize_domain: str | None
copyright: str | None
privacy_policy: str | None
input_placeholder: str | None
custom_disclaimer: str | None
customize_token_strategy: str
prompt_public: bool
show_workflow_steps: bool
use_icon_as_answer_icon: bool
class AppSiteStore(Protocol):
def update_site(
self,
*,
workspace_id: str,
app_id: str,
actor_id: str,
changes: AppSiteChanges,
) -> AppSiteCommandResult: ...
def reset_access_token(
self,
*,
workspace_id: str,
app_id: str,
actor_id: str,
) -> AppSiteCommandResult: ...
class AppSiteError(Exception):
"""Base class for framework-neutral app site failures."""
class AppSiteAppNotFoundError(AppSiteError):
def __init__(self) -> None:
super().__init__("App not found")
class AppSiteNotFoundError(AppSiteError):
def __init__(self) -> None:
super().__init__("Site not found")
class AppSiteService:
def __init__(self, *, sites: AppSiteStore) -> None:
self._sites = sites
def update(self, context: RequestContext, app_id: str, changes: AppSiteChanges) -> AppSiteCommandResult:
return self._sites.update_site(
workspace_id=self._workspace_id(context),
app_id=app_id,
actor_id=context.account_id,
changes=changes,
)
def reset_access_token(self, context: RequestContext, app_id: str) -> AppSiteCommandResult:
return self._sites.reset_access_token(
workspace_id=self._workspace_id(context),
app_id=app_id,
actor_id=context.account_id,
)
@staticmethod
def _workspace_id(context: RequestContext) -> str:
if context.active_workspace_id is None:
raise RuntimeError("Console account admission did not resolve an active workspace")
return context.active_workspace_id
@@ -52,6 +52,7 @@ from controllers.console.app import (
wraps as wraps_module,
)
from controllers.console.app.completion import ChatMessagePayload, CompletionMessagePayload
from controllers.console.app.error import AppNotFoundError
from controllers.console.app.mcp_server import MCPServerCreatePayload, MCPServerUpdatePayload
from controllers.console.app.ops_trace import TraceConfigPayload, TraceProviderQuery
from controllers.console.app.site import AppSiteUpdatePayload
@@ -64,12 +65,18 @@ from controllers.console.app.workflow_draft_variable import (
)
from controllers.console.app.workflow_statistic import WorkflowStatisticQuery
from controllers.console.app.workflow_trigger import Parser, ParserEnable
from machinery.context import RequestContext
from models import App, Site
from models.account import Account, AccountStatus
from models.engine import db
from models.enums import CustomizeTokenStrategy
from models.trigger import WorkflowWebhookTrigger
from repositories.sqlalchemy_api_workflow_run_repository import DifyAPISQLAlchemyWorkflowRunRepository
from services.app_site_service import (
AppSiteAppNotFoundError,
AppSiteChanges,
AppSiteCommandResult,
AppSiteNotFoundError,
)
APP_ID = "11111111-1111-1111-1111-111111111111"
TENANT_ID = "22222222-2222-2222-2222-222222222222"
@@ -387,27 +394,37 @@ class TestOpsTraceEndpoints:
class TestSiteEndpoints:
@staticmethod
def _add_site(session: Session) -> Site:
site = Site(
def _command_result(*, code: str = "test-code") -> AppSiteCommandResult:
return AppSiteCommandResult(
app_id=APP_ID,
title="My Site",
description="Test site",
default_language="en-US",
customize_token_strategy=CustomizeTokenStrategy.NOT_ALLOW,
code="test-code",
input_placeholder="Ask me anything",
code=code,
icon=None,
icon_background=None,
customize_domain=None,
copyright=None,
privacy_policy=None,
custom_disclaimer="",
customize_token_strategy="not_allow",
prompt_public=False,
show_workflow_steps=True,
use_icon_as_answer_icon=False,
)
session.add(site)
session.commit()
return site
def test_site_response_structure(self):
def test_site_payload_maps_to_application_changes(self):
payload = AppSiteUpdatePayload(
title="My Site",
description="Test site",
input_placeholder="Ask me anything",
)
assert payload.title == "My Site"
assert payload.input_placeholder == "Ask me anything"
assert payload.to_changes() == AppSiteChanges(
title="My Site",
description="Test site",
input_placeholder="Ask me anything",
)
def test_site_default_language_validation(self):
payload = AppSiteUpdatePayload(default_language="en-US")
@@ -415,44 +432,74 @@ class TestSiteEndpoints:
def test_app_site_update_post(
self,
database_app: Flask,
) -> None:
api = site_module.AppSite()
method = unwrap(api.post)
site = self._add_site(db.session)
services = MagicMock()
services.app_sites.update.return_value = self._command_result()
context = RequestContext("request-1", None, USER_ID, TENANT_ID)
with database_app.test_request_context("/", json={"title": "My Site", "input_placeholder": "Ask me anything"}):
with patch.object(site_module, "application_services", return_value=services):
result = method(
api,
AppSiteUpdatePayload(title="My Site", input_placeholder="Ask me anything"),
db.session,
_make_account(),
app_model=_make_app(),
context,
app_id=uuid.UUID(APP_ID),
)
db.session.refresh(site)
assert isinstance(result, dict)
assert result["title"] == "My Site"
assert result["input_placeholder"] == "Ask me anything"
assert site.input_placeholder == "Ask me anything"
assert result["access_token"] == "test-code"
assert result["code"] == "test-code"
services.app_sites.update.assert_called_once_with(
context,
APP_ID,
AppSiteChanges(title="My Site", input_placeholder="Ask me anything"),
)
def test_app_site_access_token_reset(
self,
database_app: Flask,
monkeypatch: pytest.MonkeyPatch,
) -> None:
api = site_module.AppSiteAccessTokenReset()
method = unwrap(api.post)
site = self._add_site(db.session)
monkeypatch.setattr(site_module.Site, "generate_code", lambda *_args, **_kwargs: "code")
services = MagicMock()
services.app_sites.reset_access_token.return_value = self._command_result(code="new-code")
context = RequestContext("request-1", None, USER_ID, TENANT_ID)
with database_app.test_request_context("/"):
result = method(api, db.session, _make_account(), app_model=_make_app())
with patch.object(site_module, "application_services", return_value=services):
result = method(api, context, app_id=uuid.UUID(APP_ID))
db.session.refresh(site)
assert isinstance(result, dict)
assert result["access_token"] == "code"
assert site.code == "code"
assert result["access_token"] == "new-code"
assert result["code"] == "new-code"
services.app_sites.reset_access_token.assert_called_once_with(context, APP_ID)
def test_app_site_update_maps_missing_app(self) -> None:
api = site_module.AppSite()
method = unwrap(api.post)
services = MagicMock()
services.app_sites.update.side_effect = AppSiteAppNotFoundError()
context = RequestContext("request-1", None, USER_ID, TENANT_ID)
with (
patch.object(site_module, "application_services", return_value=services),
pytest.raises(AppNotFoundError),
):
method(api, AppSiteUpdatePayload(), context, app_id=uuid.UUID(APP_ID))
def test_app_site_reset_maps_missing_site(self) -> None:
api = site_module.AppSiteAccessTokenReset()
method = unwrap(api.post)
services = MagicMock()
services.app_sites.reset_access_token.side_effect = AppSiteNotFoundError()
context = RequestContext("request-1", None, USER_ID, TENANT_ID)
with (
patch.object(site_module, "application_services", return_value=services),
pytest.raises(NotFound),
):
method(api, context, app_id=uuid.UUID(APP_ID))
class TestWorkflowEndpoints:
@@ -21,6 +21,7 @@ from models.model import AccountTrialAppRecord, DifySetup
from repositories.account_activation_repository import SQLAlchemyAccountActivationRepository
from repositories.account_integration_repository import SQLAlchemyAccountIntegrationRepository
from repositories.account_repository import SQLAlchemyAccountRepository
from repositories.app_site_command_repository import AppSiteCommandRepository
from services import recommended_app_catalog_gateway
from services.account_activation_adapters import (
BillingAccountActivationEligibility,
@@ -29,6 +30,7 @@ from services.account_activation_adapters import (
RegisterServiceInvitationTokenStore,
)
from services.account_avatar_file_gateway import SQLAlchemyAccountAvatarFileGateway
from services.app_site_service import AppSiteService
from services.auth.data_source_api_key_auth_service import DataSourceApiKeyAuthService
from services.billing_portal_service import BillingPortalService
from services.billing_service import BillingService
@@ -176,6 +178,21 @@ def test_build_application_services_wires_tag_boundary(
assert isinstance(services.tags, TagApplicationService)
def test_build_application_services_wires_app_site_boundary(
sqlite_session_factory: sessionmaker[Session],
) -> None:
services = ext_application_services.build_application_services(
database_client=sqlite_session_factory,
deployment_edition=DeploymentEdition.COMMUNITY,
initialization_password="",
redis=MagicMock(spec=RedisClientWrapper),
)
assert isinstance(services.app_sites, AppSiteService)
assert isinstance(services.app_sites._sites, AppSiteCommandRepository)
assert services.app_sites._sites._session_factory is sqlite_session_factory
def test_build_application_services_wires_billing_service(
sqlite_session: Session,
sqlite_session_factory: sessionmaker[Session],
@@ -0,0 +1,161 @@
import pytest
from sqlalchemy import delete, select
from sqlalchemy.orm import Session, sessionmaker
from models.enums import CustomizeTokenStrategy
from models.model import App, AppMode, Site
from repositories.app_site_command_repository import AppSiteCommandRepository
from services.app_site_service import AppSiteAppNotFoundError, AppSiteChanges, AppSiteNotFoundError
_APP_ID = "11111111-1111-1111-1111-111111111111"
_WORKSPACE_ID = "22222222-2222-2222-2222-222222222222"
_OTHER_WORKSPACE_ID = "33333333-3333-3333-3333-333333333333"
_ACTOR_ID = "44444444-4444-4444-4444-444444444444"
def _persist_app(session: Session, *, with_site: bool = True) -> None:
app = App(
id=_APP_ID,
tenant_id=_WORKSPACE_ID,
name="Site App",
description="",
mode=AppMode.CHAT,
icon_type=None,
icon=None,
icon_background=None,
enable_site=True,
enable_api=True,
)
session.add(app)
if with_site:
session.add(
Site(
app_id=_APP_ID,
title="Original",
description="Original description",
default_language="en-US",
input_placeholder="Original placeholder",
customize_token_strategy=CustomizeTokenStrategy.NOT_ALLOW,
prompt_public=False,
show_workflow_steps=True,
use_icon_as_answer_icon=False,
code="old-code",
)
)
session.commit()
def _repository(session_factory: sessionmaker[Session]) -> AppSiteCommandRepository:
return AppSiteCommandRepository(session_factory=session_factory)
def test_update_preserves_none_and_writes_false_and_empty_values(
sqlite_session: Session,
sqlite_session_factory: sessionmaker[Session],
) -> None:
_persist_app(sqlite_session)
result = _repository(sqlite_session_factory).update_site(
workspace_id=_WORKSPACE_ID,
app_id=_APP_ID,
actor_id=_ACTOR_ID,
changes=AppSiteChanges(
title=None,
input_placeholder="",
customize_token_strategy="allow",
show_workflow_steps=False,
),
)
assert result.title == "Original"
assert result.input_placeholder == ""
assert result.customize_token_strategy == "allow"
assert result.show_workflow_steps is False
with sqlite_session_factory() as session:
site = session.scalar(select(Site).where(Site.app_id == _APP_ID))
assert site is not None
assert site.title == "Original"
assert site.input_placeholder == ""
assert site.customize_token_strategy == CustomizeTokenStrategy.ALLOW
assert site.show_workflow_steps is False
assert site.updated_by == _ACTOR_ID
def test_update_scopes_app_to_workspace_and_distinguishes_missing_site(
sqlite_session: Session,
sqlite_session_factory: sessionmaker[Session],
) -> None:
_persist_app(sqlite_session)
repository = _repository(sqlite_session_factory)
with pytest.raises(AppSiteAppNotFoundError):
repository.update_site(
workspace_id=_OTHER_WORKSPACE_ID,
app_id=_APP_ID,
actor_id=_ACTOR_ID,
changes=AppSiteChanges(title="Leaked"),
)
with sqlite_session_factory.begin() as session:
session.execute(delete(Site).where(Site.app_id == _APP_ID))
with pytest.raises(AppSiteNotFoundError):
repository.update_site(
workspace_id=_WORKSPACE_ID,
app_id=_APP_ID,
actor_id=_ACTOR_ID,
changes=AppSiteChanges(title="Missing"),
)
def test_update_rolls_back_when_a_site_field_rejects_the_value(
sqlite_session: Session,
sqlite_session_factory: sessionmaker[Session],
) -> None:
_persist_app(sqlite_session)
with pytest.raises(ValueError, match="cannot exceed 512"):
_repository(sqlite_session_factory).update_site(
workspace_id=_WORKSPACE_ID,
app_id=_APP_ID,
actor_id=_ACTOR_ID,
changes=AppSiteChanges(title="Changed", custom_disclaimer="x" * 513),
)
with sqlite_session_factory() as session:
site = session.scalar(select(Site).where(Site.app_id == _APP_ID))
assert site is not None
assert site.title == "Original"
assert site.updated_by is None
def test_reset_access_token_uses_the_owned_transaction(
sqlite_session: Session,
sqlite_session_factory: sessionmaker[Session],
monkeypatch: pytest.MonkeyPatch,
) -> None:
_persist_app(sqlite_session)
observed_session: Session | None = None
def generate_code(length: int, *, session: Session) -> str:
nonlocal observed_session
observed_session = session
assert length == 16
assert session.in_transaction()
return "new-code"
monkeypatch.setattr(Site, "generate_code", generate_code)
result = _repository(sqlite_session_factory).reset_access_token(
workspace_id=_WORKSPACE_ID,
app_id=_APP_ID,
actor_id=_ACTOR_ID,
)
assert observed_session is not None
assert result.code == "new-code"
with sqlite_session_factory() as session:
site = session.scalar(select(Site).where(Site.app_id == _APP_ID))
assert site is not None
assert site.code == "new-code"
assert site.updated_by == _ACTOR_ID
@@ -0,0 +1,59 @@
from unittest.mock import MagicMock
import pytest
from machinery.context import RequestContext
from services.app_site_service import AppSiteChanges, AppSiteCommandResult, AppSiteService
def _command_result() -> AppSiteCommandResult:
return AppSiteCommandResult(
app_id="app-1",
code="site-code",
title="Site",
icon=None,
icon_background=None,
description=None,
default_language="en-US",
customize_domain=None,
copyright=None,
privacy_policy=None,
input_placeholder=None,
custom_disclaimer="",
customize_token_strategy="not_allow",
prompt_public=False,
show_workflow_steps=True,
use_icon_as_answer_icon=False,
)
def test_service_passes_stable_identity_to_store() -> None:
context = RequestContext("request-1", None, "account-1", "workspace-1")
changes = AppSiteChanges(title="Updated")
store = MagicMock()
store.update_site.return_value = _command_result()
store.reset_access_token.return_value = _command_result()
service = AppSiteService(sites=store)
assert service.update(context, "app-1", changes) == _command_result()
assert service.reset_access_token(context, "app-1") == _command_result()
store.update_site.assert_called_once_with(
workspace_id="workspace-1",
app_id="app-1",
actor_id="account-1",
changes=changes,
)
store.reset_access_token.assert_called_once_with(
workspace_id="workspace-1",
app_id="app-1",
actor_id="account-1",
)
def test_service_rejects_context_without_active_workspace() -> None:
context = RequestContext("request-1", None, "account-1", None)
service = AppSiteService(sites=MagicMock())
with pytest.raises(RuntimeError, match="active workspace"):
service.update(context, "app-1", AppSiteChanges())