from datetime import date, timedelta from decimal import Decimal import pytest from grid_trading.domain.calculations import CalculationError, calculate_positions, estimate_fees from grid_trading.domain.models import FeeRules, Instrument, QuoteSnapshot, Trade, TradeGroup, TradeSide def make_trade( *, trade_id: int, instrument_id: int = 1, trade_date: date, side: TradeSide, price: str, quantity: int, commission: str = "0", stamp_tax: str = "0", transfer_fee: str = "0", trade_group: TradeGroup = TradeGroup.GRID, ) -> Trade: return Trade( id=trade_id, account_id=1, instrument_id=instrument_id, trade_date=trade_date, side=side, price=Decimal(price), quantity=quantity, commission=Decimal(commission), stamp_tax=Decimal(stamp_tax), transfer_fee=Decimal(transfer_fee), trade_group=trade_group, ) def test_buy_sell_grid_profit_and_breakeven(): today = date(2026, 7, 8) instrument = Instrument(id=1, code="510300", name="沪深300ETF") trades = [ make_trade( trade_id=1, trade_date=today - timedelta(days=2), side=TradeSide.BUY, price="10", quantity=100, commission="1", trade_group=TradeGroup.GRID, ), make_trade( trade_id=2, trade_date=today - timedelta(days=1), side=TradeSide.SELL, price="11", quantity=100, commission="1", trade_group=TradeGroup.GRID, ), make_trade( trade_id=3, trade_date=today, side=TradeSide.BUY, price="9", quantity=100, commission="1", trade_group=TradeGroup.BASE, ), ] quotes = { 1: QuoteSnapshot( symbol="sh510300", code="510300", name="沪深300ETF", price=Decimal("9.50"), source="tencent", ) } [summary] = calculate_positions([instrument], trades, as_of=today, quote_snapshots=quotes) assert summary.total_quantity == 100 assert summary.base_quantity == 100 assert summary.grid_quantity == 0 assert summary.remaining_cost == Decimal("901.00") assert summary.realized_pnl == Decimal("98.00") assert summary.grid_profit == Decimal("98.00") assert summary.cost_price == Decimal("9.01") assert summary.position_breakeven_price == Decimal("8.03") assert summary.account_breakeven_price == Decimal("8.03") assert summary.current_price == Decimal("9.50") assert summary.price_source == "tencent" assert summary.market_value == Decimal("950.00") assert summary.floating_pnl == Decimal("49.00") def test_t_plus_one_available_quantity_excludes_today_buys(): today = date(2026, 7, 8) instrument = Instrument(id=1, code="600000", name="浦发银行", manual_price=Decimal("10")) trades = [ make_trade( trade_id=1, trade_date=today - timedelta(days=1), side=TradeSide.BUY, price="10", quantity=200, trade_group=TradeGroup.BASE, ), make_trade( trade_id=2, trade_date=today, side=TradeSide.BUY, price="9.8", quantity=100, trade_group=TradeGroup.BASE, ), ] [summary] = calculate_positions([instrument], trades, as_of=today) assert summary.total_quantity == 300 assert summary.available_quantity == 200 def test_quote_snapshot_supplies_current_price_for_position_value(): today = date(2026, 7, 8) instrument = Instrument(id=1, code="510300", name="沪深300ETF", manual_price=Decimal("3.90")) trades = [ make_trade( trade_id=1, trade_date=today - timedelta(days=1), side=TradeSide.BUY, price="4.00", quantity=1000, trade_group=TradeGroup.BASE, ) ] quotes = { 1: QuoteSnapshot( symbol="sh510300", code="510300", name="沪深300ETF", price=Decimal("4.12"), source="tencent", ) } [summary] = calculate_positions([instrument], trades, as_of=today, quote_snapshots=quotes) assert summary.current_price == Decimal("4.12") assert summary.price_source == "tencent" assert summary.market_value == Decimal("4120.00") def test_missing_quote_does_not_use_manual_or_last_trade_price_for_current_price(): today = date(2026, 7, 8) instrument = Instrument(id=1, code="510300", name="沪深300ETF", manual_price=Decimal("3.90")) trades = [ make_trade( trade_id=1, trade_date=today - timedelta(days=1), side=TradeSide.BUY, price="4.00", quantity=1000, trade_group=TradeGroup.BASE, ) ] [summary] = calculate_positions([instrument], trades, as_of=today) assert summary.current_price is None assert summary.price_source == "missing" assert summary.market_value is None assert summary.floating_pnl is None def test_sell_more_than_group_position_raises(): today = date(2026, 7, 8) instrument = Instrument(id=1, code="159915", name="创业板ETF") trades = [ make_trade( trade_id=1, trade_date=today - timedelta(days=1), side=TradeSide.BUY, price="2", quantity=100, trade_group=TradeGroup.GRID, ), make_trade( trade_id=2, trade_date=today, side=TradeSide.SELL, price="2.1", quantity=200, trade_group=TradeGroup.GRID, ), ] with pytest.raises(CalculationError, match="Insufficient grid position"): calculate_positions([instrument], trades, as_of=today) def test_estimate_fees_uses_min_commission_and_sell_tax(): rules = FeeRules( commission_rate=Decimal("0.00025"), min_commission=Decimal("5"), stamp_tax_rate=Decimal("0.0005"), transfer_fee_rate=Decimal("0.00001"), ) buy_fees = estimate_fees(TradeSide.BUY, Decimal("10"), 100, rules) sell_fees = estimate_fees(TradeSide.SELL, Decimal("10"), 100, rules) assert buy_fees.commission == Decimal("5.00") assert buy_fees.stamp_tax == Decimal("0.00") assert buy_fees.transfer_fee == Decimal("0.01") assert sell_fees.commission == Decimal("5.00") assert sell_fees.stamp_tax == Decimal("0.50") assert sell_fees.transfer_fee == Decimal("0.01")