127 lines
7.3 KiB
Python
127 lines
7.3 KiB
Python
"""Task-local benchmark: isolated localhost PostgreSQL only; no production writes.
|
|
|
|
Run from zhixing-server with DATABASE_URL pointing to a disposable test database.
|
|
Set PYTHONPATH to the before/after source tree to compare the same persisted data.
|
|
Use --seed once, then --run before.json / --run after.json (five reads per endpoint).
|
|
"""
|
|
import argparse
|
|
import hashlib
|
|
import json
|
|
import os
|
|
import runpy
|
|
import statistics
|
|
import time
|
|
from dataclasses import replace
|
|
from datetime import timedelta
|
|
from decimal import Decimal
|
|
from pathlib import Path
|
|
from urllib.parse import urlsplit
|
|
from unittest.mock import patch
|
|
|
|
import psycopg
|
|
from alembic import command
|
|
from alembic.config import Config
|
|
from zhixing_server.bootstrap.config import Settings, sqlalchemy_database_url
|
|
from zhixing_server.modules.sector_radar.application.details import ReadRadarDetails
|
|
from zhixing_server.modules.sector_radar.application.read import ReadSectorRadar, RadarQuery
|
|
from zhixing_server.modules.sector_radar.domain.models import PublicationStatus, SectorType, MetricKind, MetricUnit
|
|
from zhixing_server.modules.sector_radar.domain.metrics import AmountNetStrategy, RatioTurnoverStrategy, SwingEqualThreeToTenStrategy
|
|
from zhixing_server.modules.sector_radar.domain.persistence import PublicationSourceGroup as Group, PublicationSourceRecord, RankingRecord
|
|
from zhixing_server.modules.sector_radar.domain.ranking import rank_metric_observations
|
|
from zhixing_server.modules.sector_radar.domain.source import build_source_snapshot
|
|
from zhixing_server.modules.sector_radar.infrastructure.postgres import PostgresSectorRadarRepository
|
|
from zhixing_server.modules.sector_radar.presentation.http import _detail_response, _history_response, _rankings_response
|
|
|
|
parser = argparse.ArgumentParser()
|
|
parser.add_argument('--seed', action='store_true')
|
|
parser.add_argument('--run', type=Path)
|
|
args = parser.parse_args()
|
|
url = os.environ['DATABASE_URL']
|
|
assert urlsplit(url).hostname in {'localhost', '127.0.0.1'}, 'disposable local database only'
|
|
samples = runpy.run_path('tests/unit/sector_radar/test_detail_reads.py')
|
|
target, now, code = samples['TARGET'], samples['NOW'], samples['CODE']
|
|
repo = PostgresSectorRadarRepository(url)
|
|
|
|
if args.seed:
|
|
config = Config('alembic.ini')
|
|
config.set_main_option('sqlalchemy.url', sqlalchemy_database_url(url).replace('%', '%%'))
|
|
config.config_file_name = None
|
|
with patch('zhixing_server.bootstrap.config.get_settings', return_value=Settings(database_url=url)):
|
|
command.upgrade(config, 'head')
|
|
samples['seed_detail'](repo)
|
|
current = repo.get_successful_publication(target)
|
|
template = repo.load_rankings(current.publication_id)[0].observation
|
|
stocks = [f'{i:06d}.SZ' for i in range(10, 3010)]
|
|
sectors = [f'PERF{i:04d}.DC' for i in range(500)]
|
|
def save(group, rows, order, day=target):
|
|
snapshot = build_source_snapshot(api_name=group.value, params={'batch': str(order)}, rows=rows,
|
|
target_trade_date=day, observed_at=now)
|
|
repo.save_source_snapshots((snapshot,))
|
|
repo.save_publication_sources((PublicationSourceRecord(current.publication_id, group, order, snapshot),))
|
|
save(Group.STOCK_BASICS, tuple({'ts_code': s, 'symbol': s[:6], 'name': s, 'exchange': 'SZSE',
|
|
'list_status': 'L', 'list_date': '20200101'} for s in stocks), 1)
|
|
save(Group.CONCEPT_INDICES, tuple({'ts_code': s, 'name': s, 'trade_date': str(target),
|
|
'pct_change': '1.1234'} for s in sectors), 1)
|
|
save(Group.MEMBERS, tuple({'ts_code': sector, 'con_code': stocks[(i*7+j)%len(stocks)],
|
|
'name': stocks[(i*7+j)%len(stocks)], 'trade_date': str(target)}
|
|
for i, sector in enumerate(sectors) for j in range(100)), 1)
|
|
for offset in range(10):
|
|
day = target-timedelta(days=offset)
|
|
for group in (Group.DAILY, Group.MONEYFLOW_DC, Group.MONEYFLOW):
|
|
rows = tuple({'ts_code': stock, 'trade_date': str(day), 'name': stock,
|
|
'pct_chg': str(Decimal(i % 123)/100), 'amount': str(i*12),
|
|
'net_amount': str(i-1500), 'net_mf_amount': str(i-500),
|
|
'close': str(Decimal(i%1000)/10+1), 'vol': str(i*33)}
|
|
for i, stock in enumerate(stocks))
|
|
save(group, rows, offset+1, day)
|
|
versions = ((MetricKind.AMOUNT, AmountNetStrategy.metric_version, MetricUnit.CNY_100M),
|
|
(MetricKind.RATIO, RatioTurnoverStrategy.metric_version, MetricUnit.RATIO),
|
|
(MetricKind.SWING, SwingEqualThreeToTenStrategy.metric_version, MetricUnit.RATIO))
|
|
for offset in range(30):
|
|
day = target-timedelta(days=offset)
|
|
if offset < 2:
|
|
publication = repo.get_successful_publication(day)
|
|
else:
|
|
running = replace(current, publication_id=f'bench-{offset}', target_trade_date=day,
|
|
status=PublicationStatus.RUNNING, input_hash=None, finished_at=None)
|
|
repo.create_publication(running)
|
|
publication = replace(running, status=PublicationStatus.SUCCESS, input_hash='a'*64,
|
|
finished_at=now+timedelta(seconds=1))
|
|
repo.finish_publication(publication)
|
|
rankings = []
|
|
for kind, version, unit in versions:
|
|
observations = tuple(replace(template, trade_date=day, sector_code=sector, sector_name=sector,
|
|
metric_kind=kind, metric_version=version, unit=unit,
|
|
value=Decimal(i+1)) for i, sector in enumerate(sectors))
|
|
rankings.extend(RankingRecord(publication.publication_id, row)
|
|
for row in rank_metric_observations(observations))
|
|
repo.save_rankings(rankings)
|
|
with psycopg.connect(url) as connection:
|
|
connection.execute('ANALYZE')
|
|
print('dataset', connection.execute('SELECT count(*), sum(row_count), sum(octet_length(payload::text)) FROM sector_radar_source_snapshot').fetchone(),
|
|
'rankings', connection.execute('SELECT count(*) FROM sector_radar_ranking').fetchone())
|
|
|
|
if args.run:
|
|
repo.open()
|
|
reads = {
|
|
'detail': lambda: _detail_response(ReadRadarDetails(repo).detail(target, SectorType.CONCEPT, code)),
|
|
'history': lambda: _history_response(ReadRadarDetails(repo).history(target, SectorType.CONCEPT, code)),
|
|
'rankings': lambda: _rankings_response(ReadSectorRadar(repo).query(RadarQuery(trade_date=target, page_size=20))),
|
|
}
|
|
report = {}
|
|
for label, read in reads.items():
|
|
timings = []
|
|
responses = []
|
|
for _ in range(5):
|
|
start = time.perf_counter()
|
|
response = read().model_dump(mode='json')
|
|
timings.append(time.perf_counter()-start)
|
|
responses.append(response)
|
|
assert all(item == responses[0] for item in responses)
|
|
fingerprint = hashlib.sha256(json.dumps(responses[0], sort_keys=True).encode()).hexdigest()
|
|
report[label] = {'seconds': timings, 'median_seconds': statistics.median(timings),
|
|
'response_sha256': fingerprint, 'response': responses[0]}
|
|
print(label, 'median', report[label]['median_seconds'], 'sha256', fingerprint, flush=True)
|
|
args.run.write_text(json.dumps(report, ensure_ascii=False, indent=2))
|
|
repo.close()
|