Files
Grid_Trading/tests/test_services.py

301 lines
9.8 KiB
Python
Raw Normal View History

2026-07-08 17:25:37 +08:00
from datetime import date, timedelta
from decimal import Decimal
import pytest
2026-07-08 22:02:56 +08:00
from grid_trading.domain.models import Instrument, QuoteSnapshot, Trade, TradeGroup, TradeSide
2026-07-08 17:25:37 +08:00
from grid_trading.services.trading_service import TradingService
2026-07-08 22:02:56 +08:00
class FakeQuoteProvider:
def fetch_quotes(self, instruments):
return {
"510300": QuoteSnapshot(
symbol="sh510300",
code="510300",
name="沪深300ETF",
price=Decimal("4.12"),
source="tencent",
quote_time="20260708150000",
)
}
2026-07-08 17:25:37 +08:00
def test_service_creates_default_account_and_computes_summary(tmp_path):
service = TradingService(tmp_path / "grid.db")
service.ensure_defaults()
account = service.get_active_account()
assert account.name == "默认账户"
account = service.save_account(account.__class__(id=account.id, name="主账户", initial_cash=Decimal("100000")))
instrument = service.add_instrument(
Instrument(id=None, code="510300", name="沪深300ETF", market="ETF", manual_price=Decimal("4.00"))
)
service.save_trade(
Trade(
id=None,
account_id=account.id,
instrument_id=instrument.id,
trade_date=date(2026, 7, 7),
side=TradeSide.BUY,
price=Decimal("3.90"),
quantity=1000,
commission=Decimal("5"),
trade_group=TradeGroup.GRID,
)
)
service.save_trade(
Trade(
id=None,
account_id=account.id,
instrument_id=instrument.id,
trade_date=date(2026, 7, 8),
side=TradeSide.SELL,
price=Decimal("4.10"),
quantity=500,
commission=Decimal("5"),
stamp_tax=Decimal("1.03"),
trade_group=TradeGroup.GRID,
)
)
positions = service.get_position_summaries(as_of=date(2026, 7, 8))
summary = service.get_account_summary(as_of=date(2026, 7, 8))
assert len(positions) == 1
assert positions[0].total_quantity == 500
assert positions[0].current_price is None
assert positions[0].price_source == "missing"
2026-07-08 17:25:37 +08:00
assert positions[0].grid_profit == Decimal("91.47")
assert summary.cash == Decimal("98138.97")
assert summary.market_value == Decimal("0.00")
assert summary.total_assets == Decimal("98138.97")
2026-07-08 17:25:37 +08:00
def test_service_validates_lot_size_and_available_sell_quantity(tmp_path):
service = TradingService(tmp_path / "grid.db")
service.ensure_defaults()
account = service.get_active_account()
instrument = service.add_instrument(Instrument(id=None, code="600000", name="浦发银行"))
with pytest.raises(ValueError, match="100"):
service.save_trade(
Trade(
id=None,
account_id=account.id,
instrument_id=instrument.id,
trade_date=date(2026, 7, 8),
side=TradeSide.BUY,
price=Decimal("10"),
quantity=50,
trade_group=TradeGroup.BASE,
)
)
service.save_trade(
Trade(
id=None,
account_id=account.id,
instrument_id=instrument.id,
trade_date=date(2026, 7, 8),
side=TradeSide.BUY,
price=Decimal("10"),
quantity=100,
trade_group=TradeGroup.BASE,
)
)
with pytest.raises(ValueError, match="T\\+1"):
service.save_trade(
Trade(
id=None,
account_id=account.id,
instrument_id=instrument.id,
trade_date=date(2026, 7, 8),
side=TradeSide.SELL,
price=Decimal("10.1"),
quantity=100,
trade_group=TradeGroup.BASE,
)
)
service.save_trade(
Trade(
id=None,
account_id=account.id,
instrument_id=instrument.id,
trade_date=date(2026, 7, 8) - timedelta(days=1),
side=TradeSide.BUY,
price=Decimal("9.9"),
quantity=100,
trade_group=TradeGroup.BASE,
)
)
saved_sell = service.save_trade(
Trade(
id=None,
account_id=account.id,
instrument_id=instrument.id,
trade_date=date(2026, 7, 8),
side=TradeSide.SELL,
price=Decimal("10.1"),
quantity=100,
trade_group=TradeGroup.BASE,
)
)
assert saved_sell.id is not None
def test_service_rejects_historical_changes_that_break_future_sells(tmp_path):
service = TradingService(tmp_path / "grid.db")
service.ensure_defaults()
account = service.get_active_account()
instrument = service.add_instrument(Instrument(id=None, code="159915", name="创业板ETF"))
buy = service.save_trade(
Trade(
id=None,
account_id=account.id,
instrument_id=instrument.id,
trade_date=date(2026, 7, 7),
side=TradeSide.BUY,
price=Decimal("2.00"),
quantity=200,
trade_group=TradeGroup.GRID,
)
)
service.save_trade(
Trade(
id=None,
account_id=account.id,
instrument_id=instrument.id,
trade_date=date(2026, 7, 8),
side=TradeSide.SELL,
price=Decimal("2.10"),
quantity=200,
trade_group=TradeGroup.GRID,
)
)
with pytest.raises(ValueError, match="后续成交"):
service.update_trade(
Trade(
id=buy.id,
account_id=account.id,
instrument_id=instrument.id,
trade_date=date(2026, 7, 7),
side=TradeSide.BUY,
price=Decimal("2.00"),
quantity=100,
trade_group=TradeGroup.GRID,
)
)
with pytest.raises(ValueError, match="后续成交"):
service.delete_trade(buy.id)
2026-07-08 22:02:56 +08:00
def test_service_refresh_quotes_uses_realtime_price_in_summaries(tmp_path):
service = TradingService(tmp_path / "grid.db", quote_provider=FakeQuoteProvider())
service.ensure_defaults()
account = service.get_active_account()
instrument = service.add_instrument(
Instrument(id=None, code="510300", name="沪深300ETF", market="ETF", manual_price=Decimal("3.90"))
)
service.save_trade(
Trade(
id=None,
account_id=account.id,
instrument_id=instrument.id,
trade_date=date(2026, 7, 7),
side=TradeSide.BUY,
price=Decimal("4.00"),
quantity=1000,
trade_group=TradeGroup.BASE,
)
)
quotes = service.refresh_quotes()
[position] = service.get_position_summaries(as_of=date(2026, 7, 8))
summary = service.get_account_summary(as_of=date(2026, 7, 8))
assert quotes["510300"].price == Decimal("4.12")
assert position.current_price == Decimal("4.12")
assert position.price_source == "tencent"
assert summary.market_value == Decimal("4120.00")
2026-07-09 10:50:25 +08:00
def test_service_generates_grid_level_suggestions_from_realtime_price(tmp_path):
service = TradingService(tmp_path / "grid.db", quote_provider=FakeQuoteProvider())
service.ensure_defaults()
account = service.get_active_account()
instrument = service.add_instrument(Instrument(id=None, code="510300", name="沪深300ETF", market="ETF"))
service.save_trade(
Trade(
id=None,
account_id=account.id,
instrument_id=instrument.id,
trade_date=date(2026, 7, 7),
side=TradeSide.BUY,
price=Decimal("4.00"),
quantity=1000,
trade_group=TradeGroup.BASE,
)
)
service.refresh_quotes()
levels = service.get_grid_level_suggestions(instrument.id, levels=2)
assert [item.buy_price for item in levels] == [Decimal("4.00"), Decimal("3.88")]
assert [item.sell_price for item in levels] == [Decimal("4.12"), Decimal("4.00")]
assert [item.buy_amount for item in levels] == [Decimal("5000.00"), Decimal("5000.00")]
assert [item.suggested_quantity for item in levels] == [1200, 1200]
def test_service_returns_empty_grid_levels_without_realtime_price(tmp_path):
service = TradingService(tmp_path / "grid.db")
service.ensure_defaults()
instrument = service.add_instrument(Instrument(id=None, code="510300", name="沪深300ETF", market="ETF"))
assert service.get_grid_level_suggestions(instrument.id, levels=10) == []
2026-07-09 11:17:51 +08:00
def test_service_returns_open_grid_lots_with_suggested_sell_price(tmp_path):
service = TradingService(tmp_path / "grid.db", quote_provider=FakeQuoteProvider())
service.ensure_defaults()
account = service.get_active_account()
instrument = service.add_instrument(Instrument(id=None, code="510300", name="沪深300ETF", market="ETF"))
service.save_trade(
Trade(
id=None,
account_id=account.id,
instrument_id=instrument.id,
trade_date=date(2026, 7, 7),
side=TradeSide.BUY,
price=Decimal("4.00"),
quantity=1000,
trade_group=TradeGroup.GRID,
)
)
service.save_trade(
Trade(
id=None,
account_id=account.id,
instrument_id=instrument.id,
trade_date=date(2026, 7, 8),
side=TradeSide.SELL,
price=Decimal("4.12"),
quantity=400,
trade_group=TradeGroup.GRID,
)
)
service.refresh_quotes()
[lot] = service.get_open_grid_lots(instrument.id, as_of=date(2026, 7, 9))
assert lot.buy_price == Decimal("4.00")
assert lot.remaining_quantity == 600
assert lot.suggested_sell_price == Decimal("4.12")
assert lot.current_price == Decimal("4.12")
assert lot.status == "可卖"