test: use sqlite3 session in test_reset_encrypt_key_pair (#38673)

This commit is contained in:
Asuka Minato 2026-07-22 18:17:34 +09:00 committed by GitHub
parent 03fad2e041
commit 187501f53e
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194

View File

@ -1,15 +1,24 @@
"""Unit tests for the reset-encrypt-key-pair CLI command (#35396).
"""SQLite-backed tests for the reset-encrypt-key-pair CLI command (#35396).
The command must purge every table that stores ciphertext encrypted with the
tenant's asymmetric key, otherwise stale rows cause downstream API failures
such as `/console/api/workspaces/current/tool-providers` returning 500.
Tests bind the command-owned transaction to the fixture engine and assert the
committed state rather than inspecting fabricated ``Session.execute`` calls.
"""
from unittest.mock import MagicMock, patch
from types import SimpleNamespace
import pytest
from sqlalchemy import select
from sqlalchemy.orm import Session
import commands
from commands import system as system_commands
from models.provider import Provider, ProviderModel
from core.tools.entities.tool_entities import ApiProviderSchemaType
from graphon.model_runtime.entities.model_entities import ModelType
from models import Tenant
from models.provider import Provider, ProviderModel, ProviderType
from models.tools import ApiToolProvider, BuiltinToolProvider, MCPToolProvider
@ -21,17 +30,60 @@ def _invoke_reset() -> int:
return 0
def _delete_targets(session_mock: MagicMock) -> list:
"""Extract the model class targeted by each `delete(...)` call on the session."""
targets = []
for call in session_mock.execute.call_args_list:
stmt = call.args[0]
# `delete(Foo)` constructs a `Delete` statement whose entity is `Foo`.
try:
targets.append(stmt.table.name)
except AttributeError:
targets.append(repr(stmt))
return targets
TENANT_ID = "11111111-1111-1111-1111-111111111111"
OTHER_TENANT_ID = "11111111-1111-1111-1111-111111111112"
USER_ID = "22222222-2222-2222-2222-222222222222"
def _tenant(tenant_id: str, *, name: str = "Test tenant") -> Tenant:
tenant = Tenant(name=name, encrypt_public_key="old-key")
tenant.id = tenant_id
return tenant
def _encrypted_rows(tenant_id: str, *, suffix: str = "1") -> tuple[object, ...]:
"""Build one persisted credential-bearing row for every purge target."""
return (
Provider(tenant_id=tenant_id, provider_name=f"provider-{suffix}"),
ProviderModel(
tenant_id=tenant_id,
provider_name=f"provider-{suffix}",
model_name=f"model-{suffix}",
model_type=ModelType.LLM,
),
BuiltinToolProvider(
name=f"builtin-credential-{suffix}",
tenant_id=tenant_id,
user_id=USER_ID,
provider=f"builtin-{suffix}",
encrypted_credentials="ciphertext",
),
ApiToolProvider(
name=f"api-{suffix}",
icon="icon",
schema="{}",
schema_type_str=ApiProviderSchemaType.OPENAPI,
user_id=USER_ID,
tenant_id=tenant_id,
description="description",
tools_str="[]",
credentials_str="{}",
),
MCPToolProvider(
name=f"mcp-{suffix}",
server_identifier=f"server-{suffix}",
server_url="ciphertext",
server_url_hash=f"hash-{suffix}",
icon=None,
tenant_id=tenant_id,
user_id=USER_ID,
encrypted_credentials="ciphertext",
),
)
def _bind_command_to_sqlite(monkeypatch: pytest.MonkeyPatch, session: Session) -> None:
monkeypatch.setattr(system_commands, "db", SimpleNamespace(engine=session.get_bind()))
def test_reset_aborts_when_not_self_hosted(monkeypatch, capsys):
@ -44,65 +96,73 @@ def test_reset_aborts_when_not_self_hosted(monkeypatch, capsys):
assert "only for SELF_HOSTED" in captured.out
def test_reset_purges_provider_and_tool_tables_for_each_tenant(monkeypatch, capsys):
@pytest.mark.parametrize(
"sqlite_session",
[(Tenant, Provider, ProviderModel, BuiltinToolProvider, ApiToolProvider, MCPToolProvider)],
indirect=True,
)
def test_reset_purges_provider_and_tool_tables_for_each_tenant(
monkeypatch: pytest.MonkeyPatch, capsys: pytest.CaptureFixture[str], sqlite_session: Session
) -> None:
"""The command must purge LLM provider rows AND every tool provider table
that stores ciphertext encrypted under the tenant key (#35396)."""
monkeypatch.setattr(system_commands.dify_config, "EDITION", "SELF_HOSTED")
monkeypatch.setattr(system_commands, "generate_key_pair", lambda tenant_id: f"new-key-{tenant_id}")
_bind_command_to_sqlite(monkeypatch, sqlite_session)
fake_tenant = MagicMock(id="tenant-abc", encrypt_public_key="old-key")
session = MagicMock()
session.scalars.return_value.all.return_value = [fake_tenant]
tenant = _tenant(TENANT_ID)
other_tenant = _tenant(OTHER_TENANT_ID, name="Other tenant")
system_provider = Provider(
tenant_id=TENANT_ID,
provider_name="system-provider",
provider_type=ProviderType.SYSTEM,
)
sqlite_session.add_all((tenant, other_tenant, system_provider, *_encrypted_rows(TENANT_ID)))
sqlite_session.commit()
fake_sessionmaker = MagicMock()
fake_sessionmaker.begin.return_value.__enter__.return_value = session
fake_sessionmaker.begin.return_value.__exit__.return_value = False
with (
patch.object(system_commands, "db", MagicMock()),
patch.object(system_commands, "sessionmaker", return_value=fake_sessionmaker),
):
exit_code = _invoke_reset()
exit_code = _invoke_reset()
captured = capsys.readouterr()
assert exit_code == 0
assert "tenant-abc" in captured.out
assert TENANT_ID in captured.out
# New key pair generated and assigned.
assert fake_tenant.encrypt_public_key == "new-key-tenant-abc"
# Every encrypted-credential table should have been purged for this tenant.
table_names = _delete_targets(session)
expected = {
Provider.__tablename__,
ProviderModel.__tablename__,
BuiltinToolProvider.__tablename__,
ApiToolProvider.__tablename__,
MCPToolProvider.__tablename__,
}
assert expected.issubset(set(table_names)), f"missing purges: expected {expected}, got {table_names}"
sqlite_session.expire_all()
assert sqlite_session.get(Tenant, TENANT_ID).encrypt_public_key == f"new-key-{TENANT_ID}"
assert sqlite_session.get(Tenant, OTHER_TENANT_ID).encrypt_public_key == f"new-key-{OTHER_TENANT_ID}"
assert sqlite_session.scalars(select(Provider).where(Provider.provider_type == ProviderType.CUSTOM)).all() == []
assert sqlite_session.scalars(select(ProviderModel)).all() == []
assert sqlite_session.scalars(select(BuiltinToolProvider)).all() == []
assert sqlite_session.scalars(select(ApiToolProvider)).all() == []
assert sqlite_session.scalars(select(MCPToolProvider)).all() == []
assert (
sqlite_session.scalar(select(Provider).where(Provider.provider_type == ProviderType.SYSTEM)) is system_provider
)
def test_reset_iterates_all_tenants(monkeypatch, capsys):
@pytest.mark.parametrize(
"sqlite_session",
[(Tenant, Provider, ProviderModel, BuiltinToolProvider, ApiToolProvider, MCPToolProvider)],
indirect=True,
)
def test_reset_iterates_all_tenants(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session) -> None:
"""Multi-tenant deployments must purge every tenant, not just the first."""
monkeypatch.setattr(system_commands.dify_config, "EDITION", "SELF_HOSTED")
monkeypatch.setattr(system_commands, "generate_key_pair", lambda tenant_id: f"new-key-{tenant_id}")
tenants = [MagicMock(id=f"tenant-{i}", encrypt_public_key="old") for i in range(3)]
session = MagicMock()
session.scalars.return_value.all.return_value = tenants
_bind_command_to_sqlite(monkeypatch, sqlite_session)
tenant_ids = [f"11111111-1111-1111-1111-{index:012d}" for index in range(3)]
tenants = [_tenant(tenant_id, name=f"Tenant {index}") for index, tenant_id in enumerate(tenant_ids)]
for index, tenant in enumerate(tenants):
sqlite_session.add(tenant)
sqlite_session.add_all(_encrypted_rows(tenant.id, suffix=str(index)))
sqlite_session.commit()
fake_sessionmaker = MagicMock()
fake_sessionmaker.begin.return_value.__enter__.return_value = session
fake_sessionmaker.begin.return_value.__exit__.return_value = False
assert _invoke_reset() == 0
with (
patch.object(system_commands, "db", MagicMock()),
patch.object(system_commands, "sessionmaker", return_value=fake_sessionmaker),
):
_invoke_reset()
# Five purges per tenant × 3 tenants = 15 execute calls.
assert session.execute.call_count == 15
for tenant in tenants:
assert tenant.encrypt_public_key == f"new-key-{tenant.id}"
sqlite_session.expire_all()
persisted_tenants = sqlite_session.scalars(select(Tenant).order_by(Tenant.id)).all()
assert [tenant.encrypt_public_key for tenant in persisted_tenants] == [
f"new-key-{tenant_id}" for tenant_id in tenant_ids
]
for model in (Provider, ProviderModel, BuiltinToolProvider, ApiToolProvider, MCPToolProvider):
assert sqlite_session.scalars(select(model)).all() == []