feat: add trading service workflow

This commit is contained in:
王鹏
2026-07-08 17:25:37 +08:00
parent 0f4dfead6f
commit f282d4a38b
3 changed files with 366 additions and 0 deletions

View File

@@ -0,0 +1 @@
"""Application services."""

View File

@@ -0,0 +1,234 @@
from __future__ import annotations
from collections import defaultdict
from datetime import date
from decimal import Decimal
from pathlib import Path
from grid_trading.config import DEFAULT_FEE_CONFIG
from grid_trading.domain.calculations import (
CalculationError,
calculate_account_summary,
calculate_positions,
estimate_fees,
)
from grid_trading.domain.models import (
Account,
AccountSummary,
FeeEstimate,
FeeRules,
Instrument,
PositionSummary,
StrategyOverride,
StrategyTemplate,
Trade,
TradeSide,
)
from grid_trading.storage.repositories import Repository
class TradingService:
def __init__(self, db_path: str | Path):
self.repository = Repository(db_path)
self.repository.initialize()
def close(self) -> None:
self.repository.close()
def ensure_defaults(self) -> None:
if not self.repository.list_accounts():
self.repository.save_account(
Account(id=None, name="默认账户", initial_cash=Decimal("0"), notes="")
)
if self.repository.get_default_strategy_template() is None:
self.repository.save_strategy_template(
StrategyTemplate(
id=None,
name="默认网格模板",
grid_spacing_pct=Decimal("0.03"),
amount_per_grid=Decimal("5000"),
base_target_amount=Decimal("20000"),
max_position_amount=Decimal("80000"),
min_lot=100,
is_default=True,
)
)
def get_active_account(self) -> Account:
self.ensure_defaults()
accounts = self.repository.list_accounts()
if not accounts:
raise RuntimeError("No account exists after default initialization")
return accounts[0]
def save_account(self, account: Account) -> Account:
if not account.name.strip():
raise ValueError("账户名称不能为空")
if account.initial_cash < 0:
raise ValueError("初始资金不能为负数")
return self.repository.save_account(account)
def add_instrument(self, instrument: Instrument) -> Instrument:
if not instrument.code.strip():
raise ValueError("标的代码不能为空")
if not instrument.name.strip():
raise ValueError("标的名称不能为空")
if instrument.lot_size <= 0:
raise ValueError("交易单位必须大于 0")
existing = self.repository.get_instrument_by_code(instrument.code, instrument.market)
if instrument.id is None and existing is not None:
raise ValueError(f"{instrument.code} 已存在")
return self.repository.save_instrument(instrument)
def update_instrument(self, instrument: Instrument) -> Instrument:
if instrument.id is None:
raise ValueError("更新标的需要 id")
return self.add_instrument(instrument)
def list_instruments(self) -> list[Instrument]:
return self.repository.list_instruments()
def get_default_strategy_template(self) -> StrategyTemplate:
self.ensure_defaults()
template = self.repository.get_default_strategy_template()
if template is None:
raise RuntimeError("No default strategy template exists")
return template
def save_strategy_template(self, template: StrategyTemplate) -> StrategyTemplate:
if template.grid_spacing_pct <= 0:
raise ValueError("网格间距必须大于 0")
if template.amount_per_grid <= 0:
raise ValueError("每格金额必须大于 0")
if template.min_lot <= 0:
raise ValueError("最小交易单位必须大于 0")
return self.repository.save_strategy_template(template)
def save_strategy_override(self, override: StrategyOverride) -> StrategyOverride:
if override.instrument_id <= 0:
raise ValueError("策略覆盖需要有效标的")
return self.repository.save_strategy_override(override)
def estimate_trade_fees(
self,
side: TradeSide,
trade_price: Decimal,
quantity: int,
) -> FeeEstimate:
rules = FeeRules(
commission_rate=DEFAULT_FEE_CONFIG.commission_rate,
min_commission=DEFAULT_FEE_CONFIG.min_commission,
stamp_tax_rate=DEFAULT_FEE_CONFIG.stamp_tax_rate,
transfer_fee_rate=DEFAULT_FEE_CONFIG.transfer_fee_rate,
)
return estimate_fees(side, trade_price, quantity, rules)
def save_trade(self, trade: Trade) -> Trade:
instrument = self._require_instrument(trade.instrument_id)
self._require_account(trade.account_id)
self._validate_trade_input(trade, instrument)
existing_trades = [
existing
for existing in self.repository.list_trades(account_id=trade.account_id)
if existing.id != trade.id
]
if trade.side is TradeSide.SELL:
self._validate_t_plus_one_available(trade, instrument, existing_trades)
self._validate_resulting_position(trade, existing_trades)
return self.repository.save_trade(trade)
def update_trade(self, trade: Trade) -> Trade:
if trade.id is None:
raise ValueError("更新成交需要 id")
return self.save_trade(trade)
def delete_trade(self, trade_id: int) -> None:
self.repository.delete_trade(trade_id)
def list_trades(self, *, instrument_id: int | None = None) -> list[Trade]:
account = self.get_active_account()
return self.repository.list_trades(account_id=account.id, instrument_id=instrument_id)
def get_position_summaries(self, *, as_of: date | None = None) -> list[PositionSummary]:
as_of_date = as_of or date.today()
account = self.get_active_account()
instruments = self.repository.list_instruments()
trades = self.repository.list_trades(account_id=account.id)
return calculate_positions(
instruments,
trades,
as_of=as_of_date,
instrument_cashflows=self._instrument_cashflows(account.id),
)
def get_account_summary(self, *, as_of: date | None = None) -> AccountSummary:
as_of_date = as_of or date.today()
account = self.get_active_account()
positions = self.get_position_summaries(as_of=as_of_date)
trades = [
trade
for trade in self.repository.list_trades(account_id=account.id)
if trade.trade_date <= as_of_date
]
ledger_entries = [
entry
for entry in self.repository.list_cash_ledger_entries(account_id=account.id)
if entry.entry_date <= as_of_date
]
return calculate_account_summary(account, positions, ledger_entries, trades)
def _require_account(self, account_id: int) -> Account:
account = self.repository.get_account(account_id)
if account is None:
raise ValueError(f"账户不存在: {account_id}")
return account
def _require_instrument(self, instrument_id: int) -> Instrument:
instrument = self.repository.get_instrument(instrument_id)
if instrument is None:
raise ValueError(f"标的不存在: {instrument_id}")
return instrument
def _validate_trade_input(self, trade: Trade, instrument: Instrument) -> None:
if trade.price <= 0:
raise ValueError("成交价格必须大于 0")
if trade.quantity <= 0:
raise ValueError("成交数量必须大于 0")
if not instrument.allow_odd_lot and trade.quantity % instrument.lot_size != 0:
raise ValueError(f"成交数量必须是 {instrument.lot_size} 的整数倍")
if trade.total_fee < 0:
raise ValueError("成交费用不能为负数")
def _validate_t_plus_one_available(
self,
trade: Trade,
instrument: Instrument,
existing_trades: list[Trade],
) -> None:
[position] = calculate_positions(
[instrument],
[item for item in existing_trades if item.instrument_id == instrument.id],
as_of=trade.trade_date,
)
if trade.quantity > position.available_quantity:
raise ValueError(
f"T+1 可卖数量不足:可卖 {position.available_quantity},尝试卖出 {trade.quantity}"
)
def _validate_resulting_position(self, trade: Trade, existing_trades: list[Trade]) -> None:
instruments = self.repository.list_instruments()
try:
calculate_positions(
instruments,
[*existing_trades, trade],
as_of=trade.trade_date,
)
except CalculationError as exc:
raise ValueError(str(exc)) from exc
def _instrument_cashflows(self, account_id: int) -> dict[int, Decimal]:
totals: dict[int, Decimal] = defaultdict(lambda: Decimal("0"))
for entry in self.repository.list_cash_ledger_entries(account_id=account_id):
if entry.instrument_id is not None:
totals[entry.instrument_id] += entry.amount
return dict(totals)

131
tests/test_services.py Normal file
View File

@@ -0,0 +1,131 @@
from datetime import date, timedelta
from decimal import Decimal
import pytest
from grid_trading.domain.models import Instrument, Trade, TradeGroup, TradeSide
from grid_trading.services.trading_service import TradingService
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].grid_profit == Decimal("91.47")
assert summary.cash == Decimal("98138.97")
assert summary.market_value == Decimal("2000.00")
assert summary.total_assets == Decimal("100138.97")
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