Files
zhixing-system/.trellis/tasks/archive/2026-09/09-07-api-performance-diagnosis/research/benchmark.py
T

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()