mirror of
https://github.com/langgenius/dify.git
synced 2026-09-24 23:22:26 +08:00
fix(api): scope temporary capabilities to owners (#39863)
This commit is contained in:
@@ -491,6 +491,9 @@ class TestAppDslService:
|
||||
redis_key = f"{IMPORT_INFO_REDIS_KEY_PREFIX}{result.id}"
|
||||
stored = redis_client.get(redis_key)
|
||||
assert stored is not None
|
||||
pending = PendingData.model_validate_json(stored)
|
||||
assert pending.tenant_id == _DEFAULT_TENANT_ID
|
||||
assert pending.account_id == _DEFAULT_ACCOUNT_ID
|
||||
|
||||
def test_import_app_completed_uses_declared_dependencies(
|
||||
self, db_session_with_containers: Session, mock_external_service_dependencies
|
||||
@@ -606,7 +609,11 @@ class TestAppDslService:
|
||||
icon_background="#fff",
|
||||
app_id=None,
|
||||
)
|
||||
redis_client.setex(redis_key, IMPORT_INFO_REDIS_EXPIRY, pending.model_dump_json())
|
||||
redis_client.setex(
|
||||
redis_key,
|
||||
IMPORT_INFO_REDIS_EXPIRY,
|
||||
pending.model_dump_json(exclude={"tenant_id", "account_id"}),
|
||||
)
|
||||
|
||||
created_app = SimpleNamespace(
|
||||
id=str(uuid4()),
|
||||
|
||||
@@ -58,6 +58,7 @@ def _install_features(monkeypatch: pytest.MonkeyPatch, enabled: bool) -> None:
|
||||
def _make_account(account_id: str = "u1") -> Account:
|
||||
account = Account(name="Test User", email="test@example.com")
|
||||
account.id = account_id
|
||||
account._current_tenant = MagicMock(id="tenant-1")
|
||||
return account
|
||||
|
||||
|
||||
@@ -343,14 +344,13 @@ class TestAppImportConfirmApi:
|
||||
"current_account_with_tenant",
|
||||
lambda: (_make_account(), "tenant-1"),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
app_import_module.redis_client,
|
||||
"get",
|
||||
lambda *_args, **_kwargs: (
|
||||
redis_get = MagicMock(
|
||||
return_value=(
|
||||
b'{"import_mode":"yaml-content","yaml_content":"app: {}","app_id":null,'
|
||||
b'"name":null,"description":null,"icon_type":null,"icon":null,"icon_background":null}'
|
||||
),
|
||||
)
|
||||
)
|
||||
monkeypatch.setattr(app_import_module.redis_client, "get", redis_get)
|
||||
monkeypatch.setattr(app_import_module.dify_config, "RBAC_ENABLED", True)
|
||||
app_id = _install_persisting_service_result(
|
||||
monkeypatch,
|
||||
@@ -370,6 +370,7 @@ class TestAppImportConfirmApi:
|
||||
_assert_app_persistence(sqlite_app_engine, app_id, persisted=True)
|
||||
assert status == 200
|
||||
assert response["permission_keys"] == ["app.acl.view_layout", "app.acl.edit"]
|
||||
redis_get.assert_called_once_with("app_import_info:import-1")
|
||||
|
||||
def test_import_confirm_does_not_attach_permission_keys_when_overwriting_existing_app(
|
||||
self,
|
||||
|
||||
@@ -4,7 +4,7 @@ from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
from inspect import unwrap
|
||||
from unittest.mock import MagicMock, patch
|
||||
from unittest.mock import ANY, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from flask import Flask
|
||||
@@ -175,9 +175,15 @@ class TestTriggerSubscriptionBuilderApis:
|
||||
patch(
|
||||
"controllers.console.workspace.trigger_providers.TriggerSubscriptionBuilderService.get_subscription_builder_by_id",
|
||||
return_value=subscription_builder(),
|
||||
),
|
||||
) as mock_get_builder,
|
||||
):
|
||||
assert method(api, "github", "b1")["id"] == "b1"
|
||||
assert method(api, "t1", mock_user(), "github", "b1")["id"] == "b1"
|
||||
mock_get_builder.assert_called_once_with(
|
||||
tenant_id="t1",
|
||||
user_id="u1",
|
||||
provider_id=ANY,
|
||||
subscription_builder_id="b1",
|
||||
)
|
||||
|
||||
def test_verify_builder(self, app: Flask) -> None:
|
||||
api = TriggerSubscriptionBuilderVerifyApi()
|
||||
@@ -215,9 +221,16 @@ class TestTriggerSubscriptionBuilderApis:
|
||||
patch(
|
||||
"controllers.console.workspace.trigger_providers.TriggerSubscriptionBuilderService.update_trigger_subscription_builder",
|
||||
return_value=subscription_builder(),
|
||||
),
|
||||
) as mock_update_builder,
|
||||
):
|
||||
assert method(api, "t1", "github", "b1")["id"] == "b1"
|
||||
assert method(api, "t1", mock_user(), "github", "b1")["id"] == "b1"
|
||||
mock_update_builder.assert_called_once_with(
|
||||
tenant_id="t1",
|
||||
user_id="u1",
|
||||
provider_id=ANY,
|
||||
subscription_builder_id="b1",
|
||||
subscription_builder_updater=ANY,
|
||||
)
|
||||
|
||||
def test_logs(self, app: Flask) -> None:
|
||||
api = TriggerSubscriptionBuilderLogsApi()
|
||||
@@ -228,10 +241,16 @@ class TestTriggerSubscriptionBuilderApis:
|
||||
patch(
|
||||
"controllers.console.workspace.trigger_providers.TriggerSubscriptionBuilderService.list_logs",
|
||||
return_value=[request_log()],
|
||||
),
|
||||
) as mock_list_logs,
|
||||
):
|
||||
result = method(api, "github", "b1")
|
||||
result = method(api, "t1", mock_user(), "github", "b1")
|
||||
assert result["logs"][0]["id"] == "log1"
|
||||
mock_list_logs.assert_called_once_with(
|
||||
tenant_id="t1",
|
||||
user_id="u1",
|
||||
provider_id=ANY,
|
||||
subscription_builder_id="b1",
|
||||
)
|
||||
|
||||
def test_build(self, app: Flask) -> None:
|
||||
api = TriggerSubscriptionBuilderBuildApi()
|
||||
@@ -377,10 +396,17 @@ class TestTriggerOAuthApis:
|
||||
),
|
||||
patch(
|
||||
"controllers.console.workspace.trigger_providers.TriggerSubscriptionBuilderService.update_trigger_subscription_builder"
|
||||
),
|
||||
) as mock_update_builder,
|
||||
):
|
||||
resp = method(api, "github")
|
||||
assert resp.status_code == 302
|
||||
mock_update_builder.assert_called_once_with(
|
||||
tenant_id="t1",
|
||||
user_id="u1",
|
||||
provider_id=ANY,
|
||||
subscription_builder_id="b1",
|
||||
subscription_builder_updater=ANY,
|
||||
)
|
||||
|
||||
def test_oauth_callback_no_oauth_client(self, app: Flask) -> None:
|
||||
api = TriggerOAuthCallbackApi()
|
||||
|
||||
@@ -140,7 +140,7 @@ class TestGetDynamicSelectOptionsTrigger:
|
||||
result = PluginParameterService.get_dynamic_select_options(
|
||||
tenant_id="t1",
|
||||
user_id="u1",
|
||||
plugin_id="p1",
|
||||
plugin_id="org/plugin",
|
||||
provider="github",
|
||||
action="on_push",
|
||||
parameter="branch",
|
||||
@@ -149,6 +149,11 @@ class TestGetDynamicSelectOptionsTrigger:
|
||||
)
|
||||
|
||||
assert result == ["opt"]
|
||||
builder_call = mock_builder_svc.get_subscription_builder.call_args.kwargs
|
||||
assert builder_call["tenant_id"] == "t1"
|
||||
assert builder_call["user_id"] == "u1"
|
||||
assert str(builder_call["provider_id"]) == "org/plugin/github"
|
||||
assert builder_call["subscription_builder_id"] == "builder-1"
|
||||
|
||||
@patch("services.plugin.plugin_parameter_service.DynamicSelectClient")
|
||||
@patch("services.plugin.plugin_parameter_service.TriggerProviderService")
|
||||
@@ -166,7 +171,7 @@ class TestGetDynamicSelectOptionsTrigger:
|
||||
result = PluginParameterService.get_dynamic_select_options(
|
||||
tenant_id="t1",
|
||||
user_id="u1",
|
||||
plugin_id="p1",
|
||||
plugin_id="org/plugin",
|
||||
provider="github",
|
||||
action="on_push",
|
||||
parameter="branch",
|
||||
@@ -186,7 +191,7 @@ class TestGetDynamicSelectOptionsTrigger:
|
||||
PluginParameterService.get_dynamic_select_options(
|
||||
tenant_id="t1",
|
||||
user_id="u1",
|
||||
plugin_id="p1",
|
||||
plugin_id="org/plugin",
|
||||
provider="github",
|
||||
action="on_push",
|
||||
parameter="branch",
|
||||
|
||||
@@ -46,11 +46,11 @@ def service(sqlite_session: Session) -> RagPipelineDslService:
|
||||
return RagPipelineDslService(session=sqlite_session)
|
||||
|
||||
|
||||
def _account(*, tenant_id: str = "tenant-1") -> Account:
|
||||
def _account(*, tenant_id: str = "tenant-1", account_id: str = "account-1") -> Account:
|
||||
tenant = Tenant(name="Tenant")
|
||||
tenant.id = tenant_id
|
||||
account = Account(name="Account", email="account@example.com")
|
||||
account.id = "account-1"
|
||||
account.id = account_id
|
||||
account._current_tenant = tenant
|
||||
return account
|
||||
|
||||
@@ -627,6 +627,10 @@ def test_import_pending_version_stores_redis(monkeypatch: pytest.MonkeyPatch, se
|
||||
account=_account(), import_mode=ImportMode.YAML_CONTENT.value, yaml_content=_valid_dsl(version="1.0.0")
|
||||
)
|
||||
assert result.status == ImportStatus.PENDING
|
||||
assert setex.call_args.args[0] == f"app_import_info:{result.id}"
|
||||
pending = RagPipelinePendingData.model_validate_json(setex.call_args.args[2])
|
||||
assert pending.tenant_id == "tenant-1"
|
||||
assert pending.account_id == "account-1"
|
||||
setex.assert_called_once()
|
||||
|
||||
|
||||
@@ -782,14 +786,26 @@ def test_confirm_import_updates_tenant_pipeline_and_dataset(
|
||||
dataset = _dataset(sqlite_session, pipeline)
|
||||
_workflow(sqlite_session, pipeline)
|
||||
pending = RagPipelinePendingData(
|
||||
tenant_id="tenant-1",
|
||||
account_id="account-1",
|
||||
import_mode=ImportMode.YAML_CONTENT.value,
|
||||
yaml_content=_valid_dsl(name="Confirmed"),
|
||||
pipeline_id=pipeline.id,
|
||||
)
|
||||
monkeypatch.setattr(module.redis_client, "get", Mock(return_value=pending.model_dump_json()))
|
||||
redis_key = "app_import_info:import-1"
|
||||
monkeypatch.setattr(
|
||||
module.redis_client,
|
||||
"get",
|
||||
Mock(side_effect=lambda key: pending.model_dump_json() if key == redis_key else None),
|
||||
)
|
||||
delete = Mock()
|
||||
monkeypatch.setattr(module.redis_client, "delete", delete)
|
||||
monkeypatch.setattr(module.KnowledgeConfiguration, "model_validate", Mock(return_value=_knowledge_configuration()))
|
||||
for foreign_account in (_account(tenant_id="tenant-2"), _account(account_id="account-2")):
|
||||
assert service.confirm_import(import_id="import-1", account=foreign_account).status == ImportStatus.FAILED
|
||||
delete.assert_not_called()
|
||||
assert pipeline.name == "Pipeline"
|
||||
|
||||
result = service.confirm_import(import_id="import-1", account=_account())
|
||||
assert result.status == ImportStatus.COMPLETED
|
||||
assert result.pipeline_id == pipeline.id
|
||||
@@ -809,7 +825,7 @@ def test_confirm_import_updates_tenant_pipeline_and_dataset(
|
||||
assert observed_pipeline is not None
|
||||
assert observed_pipeline.name == "Confirmed"
|
||||
|
||||
delete.assert_called_once()
|
||||
delete.assert_called_once_with(redis_key)
|
||||
|
||||
|
||||
def test_export_reads_real_dataset_and_workflow_and_filters_credentials(
|
||||
|
||||
@@ -12,7 +12,7 @@ from core.workflow.llm_environment_variable import LLMEnvironmentVariable
|
||||
from models import App, AppMode
|
||||
from models.model import AppModelConfig, AppModelConfigDict, IconType
|
||||
from models.workflow import Workflow
|
||||
from services.app_dsl_service import AppDslService
|
||||
from services.app_dsl_service import AppDslService, PendingData
|
||||
from services.entities.dsl_entities import ImportStatus
|
||||
from services.errors.account import NoPermissionError
|
||||
|
||||
@@ -144,6 +144,78 @@ def test_import_app_returns_decode_error_for_invalid_yaml_url_bytes(
|
||||
assert not unbound_session.in_transaction()
|
||||
|
||||
|
||||
def test_pending_import_is_scoped_to_its_owner(monkeypatch: pytest.MonkeyPatch, unbound_session: Session) -> None:
|
||||
pending_imports: dict[str, str] = {}
|
||||
monkeypatch.setattr(
|
||||
"services.app_dsl_service.redis_client.setex",
|
||||
lambda key, _expiry, value: pending_imports.__setitem__(key, value),
|
||||
)
|
||||
service = AppDslService(session=unbound_session)
|
||||
creator = Mock(id="account-1", current_tenant_id="tenant-1")
|
||||
|
||||
pending = service.import_app(
|
||||
account=creator,
|
||||
import_mode="yaml-content",
|
||||
yaml_content="version: 99.0.0\nkind: app\napp: {name: Test, mode: workflow}\n",
|
||||
)
|
||||
|
||||
redis_key = f"app_import_info:{pending.id}"
|
||||
assert pending.status == ImportStatus.PENDING
|
||||
assert redis_key in pending_imports
|
||||
pending_data = PendingData.model_validate_json(pending_imports[redis_key])
|
||||
assert pending_data.tenant_id == "tenant-1"
|
||||
assert pending_data.account_id == "account-1"
|
||||
|
||||
monkeypatch.setattr("services.app_dsl_service.redis_client.get", pending_imports.get)
|
||||
monkeypatch.setattr("services.app_dsl_service.redis_client.delete", pending_imports.pop)
|
||||
monkeypatch.setattr(
|
||||
service,
|
||||
"_create_or_update_app",
|
||||
Mock(return_value=Mock(id="app-1", mode=AppMode.WORKFLOW)),
|
||||
)
|
||||
|
||||
for other_account in (
|
||||
Mock(id="account-1", current_tenant_id="tenant-2"),
|
||||
Mock(id="account-2", current_tenant_id="tenant-1"),
|
||||
):
|
||||
assert service.confirm_import(import_id=pending.id, account=other_account).status == ImportStatus.FAILED
|
||||
|
||||
assert service.confirm_import(import_id=pending.id, account=creator).status == ImportStatus.COMPLETED
|
||||
assert redis_key not in pending_imports
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("tenant_id", "account_id", "expected"),
|
||||
[
|
||||
("tenant-1", "account-1", True),
|
||||
(None, "account-1", False),
|
||||
("tenant-1", None, False),
|
||||
("tenant-2", "account-1", False),
|
||||
("tenant-1", "account-2", False),
|
||||
],
|
||||
)
|
||||
def test_pending_import_owner_access(
|
||||
tenant_id: str | None,
|
||||
account_id: str | None,
|
||||
expected: bool,
|
||||
) -> None:
|
||||
pending = PendingData(
|
||||
tenant_id=tenant_id,
|
||||
account_id=account_id,
|
||||
import_mode="yaml-content",
|
||||
yaml_content="",
|
||||
)
|
||||
|
||||
assert pending.is_accessible_by(tenant_id="tenant-1", account_id="account-1") is expected
|
||||
|
||||
|
||||
def test_pending_import_owner_access_accepts_legacy_json() -> None:
|
||||
pending = PendingData.model_validate_json('{"import_mode":"yaml-content","yaml_content":""}')
|
||||
|
||||
assert pending.is_accessible_by(tenant_id="tenant-1", account_id="account-1")
|
||||
assert not pending.is_accessible_by(tenant_id=None, account_id="account-1")
|
||||
|
||||
|
||||
def test_create_or_update_app_loads_existing_model_config_with_service_session(
|
||||
sqlite_session_factory: sessionmaker[Session],
|
||||
) -> None:
|
||||
|
||||
@@ -269,7 +269,7 @@ workflow:
|
||||
"""
|
||||
|
||||
result = service.import_snippet(
|
||||
account=SimpleNamespace(current_tenant_id="tenant-1"),
|
||||
account=SimpleNamespace(id="account-1", current_tenant_id="tenant-1"),
|
||||
import_mode=ImportMode.YAML_CONTENT.value,
|
||||
yaml_content=yaml_content,
|
||||
name="Override",
|
||||
@@ -278,7 +278,10 @@ workflow:
|
||||
|
||||
assert result.status == ImportStatus.PENDING
|
||||
setex.assert_called_once()
|
||||
assert setex.call_args.args[0] == f"snippet_import_info:{result.id}"
|
||||
pending = SnippetPendingData.model_validate_json(setex.call_args.args[2])
|
||||
assert pending.tenant_id == "tenant-1"
|
||||
assert pending.account_id == "account-1"
|
||||
assert pending.name == "Override"
|
||||
assert pending.description == "Override description"
|
||||
|
||||
@@ -359,7 +362,9 @@ def test_confirm_import_returns_failed_when_pending_data_missing(monkeypatch):
|
||||
service = SnippetDslService(session=SimpleNamespace())
|
||||
monkeypatch.setattr("services.snippet_dsl_service.redis_client.get", Mock(return_value=None))
|
||||
|
||||
result = service.confirm_import(import_id="missing", account=SimpleNamespace(current_tenant_id="tenant-1"))
|
||||
result = service.confirm_import(
|
||||
import_id="missing", account=SimpleNamespace(id="account-1", current_tenant_id="tenant-1")
|
||||
)
|
||||
|
||||
assert result.status == ImportStatus.FAILED
|
||||
assert result.error == "Import information expired or does not exist"
|
||||
@@ -369,13 +374,15 @@ def test_confirm_import_returns_failed_for_invalid_pending_payload(monkeypatch):
|
||||
service = SnippetDslService(session=SimpleNamespace())
|
||||
monkeypatch.setattr("services.snippet_dsl_service.redis_client.get", Mock(return_value=object()))
|
||||
|
||||
result = service.confirm_import(import_id="bad", account=SimpleNamespace(current_tenant_id="tenant-1"))
|
||||
result = service.confirm_import(
|
||||
import_id="bad", account=SimpleNamespace(id="account-1", current_tenant_id="tenant-1")
|
||||
)
|
||||
|
||||
assert result.status == ImportStatus.FAILED
|
||||
assert result.error == "Invalid import information"
|
||||
|
||||
|
||||
def test_confirm_import_creates_snippet_from_pending_data(monkeypatch):
|
||||
def test_confirm_import_is_scoped_to_its_owner(monkeypatch):
|
||||
service = SnippetDslService(session=SimpleNamespace(scalar=Mock(return_value=None)))
|
||||
account = SimpleNamespace(id="account-1", current_tenant_id="tenant-1")
|
||||
snippet = SimpleNamespace(id="snippet-new")
|
||||
@@ -391,6 +398,8 @@ workflow:
|
||||
edges: []
|
||||
"""
|
||||
pending = SnippetPendingData(
|
||||
tenant_id="tenant-1",
|
||||
account_id="account-1",
|
||||
import_mode="yaml-content",
|
||||
yaml_content=yaml_content,
|
||||
name="Override name",
|
||||
@@ -399,10 +408,21 @@ workflow:
|
||||
)
|
||||
create_or_update = Mock(return_value=snippet)
|
||||
monkeypatch.setattr(service, "_create_or_update_snippet", create_or_update)
|
||||
monkeypatch.setattr("services.snippet_dsl_service.redis_client.get", Mock(return_value=pending.model_dump_json()))
|
||||
redis_key = "snippet_import_info:import-1"
|
||||
monkeypatch.setattr(
|
||||
"services.snippet_dsl_service.redis_client.get",
|
||||
Mock(side_effect=lambda key: pending.model_dump_json() if key == redis_key else None),
|
||||
)
|
||||
redis_delete = Mock()
|
||||
monkeypatch.setattr("services.snippet_dsl_service.redis_client.delete", redis_delete)
|
||||
|
||||
for other_account in (
|
||||
SimpleNamespace(id="account-1", current_tenant_id="tenant-2"),
|
||||
SimpleNamespace(id="account-2", current_tenant_id="tenant-1"),
|
||||
):
|
||||
assert service.confirm_import(import_id="import-1", account=other_account).status == ImportStatus.FAILED
|
||||
|
||||
create_or_update.assert_not_called()
|
||||
result = service.confirm_import(import_id="import-1", account=account)
|
||||
|
||||
assert result.status == ImportStatus.COMPLETED
|
||||
@@ -414,7 +434,7 @@ workflow:
|
||||
assert kwargs["account"] is account
|
||||
assert kwargs["name"] == "Override name"
|
||||
assert kwargs["description"] == "Override description"
|
||||
redis_delete.assert_called_once_with("snippet_import_info:import-1")
|
||||
redis_delete.assert_called_once_with(redis_key)
|
||||
|
||||
|
||||
def test_confirm_import_returns_failed_for_non_mapping_yaml(monkeypatch):
|
||||
@@ -426,7 +446,9 @@ def test_confirm_import_returns_failed_for_non_mapping_yaml(monkeypatch):
|
||||
)
|
||||
monkeypatch.setattr("services.snippet_dsl_service.redis_client.get", Mock(return_value=pending.model_dump_json()))
|
||||
|
||||
result = service.confirm_import(import_id="import-1", account=SimpleNamespace(current_tenant_id="tenant-1"))
|
||||
result = service.confirm_import(
|
||||
import_id="import-1", account=SimpleNamespace(id="account-1", current_tenant_id="tenant-1")
|
||||
)
|
||||
|
||||
assert result.status == ImportStatus.FAILED
|
||||
assert result.error == "Invalid YAML format: expected a dictionary"
|
||||
@@ -445,7 +467,7 @@ def test_confirm_import_returns_failed_when_create_or_update_raises(monkeypatch)
|
||||
|
||||
result = service.confirm_import(
|
||||
import_id="import-1",
|
||||
account=SimpleNamespace(current_tenant_id="tenant-1"),
|
||||
account=SimpleNamespace(id="account-1", current_tenant_id="tenant-1"),
|
||||
)
|
||||
|
||||
assert result.status == ImportStatus.FAILED
|
||||
|
||||
@@ -0,0 +1,251 @@
|
||||
from collections.abc import Callable
|
||||
from contextlib import nullcontext
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from core.plugin.entities.plugin_daemon import CredentialType
|
||||
from core.trigger.entities.entities import SubscriptionBuilder, SubscriptionBuilderUpdater
|
||||
from core.trigger.trigger_manager import TriggerManager
|
||||
from models.provider_ids import TriggerProviderID
|
||||
from services.trigger.trigger_subscription_builder_service import TriggerSubscriptionBuilderService
|
||||
|
||||
PROVIDER_ID = TriggerProviderID("org/plugin/provider")
|
||||
|
||||
|
||||
def subscription_builder() -> SubscriptionBuilder:
|
||||
return SubscriptionBuilder(
|
||||
id="builder-1",
|
||||
name="Builder",
|
||||
tenant_id="tenant-1",
|
||||
user_id="user-1",
|
||||
provider_id=str(PROVIDER_ID),
|
||||
endpoint_id="builder-1",
|
||||
parameters={},
|
||||
properties={},
|
||||
credentials={},
|
||||
credential_type=CredentialType.UNAUTHORIZED,
|
||||
credential_expires_at=-1,
|
||||
expires_at=-1,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("field", "value"),
|
||||
[
|
||||
("id", "other-builder"),
|
||||
("tenant_id", "other-tenant"),
|
||||
("user_id", "other-user"),
|
||||
("provider_id", "org/plugin/other"),
|
||||
],
|
||||
)
|
||||
def test_get_subscription_builder_rejects_non_owner(field: str, value: str) -> None:
|
||||
builder = subscription_builder().model_copy(update={field: value})
|
||||
with patch.object(
|
||||
TriggerSubscriptionBuilderService,
|
||||
"_get_subscription_builder_by_endpoint_id",
|
||||
return_value=builder,
|
||||
):
|
||||
with pytest.raises(ValueError, match="not found"):
|
||||
TriggerSubscriptionBuilderService.get_subscription_builder(
|
||||
tenant_id="tenant-1",
|
||||
user_id="user-1",
|
||||
provider_id=PROVIDER_ID,
|
||||
subscription_builder_id="builder-1",
|
||||
)
|
||||
|
||||
|
||||
def test_get_subscription_builder_accepts_owner() -> None:
|
||||
builder = subscription_builder()
|
||||
with patch.object(
|
||||
TriggerSubscriptionBuilderService,
|
||||
"_get_subscription_builder_by_endpoint_id",
|
||||
return_value=builder,
|
||||
):
|
||||
assert (
|
||||
TriggerSubscriptionBuilderService.get_subscription_builder(
|
||||
tenant_id="tenant-1",
|
||||
user_id="user-1",
|
||||
provider_id=PROVIDER_ID,
|
||||
subscription_builder_id="builder-1",
|
||||
)
|
||||
is builder
|
||||
)
|
||||
|
||||
|
||||
def test_get_subscription_builder_returns_none_when_temporary_builder_is_absent() -> None:
|
||||
with patch.object(
|
||||
TriggerSubscriptionBuilderService,
|
||||
"_get_subscription_builder_by_endpoint_id",
|
||||
return_value=None,
|
||||
):
|
||||
assert (
|
||||
TriggerSubscriptionBuilderService.get_subscription_builder(
|
||||
tenant_id="tenant-1",
|
||||
user_id="user-1",
|
||||
provider_id=PROVIDER_ID,
|
||||
subscription_builder_id="builder-1",
|
||||
)
|
||||
is None
|
||||
)
|
||||
|
||||
|
||||
def test_get_subscription_builder_rejects_mismatched_endpoint() -> None:
|
||||
builder = subscription_builder().model_copy(update={"endpoint_id": "other-builder"})
|
||||
with patch(
|
||||
"services.trigger.trigger_subscription_builder_service.redis_client.get",
|
||||
return_value=builder.model_dump_json(),
|
||||
):
|
||||
assert TriggerSubscriptionBuilderService._get_subscription_builder_by_endpoint_id("builder-1") is None
|
||||
|
||||
|
||||
def test_get_subscription_builder_accepts_matching_endpoint() -> None:
|
||||
builder = subscription_builder()
|
||||
with patch(
|
||||
"services.trigger.trigger_subscription_builder_service.redis_client.get",
|
||||
return_value=builder.model_dump_json(),
|
||||
):
|
||||
assert TriggerSubscriptionBuilderService._get_subscription_builder_by_endpoint_id("builder-1") == builder
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("operation", "needs_updater"),
|
||||
[
|
||||
(TriggerSubscriptionBuilderService.update_trigger_subscription_builder, True),
|
||||
(TriggerSubscriptionBuilderService.update_and_verify_builder, True),
|
||||
(TriggerSubscriptionBuilderService.update_and_build_builder, True),
|
||||
(TriggerSubscriptionBuilderService.list_logs, False),
|
||||
(TriggerSubscriptionBuilderService.get_subscription_builder_by_id, False),
|
||||
],
|
||||
)
|
||||
def test_owner_scoped_operations_reject_missing_builder_before_side_effects(
|
||||
operation: Callable[..., object], needs_updater: bool
|
||||
) -> None:
|
||||
kwargs: dict[str, object] = {
|
||||
"tenant_id": "tenant-1",
|
||||
"user_id": "user-1",
|
||||
"provider_id": PROVIDER_ID,
|
||||
"subscription_builder_id": "builder-1",
|
||||
}
|
||||
if needs_updater:
|
||||
kwargs["subscription_builder_updater"] = SubscriptionBuilderUpdater(name="Updated")
|
||||
|
||||
with (
|
||||
patch.object(TriggerManager, "get_trigger_provider", return_value=Mock()),
|
||||
patch.object(TriggerSubscriptionBuilderService, "acquire_builder_lock", return_value=nullcontext()),
|
||||
patch.object(
|
||||
TriggerSubscriptionBuilderService,
|
||||
"get_subscription_builder",
|
||||
return_value=None,
|
||||
) as get_subscription_builder,
|
||||
patch("services.trigger.trigger_subscription_builder_service.redis_client.setex") as setex,
|
||||
patch("services.trigger.trigger_subscription_builder_service.redis_client.delete") as delete,
|
||||
):
|
||||
with pytest.raises(ValueError, match="not found"):
|
||||
operation(**kwargs)
|
||||
|
||||
get_subscription_builder.assert_called_once_with(
|
||||
tenant_id="tenant-1",
|
||||
user_id="user-1",
|
||||
provider_id=PROVIDER_ID,
|
||||
subscription_builder_id="builder-1",
|
||||
)
|
||||
setex.assert_not_called()
|
||||
delete.assert_not_called()
|
||||
|
||||
|
||||
def test_update_and_build_uses_the_owned_builder_without_refetching() -> None:
|
||||
builder = subscription_builder()
|
||||
cache_key = TriggerSubscriptionBuilderService.encode_cache_key(builder.id)
|
||||
|
||||
with (
|
||||
patch.object(TriggerManager, "get_trigger_provider", return_value=Mock()),
|
||||
patch.object(TriggerSubscriptionBuilderService, "acquire_builder_lock", return_value=nullcontext()),
|
||||
patch.object(
|
||||
TriggerSubscriptionBuilderService,
|
||||
"get_subscription_builder",
|
||||
return_value=builder,
|
||||
) as get_subscription_builder,
|
||||
patch("services.trigger.trigger_subscription_builder_service.redis_client.setex") as setex,
|
||||
patch(
|
||||
"services.trigger.trigger_subscription_builder_service.TriggerProviderService.add_trigger_subscription"
|
||||
) as add_subscription,
|
||||
patch("services.trigger.trigger_subscription_builder_service.redis_client.delete") as delete,
|
||||
):
|
||||
TriggerSubscriptionBuilderService.update_and_build_builder(
|
||||
tenant_id="tenant-1",
|
||||
user_id="user-1",
|
||||
provider_id=PROVIDER_ID,
|
||||
subscription_builder_id=builder.id,
|
||||
subscription_builder_updater=SubscriptionBuilderUpdater(name="Updated"),
|
||||
)
|
||||
|
||||
get_subscription_builder.assert_called_once_with(
|
||||
tenant_id="tenant-1",
|
||||
user_id="user-1",
|
||||
provider_id=PROVIDER_ID,
|
||||
subscription_builder_id=builder.id,
|
||||
)
|
||||
setex.assert_called_once_with(cache_key, 30 * 60, builder.model_dump_json())
|
||||
add_subscription.assert_called_once()
|
||||
subscription_call = add_subscription.call_args.kwargs
|
||||
assert subscription_call["subscription_id"] == builder.id
|
||||
assert subscription_call["tenant_id"] == "tenant-1"
|
||||
assert subscription_call["user_id"] == "user-1"
|
||||
assert subscription_call["provider_id"] == PROVIDER_ID
|
||||
assert subscription_call["endpoint_id"] == builder.endpoint_id
|
||||
assert subscription_call["name"] == "Updated"
|
||||
delete.assert_called_once_with(cache_key)
|
||||
|
||||
|
||||
def test_list_logs_uses_the_owned_builder_endpoint() -> None:
|
||||
builder = subscription_builder()
|
||||
logs_key = f"trigger:subscription:builder:logs:{builder.endpoint_id}"
|
||||
|
||||
with (
|
||||
patch.object(TriggerSubscriptionBuilderService, "get_subscription_builder", return_value=builder),
|
||||
patch("services.trigger.trigger_subscription_builder_service.redis_client.get", return_value=None) as redis_get,
|
||||
):
|
||||
assert (
|
||||
TriggerSubscriptionBuilderService.list_logs(
|
||||
tenant_id="tenant-1",
|
||||
user_id="user-1",
|
||||
provider_id=PROVIDER_ID,
|
||||
subscription_builder_id=builder.id,
|
||||
)
|
||||
== []
|
||||
)
|
||||
|
||||
redis_get.assert_called_once_with(logs_key)
|
||||
|
||||
|
||||
def test_process_validation_endpoint_uses_the_public_capability() -> None:
|
||||
builder = subscription_builder()
|
||||
request = Mock()
|
||||
response = Mock()
|
||||
controller = Mock()
|
||||
controller.dispatch.return_value = Mock(response=response)
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
TriggerSubscriptionBuilderService,
|
||||
"_get_subscription_builder_by_endpoint_id",
|
||||
return_value=builder,
|
||||
) as get_by_endpoint,
|
||||
patch.object(TriggerSubscriptionBuilderService, "get_subscription_builder") as get_owned,
|
||||
patch.object(TriggerManager, "get_trigger_provider", return_value=controller) as get_provider,
|
||||
patch.object(TriggerSubscriptionBuilderService, "append_log") as append_log,
|
||||
):
|
||||
assert (
|
||||
TriggerSubscriptionBuilderService.process_builder_validation_endpoint(builder.endpoint_id, request)
|
||||
is response
|
||||
)
|
||||
|
||||
get_by_endpoint.assert_called_once_with(builder.endpoint_id)
|
||||
get_owned.assert_not_called()
|
||||
get_provider.assert_called_once()
|
||||
provider_call = get_provider.call_args.kwargs
|
||||
assert provider_call["tenant_id"] == builder.tenant_id
|
||||
assert str(provider_call["provider_id"]) == str(PROVIDER_ID)
|
||||
controller.dispatch.assert_called_once()
|
||||
append_log.assert_called_once()
|
||||
Reference in New Issue
Block a user