test_runner.py 4.02 KB
from unittest.mock import MagicMock

from etl_to_crawler import runner


class _Cursor:
    def __enter__(self):
        return self

    def __exit__(self, exc_type, exc, tb):
        return False

    def execute(self, *args, **kwargs):
        return None


class _PgConnection:
    def __init__(self):
        self.cur = _Cursor()
        self.commits = 0
        self.closed = False

    def cursor(self):
        return self.cur

    def commit(self):
        self.commits += 1

    def close(self):
        self.closed = True


class _Connection:
    def close(self):
        return None


def test_run_imports_all_platform_records_but_writes_one_primary_yinyan_relation(monkeypatch):
    pg_conn = _PgConnection()
    processors = {
        '1': MagicMock(return_value={'platform': 'qq', 'platform_song_id': 100, 'mid': 'qq-mid', 'title': '歌'}),
        '2': MagicMock(return_value={'platform': 'kugou', 'platform_song_id': 200, 'hash': 'kg-hash', 'title': '歌'}),
    }
    yinyan_writer = MagicMock()

    monkeypatch.setattr(runner, 'get_hk_songs_conn', lambda: _Connection())
    monkeypatch.setattr(runner, 'get_source_conn', lambda: _Connection())
    monkeypatch.setattr(runner, 'get_spider_conn', lambda: _Connection())
    monkeypatch.setattr(runner, 'get_pg_conn', lambda: pg_conn)
    monkeypatch.setattr(runner, 'get_oss_bucket', lambda: object())
    monkeypatch.setattr(runner, 'iter_hk_songs_batches', lambda conn, batch_size, start_after_id=0: [[{
        'source_song_id': 10,
        'name': '歌',
        'audio_url': 'https://example.com/a.mp3',
        'singer': '歌手',
    }]])
    monkeypatch.setattr(runner, 'fetch_platform_records', lambda conn, song_ids: [
        {
            'source_song_id': 10,
            'record_id': 100,
            'platform': '1',
            'platform_unique_key': 'qq-mid',
            'platform_mid': '100',
            'album_audio_id': None,
            'is_main_version': 0,
            'is_high': 1,
            'pub_time': '2020-01-01',
        },
        {
            'source_song_id': 10,
            'record_id': 200,
            'platform': '2',
            'platform_unique_key': '200',
            'platform_mid': 'kg-hash',
            'album_audio_id': None,
            'is_main_version': 1,
            'is_high': 0,
            'pub_time': '2021-01-01',
        },
    ])
    monkeypatch.setattr(runner, '_PROCESSORS', processors)
    monkeypatch.setattr(runner, 'upsert_yinyan_song_records', yinyan_writer)

    runner.run(['1', '2'])

    processors['1'].assert_called_once()
    processors['2'].assert_called_once()
    yinyan_writer.assert_called_once_with(pg_conn.cur, [(10, 200)])
    assert pg_conn.commits == 1


def test_run_resume_starts_after_saved_id_and_updates_state_after_commit(monkeypatch, tmp_path):
    pg_conn = _PgConnection()
    state_file = tmp_path / 'etl_state.json'
    state_file.write_text('{"all": {"last_hk_songs_id": 40}}', encoding='utf-8')
    seen_start_ids = []

    monkeypatch.setattr(runner, 'get_hk_songs_conn', lambda: _Connection())
    monkeypatch.setattr(runner, 'get_source_conn', lambda: _Connection())
    monkeypatch.setattr(runner, 'get_spider_conn', lambda: _Connection())
    monkeypatch.setattr(runner, 'get_pg_conn', lambda: pg_conn)
    monkeypatch.setattr(runner, 'get_oss_bucket', lambda: object())

    def fake_batches(conn, batch_size, start_after_id=0):
        seen_start_ids.append(start_after_id)
        return iter([[
            {'id': 50, 'source_song_id': 10, 'name': '歌', 'audio_url': 'https://example.com/a.mp3', 'singer': '歌手'},
            {'id': 60, 'source_song_id': 20, 'name': '歌2', 'audio_url': 'https://example.com/b.mp3', 'singer': '歌手2'},
        ]])

    monkeypatch.setattr(runner, 'iter_hk_songs_batches', fake_batches)
    monkeypatch.setattr(runner, 'fetch_platform_records', lambda conn, song_ids: [])

    runner.run(['1', '2', '4'], resume=True, state_file=state_file, state_key='all')

    assert seen_start_ids == [40]
    assert pg_conn.commits == 1
    assert '"last_hk_songs_id": 60' in state_file.read_text(encoding='utf-8')