perf(selection): 优化选股执行性能
This commit is contained in:
@@ -1,5 +1,7 @@
|
||||
"""PostgreSQL reader contract tests using a fake connection."""
|
||||
|
||||
from collections.abc import Generator
|
||||
from contextlib import contextmanager
|
||||
from datetime import date
|
||||
from decimal import Decimal
|
||||
from typing import cast
|
||||
@@ -7,6 +9,8 @@ from typing import cast
|
||||
import psycopg
|
||||
import pytest
|
||||
|
||||
from zhixing_server.modules.selection.domain.runs import SelectionStock
|
||||
from zhixing_server.modules.selection.infrastructure.postgres_pool import SelectionPostgresPool
|
||||
from zhixing_server.modules.selection.infrastructure.postgres_reader import (
|
||||
PostgresMarketDataReader,
|
||||
SelectionMarketDataNotReady,
|
||||
@@ -39,6 +43,23 @@ class FakeResult:
|
||||
return self.rows
|
||||
|
||||
|
||||
class Pool:
|
||||
def __init__(self, connection: FakeConnection) -> None:
|
||||
self.connection_value = connection
|
||||
self.opened = 0
|
||||
|
||||
def open(self, *, wait: bool = True) -> None:
|
||||
assert wait is True
|
||||
self.opened += 1
|
||||
|
||||
def close(self) -> None:
|
||||
return None
|
||||
|
||||
@contextmanager
|
||||
def connection(self) -> Generator[FakeConnection, None, None]:
|
||||
yield self.connection_value
|
||||
|
||||
|
||||
def test_reader_parameterizes_target_and_maps_left_join(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
connection = FakeConnection(
|
||||
[
|
||||
@@ -86,6 +107,64 @@ def test_reader_parameterizes_target_and_maps_left_join(monkeypatch: pytest.Monk
|
||||
assert "trade_date <= %s" in cast(str, connection.query)
|
||||
|
||||
|
||||
def test_reader_batches_qfq_rows_by_stock_without_historical_basic_join() -> None:
|
||||
connection = FakeConnection(
|
||||
[
|
||||
(
|
||||
"600000.SH",
|
||||
"浦发银行",
|
||||
date(2024, 1, 2),
|
||||
"8",
|
||||
"8.5",
|
||||
"7.8",
|
||||
"8.2",
|
||||
"900",
|
||||
),
|
||||
(
|
||||
"000001.SZ",
|
||||
"平安银行",
|
||||
date(2024, 1, 3),
|
||||
"10.5",
|
||||
"11",
|
||||
"10",
|
||||
"10.8",
|
||||
"1200",
|
||||
),
|
||||
(
|
||||
"000001.SZ",
|
||||
"平安银行",
|
||||
date(2024, 1, 2),
|
||||
"10",
|
||||
"11",
|
||||
"9",
|
||||
"10.5",
|
||||
"1000",
|
||||
),
|
||||
]
|
||||
)
|
||||
pool = Pool(connection)
|
||||
owner = SelectionPostgresPool("postgresql://test", max_connections=6, pool=pool)
|
||||
|
||||
histories = PostgresMarketDataReader("postgresql://test", pool=owner).load_histories(
|
||||
(
|
||||
SelectionStock("000001.SZ", "平安银行"),
|
||||
SelectionStock("600000.SH", "浦发银行"),
|
||||
),
|
||||
date(2024, 1, 3),
|
||||
)
|
||||
|
||||
assert [history.ts_code for history in histories] == ["000001.SZ", "600000.SH"]
|
||||
assert [bar.trade_date for bar in histories[0].bars] == [
|
||||
date(2024, 1, 2),
|
||||
date(2024, 1, 3),
|
||||
]
|
||||
assert histories[0].daily_basic == {}
|
||||
assert connection.parameters == (["000001.SZ", "600000.SH"], date(2024, 1, 3))
|
||||
assert "bar.ts_code = ANY(%s)" in cast(str, connection.query)
|
||||
assert "market_daily_basic" not in cast(str, connection.query)
|
||||
assert pool.opened == 1
|
||||
|
||||
|
||||
class SourceConnection:
|
||||
def __init__(self, source_row: tuple[object, ...] | None) -> None:
|
||||
self.source_row = source_row
|
||||
|
||||
Reference in New Issue
Block a user