feat(selection): 迁移知行B1选股策略
This commit is contained in:
@@ -0,0 +1,84 @@
|
||||
"""PostgreSQL reader contract tests using a fake connection."""
|
||||
|
||||
from datetime import date
|
||||
from typing import cast
|
||||
|
||||
import psycopg
|
||||
import pytest
|
||||
|
||||
from zhixing_server.modules.selection.infrastructure.postgres_reader import (
|
||||
PostgresMarketDataReader,
|
||||
)
|
||||
|
||||
|
||||
class FakeConnection:
|
||||
def __init__(self, rows: list[tuple[object, ...]]) -> None:
|
||||
self.rows = rows
|
||||
self.query: str | None = None
|
||||
self.parameters: tuple[object, ...] | None = None
|
||||
|
||||
def __enter__(self) -> "FakeConnection":
|
||||
return self
|
||||
|
||||
def __exit__(self, *args: object) -> None:
|
||||
return None
|
||||
|
||||
def execute(self, query: str, parameters: tuple[object, ...]) -> "FakeResult":
|
||||
self.query = query
|
||||
self.parameters = parameters
|
||||
return FakeResult(self.rows)
|
||||
|
||||
|
||||
class FakeResult:
|
||||
def __init__(self, rows: list[tuple[object, ...]]) -> None:
|
||||
self.rows = rows
|
||||
|
||||
def fetchall(self) -> list[tuple[object, ...]]:
|
||||
return self.rows
|
||||
|
||||
|
||||
def test_reader_parameterizes_target_and_maps_left_join(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
connection = FakeConnection(
|
||||
[
|
||||
(
|
||||
"000001.SZ",
|
||||
"平安银行",
|
||||
date(2024, 1, 2),
|
||||
"10",
|
||||
"11",
|
||||
"9",
|
||||
"10.5",
|
||||
"1000",
|
||||
None,
|
||||
None,
|
||||
),
|
||||
(
|
||||
"000001.SZ",
|
||||
"平安银行",
|
||||
date(2024, 1, 3),
|
||||
"10.5",
|
||||
"11",
|
||||
"10",
|
||||
"10.8",
|
||||
"1200",
|
||||
"1.2",
|
||||
"100000",
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
def connect(database_url: str) -> FakeConnection:
|
||||
assert database_url == "postgresql://test"
|
||||
return connection
|
||||
|
||||
monkeypatch.setattr(psycopg, "connect", connect)
|
||||
history = PostgresMarketDataReader("postgresql://test").load_history(
|
||||
"000001.SZ", date(2024, 1, 3)
|
||||
)
|
||||
|
||||
assert [bar.trade_date for bar in history.bars] == [date(2024, 1, 2), date(2024, 1, 3)]
|
||||
assert history.daily_basic[date(2024, 1, 2)].turnover_rate is None
|
||||
assert history.daily_basic[date(2024, 1, 3)].total_mv == 100000.0
|
||||
assert connection.parameters == ("000001.SZ", date(2024, 1, 3))
|
||||
assert "source_adj = 'qfq'" in cast(str, connection.query)
|
||||
assert "trade_date <= %s" in cast(str, connection.query)
|
||||
Reference in New Issue
Block a user