| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307 |
- from flask import Flask
- from app.core.data_source.models import (
- DataSourceCredential,
- DataSourceDefinition,
- SealedCredential,
- )
- UID = "01900000-0000-7000-8000-000000000010"
- class FakeSession:
- def __init__(self):
- self.commits = 0
- self.rollbacks = 0
- def commit(self):
- self.commits += 1
- def rollback(self):
- self.rollbacks += 1
- class FakeDefinitions:
- def __init__(self, existing=None, fail_save=False):
- self.existing = existing
- self.saved = None
- self.deleted = None
- self.fail_save = fail_save
- def get(self, uid):
- return self.existing if uid == UID else None
- def list(self, _filters):
- return [self.existing] if self.existing else []
- def save(self, definition):
- if self.fail_save:
- raise RuntimeError("neo4j unavailable")
- self.saved = definition
- self.existing = definition
- return definition
- def delete(self, uid):
- self.deleted = uid
- return True
- class FakeCredentials:
- def __init__(self):
- self.created = []
- self.compensations = []
- self.revoked = []
- def get_active(self, session, uid, version=None):
- del session, uid, version
- return DataSourceCredential("existing-reader", "existing-secret")
- def create_version(
- self,
- session,
- *,
- data_source_uid,
- credential,
- actor_uid,
- ):
- del session, actor_uid
- version = len(self.created) + 2
- self.created.append(credential)
- return SealedCredential(
- id=UID,
- data_source_uid=data_source_uid,
- credential_version=version,
- encrypted_payload=b"cipher",
- nonce=b"n" * 12,
- key_version="v1",
- )
- def compensate_failed_activation(
- self,
- session,
- *,
- data_source_uid,
- failed_version,
- restore_version,
- actor_uid,
- ):
- del session, actor_uid
- self.compensations.append(
- (data_source_uid, failed_version, restore_version)
- )
- def revoke_all(self, session, *, data_source_uid, actor_uid):
- del session, actor_uid
- self.revoked.append(data_source_uid)
- return 1
- class FakeAdapter:
- def __init__(self):
- self.tests = []
- def validate_options(self, options):
- return dict(options)
- def test_connection(
- self,
- definition,
- credential,
- *,
- query_timeout,
- ):
- self.tests.append((definition, credential, query_timeout))
- class FakeManager:
- def __init__(self):
- self.invalidations = []
- def invalidate(self, uid, reason):
- self.invalidations.append((uid, reason))
- def existing_definition():
- return DataSourceDefinition(
- uid=UID,
- name_en="warehouse",
- database_type="postgresql",
- host="old-db",
- port=5432,
- database="analytics",
- credential_ref=UID,
- credential_version=1,
- )
- def make_service(existing=None, fail_save=False):
- from app.core.data_source.service import DataSourceService
- session = FakeSession()
- definitions = FakeDefinitions(existing=existing, fail_save=fail_save)
- credentials = FakeCredentials()
- adapter = FakeAdapter()
- manager = FakeManager()
- events = []
- service = DataSourceService(
- definitions=definitions,
- credentials=credentials,
- platform_session=lambda: session,
- adapter_resolver=lambda _database_type: adapter,
- connection_manager=manager,
- settings_resolver=lambda _overrides: {"query_timeout": 30},
- outbox_enqueuer=lambda *_args, **kwargs: events.append(kwargs),
- )
- return (
- service,
- session,
- definitions,
- credentials,
- adapter,
- manager,
- events,
- )
- def test_create_stores_secret_only_in_credential_store():
- (
- service,
- session,
- definitions,
- credentials,
- adapter,
- manager,
- events,
- ) = make_service()
- result, created = service.save(
- {
- "name_en": "warehouse",
- "name_zh": "数仓",
- "type": "postgresql",
- "host": "source-postgres",
- "port": 5432,
- "database": "analytics",
- "schema": "public",
- "username": "reader",
- "password": "source-pass",
- },
- actor_uid=None,
- )
- assert created is True
- assert result.uid
- assert definitions.saved.credential_version == 2
- assert credentials.created[0].password == "source-pass"
- assert "password" not in repr(definitions.saved)
- assert adapter.tests[0][2] == 30
- assert manager.invalidations == []
- assert session.commits == 1
- assert events[0]["event_type"] == "datasource.credential_version_created"
- def test_update_reuses_omitted_credentials_and_invalidates_old_pool():
- (
- service,
- _session,
- _definitions,
- credentials,
- _adapter,
- manager,
- _events,
- ) = make_service(existing=existing_definition())
- result, created = service.save(
- {
- "uid": UID,
- "name_en": "warehouse",
- "type": "postgresql",
- "host": "new-db",
- "port": 5432,
- "database": "analytics",
- },
- actor_uid=None,
- )
- assert created is False
- assert result.credential_version == 2
- assert credentials.created[0] == DataSourceCredential(
- "existing-reader",
- "existing-secret",
- )
- assert manager.invalidations == [(UID, "configuration_changed")]
- def test_neo4j_failure_restores_prior_active_credential():
- (
- service,
- session,
- _definitions,
- credentials,
- _adapter,
- _manager,
- _events,
- ) = make_service(existing=existing_definition(), fail_save=True)
- try:
- service.save(
- {
- "uid": UID,
- "name_en": "warehouse",
- "type": "postgresql",
- "host": "new-db",
- "port": 5432,
- "database": "analytics",
- },
- actor_uid=None,
- )
- except Exception:
- pass
- else: # pragma: no cover - failure is required
- raise AssertionError("save should fail")
- assert credentials.compensations == [(UID, 2, 1)]
- assert session.commits == 2
- def test_routes_use_service_and_return_real_http_status(monkeypatch):
- from app.api.data_source import routes
- class FakeService:
- def save(self, payload, actor_uid):
- del payload, actor_uid
- return existing_definition(), True
- def serialize(self, definition):
- return {
- "uid": definition.uid,
- "name_en": definition.name_en,
- "credential_configured": True,
- }
- monkeypatch.setattr(
- routes,
- "get_data_source_service",
- lambda: FakeService(),
- )
- app = Flask(__name__)
- app.register_blueprint(routes.bp, url_prefix="/api/datasource")
- response = app.test_client().post(
- "/api/datasource/save",
- json={"name_en": "warehouse"},
- )
- assert response.status_code == 201
- assert response.get_json()["data"]["uid"] == UID
- def test_datasource_routes_do_not_construct_engines_directly():
- source = open(
- "app/api/data_source/routes.py",
- encoding="utf-8",
- ).read()
- assert "create_engine" not in source
- assert "URL.create" not in source
|