Compare commits
6 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 2ee6043592 | |||
| 102f602aee | |||
| 36db2ed5a6 | |||
| 6220fc5943 | |||
| 8e7e7282de | |||
| 861c7faa8c |
@@ -21,6 +21,8 @@ export interface PaperCreate {
|
||||
pool?: string
|
||||
max_pool?: number
|
||||
benchmark?: string
|
||||
// 撮合引擎(影子柜台 P1):eod_replay=日终回放 / shadow=影子柜台
|
||||
engine?: string
|
||||
}
|
||||
|
||||
export interface PaperAccount {
|
||||
@@ -30,10 +32,12 @@ export interface PaperAccount {
|
||||
interval: string
|
||||
status: string
|
||||
strategy_type?: string
|
||||
engine?: string
|
||||
symbols?: string
|
||||
initial_capital?: number
|
||||
start_date?: string
|
||||
end_date?: string
|
||||
created_at?: string
|
||||
last_run_date?: string | null
|
||||
next_run_at?: string | null
|
||||
checkpoint_date?: string | null
|
||||
|
||||
@@ -13,6 +13,8 @@ export interface PortfolioBacktestReq {
|
||||
stamp_duty_rate?: number
|
||||
min_commission?: number
|
||||
slippage?: number
|
||||
// K线周期(组合回放暂仅日线 d)
|
||||
interval?: string
|
||||
}
|
||||
|
||||
export interface EquityPoint {
|
||||
|
||||
@@ -34,6 +34,17 @@
|
||||
background: rgba(255, 176, 0, 0.1);
|
||||
border-color: rgba(255, 176, 0, 0.3);
|
||||
}
|
||||
.chip-shadow {
|
||||
color: var(--brand);
|
||||
background: rgba(0, 229, 255, 0.1);
|
||||
border-color: rgba(0, 229, 255, 0.3);
|
||||
margin-left: 4px;
|
||||
}
|
||||
.chip-cta {
|
||||
color: #6ee7a8;
|
||||
background: rgba(110, 231, 168, 0.1);
|
||||
border-color: rgba(110, 231, 168, 0.3);
|
||||
}
|
||||
.chip-portfolio {
|
||||
color: var(--amber);
|
||||
background: rgba(255, 176, 0, 0.08);
|
||||
|
||||
@@ -6,6 +6,7 @@ import type { EChartsCoreOption } from 'echarts'
|
||||
import { useChart } from '@/composables/useChart'
|
||||
import { darkTitle, darkTooltip, darkGrid, darkAxis } from '@/utils/echartsDark'
|
||||
import { getTaskParams } from '@/api/backtest'
|
||||
import { INTERVAL_OPTIONS } from '@/constants/intervals'
|
||||
import {
|
||||
postPortfolioBacktest,
|
||||
getPortfolioResult,
|
||||
@@ -60,6 +61,7 @@ const form = reactive({
|
||||
stamp_duty_rate: 0.001,
|
||||
min_commission: 5,
|
||||
slippage: 0.001,
|
||||
interval: 'd',
|
||||
})
|
||||
watch(
|
||||
() => [form.start, form.end],
|
||||
@@ -81,6 +83,7 @@ const strategyOptions = [
|
||||
{ label: '牛熊动量', value: 'momentum_timing' },
|
||||
{ label: '价值精选', value: 'value_selection' },
|
||||
{ label: '小市值轮动', value: 'small_cap' },
|
||||
{ label: '通路测试(影子vs实盘双轨)', value: 'channel_test' },
|
||||
]
|
||||
const BENCHMARK_OPTIONS = [
|
||||
{ label: '沪深300', value: '000300.XSHG' },
|
||||
@@ -360,6 +363,7 @@ async function onSubmit(): Promise<void> {
|
||||
stamp_duty_rate: Number(form.stamp_duty_rate),
|
||||
min_commission: Number(form.min_commission),
|
||||
slippage: Number(form.slippage),
|
||||
interval: form.interval,
|
||||
})
|
||||
ElMessage.success('回测已提交,后台运行中')
|
||||
router.push('/backtest/history') // 跳任务中心(历史任务页)看进度/结果
|
||||
@@ -398,6 +402,7 @@ async function prefillFromTask(tid: string): Promise<void> {
|
||||
if (p.stamp_duty_rate != null) form.stamp_duty_rate = Number(p.stamp_duty_rate)
|
||||
if (p.min_commission != null) form.min_commission = Number(p.min_commission)
|
||||
if (p.slippage != null) form.slippage = Number(p.slippage)
|
||||
if (p.interval) form.interval = String(p.interval)
|
||||
} catch {
|
||||
/* 预填失败,走默认 */
|
||||
}
|
||||
@@ -444,6 +449,15 @@ function fmtNum(v: number | null | undefined, digits = 2): string {
|
||||
<el-input-number v-model="form.max_pool" :min="0" :step="10" :controls="false" style="width: 220px" />
|
||||
<span class="muted form-hint">0=全市场不限, N=前N只(MVP验证用, 默认30)</span>
|
||||
</el-form-item>
|
||||
<el-form-item label="K 线周期">
|
||||
<el-select v-model="form.interval" style="width: 220px">
|
||||
<el-option v-for="o in INTERVAL_OPTIONS" :key="o.value" :value="o.value" :label="o.label" :disabled="o.value !== 'd'">
|
||||
<span>{{ o.label }}</span>
|
||||
<span v-if="o.value !== 'd'" class="muted" style="float: right; font-size: 12px">影子柜台实时可用·组合回放暂仅日线</span>
|
||||
</el-option>
|
||||
</el-select>
|
||||
<span class="muted form-hint">组合回放暂仅日线(本地数据);分钟档影子柜台/实盘可用</span>
|
||||
</el-form-item>
|
||||
</el-form>
|
||||
</el-card>
|
||||
|
||||
|
||||
@@ -33,6 +33,7 @@ const PORTFOLIO_LABELS: Record<string, string> = {
|
||||
momentum_timing: '牛熊动量',
|
||||
value_selection: '价值精选',
|
||||
small_cap: '小市值轮动',
|
||||
channel_test: '通路测试(影子vs实盘双轨)',
|
||||
}
|
||||
|
||||
onMounted(async () => {
|
||||
@@ -287,6 +288,15 @@ async function onSubmit(): Promise<void> {
|
||||
<el-option v-for="b in BENCH_OPTIONS" :key="b.value" :label="b.label" :value="b.value" />
|
||||
</el-select>
|
||||
</el-form-item>
|
||||
<el-form-item label="K 线周期">
|
||||
<el-select v-model="form.interval" style="width: 220px">
|
||||
<el-option v-for="o in INTERVAL_OPTIONS" :key="o.value" :value="o.value" :label="o.label">
|
||||
<span>{{ o.label }}</span>
|
||||
<span v-if="!o.replayOk" class="muted" style="float: right; font-size: 12px">影子柜台实时可用·回放暂无本地数据</span>
|
||||
</el-option>
|
||||
</el-select>
|
||||
<span class="muted form-hint">miniQMT 成品K线,默认日线</span>
|
||||
</el-form-item>
|
||||
</el-form>
|
||||
</el-card>
|
||||
|
||||
|
||||
@@ -44,7 +44,7 @@ const stats = computed(() => ({
|
||||
failed: accounts.value.filter((a) => a.status === 'failed').length,
|
||||
}))
|
||||
|
||||
const modeLabel: Record<string, string> = { replay: '回放', live: '实走' }
|
||||
const modeLabel: Record<string, string> = { replay: '回放', live: '实走', shadow: '影子' }
|
||||
const statusLabel: Record<string, string> = { running: '运行中', done: '完成', failed: '失败', created: '已创建', pending: '排队', stopped: '已停止' }
|
||||
|
||||
function pct(v: number | null | undefined): string {
|
||||
@@ -57,11 +57,26 @@ function num(v: number | null | undefined): string {
|
||||
}
|
||||
function symbolsShort(s: string | undefined): string {
|
||||
if (!s) return '—'
|
||||
const arr = typeof s === 'string' ? s.split(',') : []
|
||||
return arr.length > 3 ? arr.slice(0, 3).join(',') + ` +${arr.length - 3}` : s
|
||||
let arr: string[] = []
|
||||
try {
|
||||
arr = JSON.parse(s)
|
||||
} catch {
|
||||
arr = String(s).split(',').map((x) => x.trim())
|
||||
}
|
||||
if (!arr.length) return '—'
|
||||
return arr.length > 3 ? arr.slice(0, 3).join(',') + ` +${arr.length - 3}` : arr.join(',')
|
||||
}
|
||||
const typeLabel: Record<string, string> = { cta: '个股', portfolio: '组合' }
|
||||
function createdShort(v: string | undefined): string {
|
||||
if (!v) return '—'
|
||||
// created_at 存 UTC isoformat → 转北京时间显示
|
||||
const d = new Date(v)
|
||||
if (Number.isNaN(d.getTime())) return v.replace('T', ' ').slice(0, 16)
|
||||
return new Date(d.getTime() + 8 * 3600 * 1000)
|
||||
.toISOString().replace('T', ' ').slice(0, 16)
|
||||
}
|
||||
function open(a: PaperAccount): void {
|
||||
router.push(a.mode === 'live' ? `/paper/live/${a.id}` : `/paper/result/${a.id}`)
|
||||
router.push(a.mode === 'replay' ? `/paper/result/${a.id}` : `/paper/live/${a.id}`)
|
||||
}
|
||||
function goNew(): void {
|
||||
router.push('/paper/new')
|
||||
@@ -170,6 +185,7 @@ async function saveEdit(): Promise<void> {
|
||||
<el-select v-model="modeFilter" placeholder="全部模式" clearable style="width: 130px">
|
||||
<el-option label="回放" value="replay" />
|
||||
<el-option label="实走" value="live" />
|
||||
<el-option label="影子" value="shadow" />
|
||||
</el-select>
|
||||
<el-select v-model="statusFilter" placeholder="全部状态" clearable style="width: 130px">
|
||||
<el-option label="运行中" value="running" />
|
||||
@@ -181,48 +197,55 @@ async function saveEdit(): Promise<void> {
|
||||
|
||||
<el-card class="table-card" shadow="never">
|
||||
<el-table :data="filtered" size="small" empty-text="暂无模拟盘">
|
||||
<el-table-column label="名称 / ID" min-width="160">
|
||||
<el-table-column label="名称 / ID" min-width="150">
|
||||
<template #default="{ row }">
|
||||
<div class="cell-name">{{ row.name }}</div>
|
||||
<div class="cell-id mono">#{{ row.id }}</div>
|
||||
</template>
|
||||
</el-table-column>
|
||||
<el-table-column label="模式" width="80">
|
||||
<template #default="{ row }"><span class="chip chip-paper">{{ modeLabel[row.mode] || row.mode }}</span></template>
|
||||
<el-table-column label="类型" width="76">
|
||||
<template #default="{ row }">
|
||||
<span class="chip" :class="row.strategy_type === 'portfolio' ? 'chip-portfolio' : 'chip-cta'">
|
||||
{{ typeLabel[row.strategy_type] || row.strategy_type || '个股' }}
|
||||
</span>
|
||||
</template>
|
||||
</el-table-column>
|
||||
<el-table-column label="频率" width="70">
|
||||
<template #default="{ row }"><span class="mono">{{ row.interval }}</span></template>
|
||||
<el-table-column label="模式" width="92">
|
||||
<template #default="{ row }">
|
||||
<span class="chip chip-paper" :class="{ 'chip-shadow': row.mode === 'shadow' }">{{ modeLabel[row.mode] || row.mode }}</span>
|
||||
<span v-if="row.strategy_type === 'portfolio' && row.mode === 'live'" class="chip chip-shadow" title="日终回放=每晚 20:30 全量重放结算">日终</span>
|
||||
</template>
|
||||
</el-table-column>
|
||||
<el-table-column prop="instance" label="实例" min-width="160">
|
||||
<template #default="{ row }"><span class="mono">{{ row.instance || '—' }}</span></template>
|
||||
<el-table-column label="频率" width="64">
|
||||
<template #default="{ row }"><span class="mono muted">{{ row.interval }}</span></template>
|
||||
</el-table-column>
|
||||
<el-table-column label="标的" min-width="120">
|
||||
<el-table-column label="标的" min-width="130">
|
||||
<template #default="{ row }"><span class="mono">{{ symbolsShort(row.symbols) }}</span></template>
|
||||
</el-table-column>
|
||||
<el-table-column label="收益率" width="100" align="right">
|
||||
<el-table-column label="收益率" width="86" align="right">
|
||||
<template #default="{ row }">
|
||||
<span class="mono" :class="(row.total_return ?? 0) >= 0 ? 'up' : 'down'">{{ pct(row.total_return) }}</span>
|
||||
</template>
|
||||
</el-table-column>
|
||||
<el-table-column label="最新净值" width="120" align="right">
|
||||
<el-table-column label="最新净值" width="104" align="right">
|
||||
<template #default="{ row }"><span class="mono">{{ num(row.latest_equity) }}</span></template>
|
||||
</el-table-column>
|
||||
<el-table-column label="起止" min-width="180">
|
||||
<template #default="{ row }"><span class="mono muted">{{ row.start_date }} ~ {{ row.end_date }}</span></template>
|
||||
<el-table-column label="创建时间" width="142">
|
||||
<template #default="{ row }"><span class="mono muted">{{ createdShort(row.created_at) }}</span></template>
|
||||
</el-table-column>
|
||||
<el-table-column label="状态" width="84">
|
||||
<el-table-column label="状态" width="80">
|
||||
<template #default="{ row }"><span class="chip" :class="`st-${row.status}`">{{ statusLabel[row.status] || row.status }}</span></template>
|
||||
</el-table-column>
|
||||
<el-table-column label="操作" width="200" fixed="right">
|
||||
<el-table-column label="操作" width="176" fixed="right">
|
||||
<template #default="{ row }">
|
||||
<el-button link type="primary" @click="open(row)">查看</el-button>
|
||||
<el-button
|
||||
v-if="row.mode === 'live' && row.status === 'running'"
|
||||
v-if="row.mode !== 'replay' && row.status === 'running'"
|
||||
link type="warning" :loading="actionLoading === row.id"
|
||||
@click="onStop(row)"
|
||||
>停止</el-button>
|
||||
<el-button
|
||||
v-if="row.mode === 'live' && row.status === 'stopped'"
|
||||
v-if="row.mode !== 'replay' && row.status === 'stopped'"
|
||||
link type="success" :loading="actionLoading === row.id"
|
||||
@click="onResume(row)"
|
||||
>恢复</el-button>
|
||||
|
||||
@@ -32,6 +32,7 @@ const PORTFOLIO_LABELS: Record<string, string> = {
|
||||
momentum_timing: '牛熊动量',
|
||||
value_selection: '价值精选',
|
||||
small_cap: '小市值轮动',
|
||||
channel_test: '通路测试(影子vs实盘双轨)',
|
||||
}
|
||||
|
||||
onMounted(async () => {
|
||||
@@ -84,7 +85,8 @@ watch(
|
||||
|
||||
const modes = [
|
||||
{ value: 'replay', label: '回放', desc: '历史重放,立即出结果' },
|
||||
{ value: 'live', label: '实走', desc: '每日定时 step,长期跟踪' },
|
||||
{ value: 'live', label: '实走', desc: '每日 20:30 定时结算,长期跟踪' },
|
||||
{ value: 'shadow', label: '影子', desc: 'VPS 影子柜台盘中实时本地撮合' },
|
||||
]
|
||||
const sessions = [
|
||||
{ value: 'next_open', label: '次日开盘', desc: '收盘型信号,T+1 开盘撮合' },
|
||||
@@ -117,15 +119,15 @@ const isPortfolio = computed(() => strategyType.value === 'portfolio')
|
||||
const portfolioStrategy = ref('all_weather')
|
||||
|
||||
function setMode(m: string): void {
|
||||
// 组合类型锁定实走:回放卡片不可选(历史回放走「组合回测」页)
|
||||
// 组合类型不支持回放(历史回放走「组合回测」页)
|
||||
if (isPortfolio.value && m === 'replay') return
|
||||
form.value.mode = m
|
||||
}
|
||||
|
||||
// 组合类型锁定实走(历史回放走「组合回测」页)
|
||||
// 组合类型默认实走(历史回放走「组合回测」页)
|
||||
watch(strategyType, (t) => {
|
||||
if (t === 'portfolio') {
|
||||
form.value.mode = 'live'
|
||||
if (form.value.mode === 'replay') form.value.mode = 'live'
|
||||
if (!portfolioStrategy.value && portfolioOptions.value.length) {
|
||||
portfolioStrategy.value = portfolioOptions.value[0]
|
||||
}
|
||||
@@ -138,7 +140,6 @@ async function onSubmit(): Promise<void> {
|
||||
const payload: PaperCreate = { ...form.value }
|
||||
if (strategyType.value === 'portfolio') {
|
||||
payload.strategy_type = 'portfolio'
|
||||
payload.mode = 'live'
|
||||
payload.symbols = [poolForm.pool]
|
||||
payload.strategies = [{
|
||||
name: portfolioStrategy.value,
|
||||
@@ -150,9 +151,15 @@ async function onSubmit(): Promise<void> {
|
||||
payload.max_pool = Number(poolForm.max_pool)
|
||||
payload.benchmark = poolForm.benchmark
|
||||
}
|
||||
// 影子模式:engine=shadow(VPS 柜台盘中实时撮合);其余 eod_replay(回放/日终)
|
||||
payload.engine = payload.mode === 'shadow' ? 'shadow' : 'eod_replay'
|
||||
const aid = await createPaper(payload)
|
||||
ElMessage.success(`已创建模拟盘 #${aid}(今晚 20:30 起每日结算)`)
|
||||
router.push(payload.mode === 'live' ? `/paper/live/${aid}` : `/paper/result/${aid}`)
|
||||
ElMessage.success(payload.mode === 'shadow'
|
||||
? `已创建影子柜台模拟盘 #${aid}(VPS 柜台运行期间盘中实时结算)`
|
||||
: payload.mode === 'live'
|
||||
? `已创建模拟盘 #${aid}(今晚 20:30 起每日结算)`
|
||||
: `已创建回放模拟盘 #${aid}`)
|
||||
router.push(payload.mode === 'replay' ? `/paper/result/${aid}` : `/paper/live/${aid}`)
|
||||
} catch (e: unknown) {
|
||||
ElMessage.error(e instanceof Error ? e.message : '创建失败')
|
||||
} finally {
|
||||
@@ -206,15 +213,15 @@ function onSymbols(v: string): void {
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div v-if="!isPortfolio" class="seg-label" style="margin-top:16px">K 线周期</div>
|
||||
<el-select v-if="!isPortfolio" v-model="form.interval" style="width: 220px">
|
||||
<el-option v-for="i in INTERVAL_OPTIONS" :key="i.value"
|
||||
<div class="seg-label" style="margin-top:16px">K 线周期</div>
|
||||
<el-select v-model="form.interval" style="width: 220px">
|
||||
<el-option v-for="i in INTERVAL_OPTIONS" :key="i.value"
|
||||
:value="i.value" :label="i.label">
|
||||
<span>{{ i.label }}</span>
|
||||
<span v-if="!i.replayOk" class="muted" style="float:right;font-size:12px">影子柜台实时可用·回放暂无本地数据</span>
|
||||
</el-option>
|
||||
</el-select>
|
||||
<div v-if="!isPortfolio" class="muted" style="font-size:12px;margin-top:4px">
|
||||
<div class="muted" style="font-size:12px;margin-top:4px">
|
||||
{{ INTERVAL_OPTIONS.find(o => o.value === form.interval)?.matchHint || (form.interval === 'd' ? '' : '分钟级周期需标的分钟数据') }}
|
||||
</div>
|
||||
</el-card>
|
||||
|
||||
@@ -43,7 +43,7 @@ class LiveCreateRequest(BaseModel):
|
||||
strategy_class: str = "AShareDoubleMaStrategy"
|
||||
strategy_name: str
|
||||
setting: dict = {}
|
||||
interval: str = "15m"
|
||||
interval: str = "" # 空=按类型给默认(cta→15m, portfolio→d);前端下拉显式传
|
||||
initial_capital: float = 1_000_000
|
||||
connect_wait_sec: int = 10
|
||||
init_wait_sec: int = 60
|
||||
@@ -83,9 +83,11 @@ def create_live(req: LiveCreateRequest):
|
||||
payload.setdefault("max_pool", 30)
|
||||
payload.setdefault("benchmark", "000300.XSHG")
|
||||
payload["vt_symbol"] = payload["pool"]
|
||||
payload["interval"] = "d"
|
||||
# 周期由前端下拉传(miniQMT 成品K线档位);空=默认日线
|
||||
payload["interval"] = payload.get("interval") or "d"
|
||||
else:
|
||||
# CTA 实盘:标的允许只写 6 位码,后端自动补交易所后缀
|
||||
payload["interval"] = payload.get("interval") or "15m"
|
||||
payload["vt_symbol"] = _normalize_vt_symbol(payload.get("vt_symbol", ""))
|
||||
# mini_path 兜底:req → env SANGUO_QMT_PATH → 内置默认(空值会导致 connect=-1)
|
||||
if not payload.get("mini_path"):
|
||||
|
||||
@@ -56,6 +56,8 @@ class PaperCreateRequest(BaseModel):
|
||||
# 组合策略实走(E1):strategy_type=portfolio 时 mode 必须 live,
|
||||
# strategies[0].name=组合策略名,pool/max_pool/benchmark 进 params
|
||||
strategy_type: str = "cta"
|
||||
# 撮合引擎(影子柜台 P1):eod_replay=日终回放(NAS 20:30) / shadow=影子柜台(VPS 盘中实时)
|
||||
engine: str = "eod_replay"
|
||||
pool: str = "hs300_subset"
|
||||
max_pool: int = 30
|
||||
benchmark: str = "000300.XSHG"
|
||||
@@ -69,9 +71,10 @@ def create_paper(req: PaperCreateRequest):
|
||||
db = _db_path["path"] or ":memory:"
|
||||
init_db(db)
|
||||
if req.strategy_type == "portfolio":
|
||||
if req.mode != "live":
|
||||
raise HTTPException(400, "组合策略模拟盘仅支持实走(live)模式;历史回放请用「组合回测」")
|
||||
if req.mode not in ("live", "shadow"):
|
||||
raise HTTPException(400, "组合策略模拟盘仅支持实走(live)/影子(shadow)模式;历史回放请用「组合回测」")
|
||||
payload = req.model_dump()
|
||||
payload["engine"] = "shadow" if req.mode == "shadow" else "eod_replay"
|
||||
payload["symbols"] = [req.pool]
|
||||
payload["strategies"] = [{
|
||||
"name": (req.strategies[0].name if req.strategies else "all_weather"),
|
||||
@@ -82,7 +85,9 @@ def create_paper(req: PaperCreateRequest):
|
||||
update_account_status(db, aid, "running")
|
||||
return {"account_id": aid, "status": "running"}
|
||||
|
||||
aid = save_account(db, req.model_dump())
|
||||
cta_payload = req.model_dump()
|
||||
cta_payload["engine"] = "shadow" if req.mode == "shadow" else "eod_replay"
|
||||
aid = save_account(db, cta_payload)
|
||||
status = "created"
|
||||
if req.mode == "replay": # 回放后台线程跑,create 立即返回(避免阻塞 worker 502)
|
||||
def _bg():
|
||||
@@ -94,7 +99,7 @@ def create_paper(req: PaperCreateRequest):
|
||||
update_account_status(db, aid, "failed", str(e))
|
||||
threading.Thread(target=_bg, daemon=True).start()
|
||||
status = "running"
|
||||
elif req.mode == "live": # 实走:全局 job 每日 20:30 遍历 step(不跑回放)
|
||||
elif req.mode in ("live", "shadow"): # 实走:每日 20:30 step;影子:等 VPS 影子柜台进程接管
|
||||
from sanguo_trader.persistence import update_account_status
|
||||
update_account_status(db, aid, "running")
|
||||
status = "running"
|
||||
|
||||
@@ -43,11 +43,14 @@ class PortfolioBacktestRequest(BaseModel):
|
||||
stamp_duty_rate: float = Field(default=0.001, description="印花税率卖出(千1=0.001)")
|
||||
min_commission: float = Field(default=5.0, description="单笔最低佣金(元)")
|
||||
slippage: float = Field(default=0.0, description="滑点比率(万10=0.001,0=不加)")
|
||||
interval: str = Field(default="d", description="K线周期:组合回放暂仅日线(d)")
|
||||
|
||||
|
||||
@router.post("/portfolio/backtest", dependencies=[Depends(verify_token)])
|
||||
async def run_portfolio_backtest(req: PortfolioBacktestRequest):
|
||||
"""异步提交组合回测任务,返回 task_id。前端轮询 GET /task/{id} 再取结果。"""
|
||||
if req.interval != "d":
|
||||
raise HTTPException(400, "组合回放暂仅支持日线(影子柜台将支持全周期分钟档)")
|
||||
tid = await get_orchestrator().submit_portfolio(
|
||||
start=req.start_date,
|
||||
end=req.end_date,
|
||||
@@ -60,6 +63,7 @@ async def run_portfolio_backtest(req: PortfolioBacktestRequest):
|
||||
stamp_duty_rate=req.stamp_duty_rate,
|
||||
min_commission=req.min_commission,
|
||||
slippage=req.slippage,
|
||||
interval=req.interval,
|
||||
)
|
||||
return {"task_id": tid}
|
||||
|
||||
|
||||
@@ -119,6 +119,7 @@ def run_portfolio_task(spec: dict) -> Any:
|
||||
"stamp_duty_rate": spec.get("stamp_duty_rate", 0.001),
|
||||
"min_commission": spec.get("min_commission", 5.0),
|
||||
"slippage": spec.get("slippage", 0.0),
|
||||
"interval": spec.get("interval", "d"),
|
||||
},
|
||||
start=period.get("start", ""),
|
||||
end=period.get("end", ""),
|
||||
|
||||
@@ -141,7 +141,8 @@ class Orchestrator:
|
||||
commission_rate: float = 0.0003,
|
||||
stamp_duty_rate: float = 0.001,
|
||||
min_commission: float = 5.0,
|
||||
slippage: float = 0.0) -> str:
|
||||
slippage: float = 0.0,
|
||||
interval: str = "d") -> str:
|
||||
"""Submit a portfolio backtest task asynchronously.
|
||||
|
||||
Runs runner_backtest as a subprocess (3600s hard cap) inside the
|
||||
@@ -165,6 +166,7 @@ class Orchestrator:
|
||||
stamp_duty_rate=stamp_duty_rate,
|
||||
min_commission=min_commission,
|
||||
slippage=slippage,
|
||||
interval=interval,
|
||||
db_path=self.db_path,
|
||||
file_dir=file_dir,
|
||||
)
|
||||
|
||||
@@ -23,6 +23,7 @@ def _build_live_strategy(provider):
|
||||
"""env 配置 → StrategyTemplate 实例(对齐 runner_backtest._build_strategy)。"""
|
||||
from sanguo_portfolio.strategies import (
|
||||
AllWeatherConfig, AllWeatherStrategy,
|
||||
ChannelTestConfig, ChannelTestStrategy,
|
||||
MomentumTimingConfig, MomentumTimingStrategy,
|
||||
SmallCapConfig, SmallCapStrategy,
|
||||
ValueSelectionConfig, ValueSelectionStrategy,
|
||||
@@ -39,6 +40,8 @@ def _build_live_strategy(provider):
|
||||
provider=provider, config=ValueSelectionConfig(max_pool=max_pool)),
|
||||
"small_cap": lambda: SmallCapStrategy(
|
||||
provider=provider, config=SmallCapConfig(max_pool=max_pool)),
|
||||
"channel_test": lambda: ChannelTestStrategy(
|
||||
provider=provider, config=ChannelTestConfig()),
|
||||
}
|
||||
if name not in factories:
|
||||
raise ValueError(
|
||||
|
||||
@@ -56,8 +56,8 @@ def parse_args() -> argparse.Namespace:
|
||||
p.add_argument("--slippage", type=float, default=0.0, help="滑点比率(万10=0.001,0=不加)")
|
||||
p.add_argument(
|
||||
"--strategy", default="all_weather",
|
||||
choices=["all_weather", "momentum_timing", "value_selection", "small_cap"],
|
||||
help="策略: all_weather(全天候轮动) / momentum_timing(牛熊分界+取强舍弱+均线动量) / value_selection(价值精选6条月度调仓) / small_cap(小市值20只轮动,无对冲)",
|
||||
choices=["all_weather", "momentum_timing", "value_selection", "small_cap", "channel_test"],
|
||||
help="策略: all_weather / momentum_timing / value_selection / small_cap / channel_test(通路测试,影子vs实盘双轨验证)",
|
||||
)
|
||||
p.add_argument(
|
||||
"--provider", default="local", choices=["local", "baostock", "miniqmt", "unified"],
|
||||
@@ -176,6 +176,9 @@ def _build_strategy(args: argparse.Namespace, provider: Any) -> Any:
|
||||
provider=provider,
|
||||
config=SmallCapConfig(max_pool=args.max_pool),
|
||||
)
|
||||
if name == "channel_test":
|
||||
from .strategies import ChannelTestConfig, ChannelTestStrategy
|
||||
return ChannelTestStrategy(provider=provider, config=ChannelTestConfig())
|
||||
raise ValueError(
|
||||
f"未知 strategy: {name}(支持: all_weather / momentum_timing / value_selection / small_cap)"
|
||||
)
|
||||
@@ -375,6 +378,10 @@ def _write_result_md(result: Dict[str, Any], path: str, args: argparse.Namespace
|
||||
logger.warning("写结果文件失败: %s", exc)
|
||||
|
||||
|
||||
def _raise_interval(interval: Any) -> str:
|
||||
raise ValueError(f"组合回测暂仅支持日线(interval=d),收到: {interval}")
|
||||
|
||||
|
||||
def run_backtest_json(params: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""JSON 入口(供 SSH 触发,前端 MVP 用)。
|
||||
|
||||
@@ -402,7 +409,9 @@ def run_backtest_json(params: Dict[str, Any]) -> Dict[str, Any]:
|
||||
end=params.get("end_date", "2024-02-29"),
|
||||
cash=float(params.get("initial_cash", 1_000_000.0)),
|
||||
benchmark=params.get("benchmark", "000300.XSHG"),
|
||||
frequency="day",
|
||||
# 周期:目前组合回放只有日线有本地数据;分钟级等数据层补齐后放开
|
||||
frequency="day" if params.get("interval", "d") in ("", "d", "day") else (
|
||||
_raise_interval(params.get("interval"))),
|
||||
strategy=strategy_name,
|
||||
provider=params.get("provider", "local"),
|
||||
provider_config=params.get("provider_config", "{}"),
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
"""sanguo_portfolio 策略层。"""
|
||||
from .all_weather import AllWeatherConfig, AllWeatherStrategy, BrokerFacade
|
||||
from .channel_test import ChannelTestConfig, ChannelTestStrategy
|
||||
from .momentum_timing import MomentumTimingConfig, MomentumTimingStrategy
|
||||
from .small_cap import SmallCapConfig, SmallCapStrategy
|
||||
from .value_selection import ValueSelectionConfig, ValueSelectionStrategy
|
||||
@@ -8,6 +9,8 @@ __all__ = [
|
||||
"AllWeatherStrategy",
|
||||
"AllWeatherConfig",
|
||||
"BrokerFacade",
|
||||
"ChannelTestStrategy",
|
||||
"ChannelTestConfig",
|
||||
"MomentumTimingStrategy",
|
||||
"MomentumTimingConfig",
|
||||
"SmallCapStrategy",
|
||||
|
||||
@@ -0,0 +1,116 @@
|
||||
"""通路测试策略(影子柜台 vs 实盘 双轨验证专用,docs/design/paper-shadow-desk-design.md §8)。
|
||||
|
||||
不是为赚钱,是为**把买卖全通路在真实 miniQMT 数据上跑通**,让 ShadowBroker 与
|
||||
QmtBroker 在同一段时间、同一批订单上各自成交,事后对账验证通路正确性。
|
||||
|
||||
设计目标(每个调仓日都尽量触发):
|
||||
- 买:用 order_target_value 等权买入若干只 → 走买入 + 整手(100) + 资金扣减 + 均价
|
||||
- 卖:次日先清掉上轮非目标持仓 → 走卖出 + 印花税 + T+1 可卖
|
||||
- 轮换:目标集每个周期偏移一格 → 既有卖(旧)又有买(新),长期跑必两端都覆盖
|
||||
- T+1 拒单探针:买入后立即试卖同一只(当日) → 两端都应被 T+1 拒,验证拒单通路一致
|
||||
- 上涨/下跌/停牌/涨跌停:靠自然行情出现,差异进双轨对账报告(P1.3 涨跌停拦截随影子柜台完善)
|
||||
|
||||
universe 默认几只高流动性 ETF/蓝筹(miniQMT 必有数据、好成交),可经 env/max_pool 调。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, List, Optional
|
||||
|
||||
from .all_weather import BrokerFacade, _available_cash, _get_positions
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# 默认 universe:流动性好的 ETF + 蓝筹,miniQMT 必有数据,实盘也容易成交
|
||||
_DEFAULT_UNIVERSE: List[str] = [
|
||||
"510300.XSHG", # 沪深300ETF
|
||||
"510050.XSHG", # 上证50ETF
|
||||
"159915.XSHE", # 创业板ETF
|
||||
"510500.XSHG", # 中证500ETF
|
||||
"588000.XSHG", # 科创50ETF
|
||||
]
|
||||
|
||||
|
||||
@dataclass
|
||||
class ChannelTestConfig:
|
||||
universe: List[str] = field(default_factory=lambda: list(_DEFAULT_UNIVERSE))
|
||||
hold_n: int = 2 # 每轮等权持有几只
|
||||
period: int = 1 # 每 N 个交易日轮换一次(1=每日)
|
||||
probe_t1: bool = True # 是否做 T+1 当日卖探针(验证拒单通路)
|
||||
benchmark: str = "000300.XSHG"
|
||||
|
||||
|
||||
class ChannelTestStrategy:
|
||||
"""通路测试策略:周期性等权轮换 + T+1 拒单探针。"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
provider: Any,
|
||||
broker: Optional[BrokerFacade] = None,
|
||||
config: Optional[ChannelTestConfig] = None,
|
||||
) -> None:
|
||||
self.provider = provider
|
||||
self.broker = broker or BrokerFacade()
|
||||
self.config = config or ChannelTestConfig()
|
||||
self._day = 0
|
||||
|
||||
def initialize(self, context: Any) -> None:
|
||||
b = self.broker
|
||||
b.set_benchmark(self.config.benchmark)
|
||||
b.set_option("use_real_price", True)
|
||||
b.set_option("avoid_future_data", True)
|
||||
b.run_daily(self.rotate, "9:30")
|
||||
|
||||
# ---------------- 轮换主流程 ----------------
|
||||
def _target_set(self) -> List[str]:
|
||||
"""按 day 偏移在 universe 里取 hold_n 只(循环),保证每轮目标变。"""
|
||||
u = self.config.universe or _DEFAULT_UNIVERSE
|
||||
n = max(1, self.config.hold_n)
|
||||
if len(u) <= n:
|
||||
return list(u)
|
||||
offset = (self._day // max(1, self.config.period)) % len(u)
|
||||
# 取从 offset 起的 n 只(环绕)
|
||||
return [u[(offset + i) % len(u)] for i in range(n)]
|
||||
|
||||
def rotate(self, context: Any) -> None:
|
||||
self._day += 1
|
||||
if (self._day - 1) % max(1, self.config.period) != 0:
|
||||
return
|
||||
positions = _get_positions(context)
|
||||
target = self._target_set()
|
||||
target_set = set(target)
|
||||
logger.info("[channel_test] day=%d target=%s holding=%s",
|
||||
self._day, target, list(positions.keys()))
|
||||
|
||||
# 1) 卖:清掉不在目标里的持仓(走卖出通路)
|
||||
for code in list(positions.keys()):
|
||||
if code not in target_set:
|
||||
logger.info("[channel_test] 卖出 %s", code)
|
||||
self.broker.order_target_value(code, 0)
|
||||
|
||||
# 2) 买/调:目标等权(走买入通路 + 整手)
|
||||
cash = _available_cash(context)
|
||||
total = cash + sum(_safe_value(p) for p in positions.values())
|
||||
per = total / max(1, len(target))
|
||||
for code in target:
|
||||
logger.info("[channel_test] 调仓 %s → target_value=%.2f", code, per)
|
||||
self.broker.order_target_value(code, per)
|
||||
|
||||
# 3) T+1 拒单探针:当日买入立即试卖 → 两端都应 T+1 拒(验证拒单通路)
|
||||
if self.config.probe_t1 and target:
|
||||
probe_code = target[0]
|
||||
try:
|
||||
self.broker.order_target_value(probe_code, 0)
|
||||
logger.info("[channel_test] T+1 探针 %s 当日卖已提交(预期被拒)", probe_code)
|
||||
except Exception as exc: # noqa: BLE001 - 探针失败不阻断主流程
|
||||
logger.debug("[channel_test] T+1 探针异常(正常): %s", exc)
|
||||
|
||||
|
||||
def _safe_value(pos: Any) -> float:
|
||||
"""从 position 对象取市值,兼容多种属性名。"""
|
||||
for attr in ("value", "market_value", "total_value"):
|
||||
v = getattr(pos, attr, None)
|
||||
if isinstance(v, (int, float)):
|
||||
return float(v)
|
||||
return 0.0
|
||||
@@ -74,6 +74,11 @@ def init_db(db_path: str) -> None:
|
||||
conn.execute("ALTER TABLE paper_accounts ADD COLUMN strategy_type TEXT DEFAULT 'cta'")
|
||||
except sqlite3.OperationalError:
|
||||
pass # 列已存在
|
||||
# 迁移:老库补 engine 列(影子柜台 P1:eod_replay=日终回放 / shadow=影子柜台)
|
||||
try:
|
||||
conn.execute("ALTER TABLE paper_accounts ADD COLUMN engine TEXT DEFAULT 'eod_replay'")
|
||||
except sqlite3.OperationalError:
|
||||
pass # 列已存在
|
||||
conn.execute("PRAGMA journal_mode=WAL")
|
||||
conn.commit()
|
||||
|
||||
@@ -85,8 +90,8 @@ def save_account(db_path: str, account: dict[str, Any]) -> int:
|
||||
(task_id, owner_id, name, strategy_type, mode, interval, symbols, strategies,
|
||||
initial_capital, rate, slippage, size, pricetick,
|
||||
stamp_duty_rate, transfer_fee_rate, min_commission,
|
||||
status, start_date, end_date, created_at, updated_at)
|
||||
VALUES (?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?)""",
|
||||
status, start_date, end_date, engine, created_at, updated_at)
|
||||
VALUES (?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?)""",
|
||||
(
|
||||
account.get("task_id"), account.get("owner_id", "admin"),
|
||||
account.get("name"), account.get("strategy_type", "cta"),
|
||||
@@ -102,6 +107,7 @@ def save_account(db_path: str, account: dict[str, Any]) -> int:
|
||||
account.get("status", "pending"),
|
||||
account.get("start_date") or account.get("start"),
|
||||
account.get("end_date") or account.get("end"),
|
||||
account.get("engine", "eod_replay"),
|
||||
_now(), _now(),
|
||||
),
|
||||
)
|
||||
|
||||
@@ -0,0 +1,58 @@
|
||||
"""影子柜台 CLI 入口(单实例文件锁防双开重复撮合)。
|
||||
|
||||
用法:
|
||||
python -m sanguo_trader.shadow # 读 env(见 runner.py docstring)
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
LOCK_FILE = Path(
|
||||
os.environ.get("SANGUO_SHADOW_LOCK")
|
||||
or Path.home() / ".sanguo_shadow_desk.lock"
|
||||
)
|
||||
|
||||
|
||||
def _acquire_lock() -> "object | None":
|
||||
"""单实例锁(Windows/msvcrt 与 POSIX/fcntl 双兼容)。"""
|
||||
LOCK_FILE.parent.mkdir(parents=True, exist_ok=True)
|
||||
try:
|
||||
import fcntl # POSIX
|
||||
|
||||
fh = open(LOCK_FILE, "w")
|
||||
fcntl.flock(fh, fcntl.LOCK_EX | fcntl.LOCK_NB)
|
||||
return fh
|
||||
except ImportError:
|
||||
pass
|
||||
except OSError:
|
||||
return None # 已有实例在跑
|
||||
try:
|
||||
import msvcrt # Windows
|
||||
|
||||
fh = open(LOCK_FILE, "w")
|
||||
msvcrt.locking(fh.fileno(), msvcrt.LK_NBLCK, 1)
|
||||
return fh
|
||||
except (ImportError, OSError):
|
||||
return None
|
||||
|
||||
|
||||
def main() -> int:
|
||||
logging.basicConfig(
|
||||
level=logging.INFO,
|
||||
format="%(asctime)s %(levelname)s %(name)s: %(message)s",
|
||||
)
|
||||
lock = _acquire_lock()
|
||||
if lock is None:
|
||||
print("影子柜台已在运行(锁占用),本次启动退出。", flush=True)
|
||||
return 0
|
||||
from .runner import run_shadow
|
||||
|
||||
run_shadow()
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
@@ -0,0 +1,228 @@
|
||||
"""影子柜台本地模拟 broker(P1,docs/design/paper-shadow-desk-design.md §3.2)。
|
||||
|
||||
挂在 bullet_trade LiveEngine 的 broker_factory 上:策略下单不出门,
|
||||
由本 broker 以「下单时刻实时价 ± 滑点」本地撮合,A 股费用/整手/T+1 对齐。
|
||||
|
||||
与实盘(QmtBroker)同接口(BrokerBase) → 同一个 LiveEngine 两种柜台,
|
||||
这是「双轨一致性验证」(§8)的基础:同策略同参数分别接真/假 broker 并跑对账。
|
||||
|
||||
价格来源由 price_getter 注入(通常=数据 provider 最新收盘/实时价),
|
||||
成交回调 on_trade 注入(落 paper_trades 表)。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
from typing import Any, Callable, Dict, List, Optional
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
LOT = 100 # A 股整手
|
||||
|
||||
|
||||
class ShadowBroker: # noqa: R0903 - 仅实现 BrokerBase 协议(bullet_trade duck-typed)
|
||||
"""本地虚拟账户撮合台。不继承 BrokerBase(避免硬依赖 bullet_trade import 顺序),
|
||||
LiveEngine 按 duck-typed 协议调用。"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
initial_cash: float = 1_000_000.0,
|
||||
*,
|
||||
commission_rate: float = 0.0003,
|
||||
stamp_duty_rate: float = 0.001,
|
||||
min_commission: float = 5.0,
|
||||
slippage: float = 0.0,
|
||||
price_getter: Optional[Callable[[str], Optional[float]]] = None,
|
||||
on_trade: Optional[Callable[[Dict[str, Any]], None]] = None,
|
||||
now_provider: Optional[Callable[[], datetime]] = None,
|
||||
) -> None:
|
||||
self.initial_cash = float(initial_cash)
|
||||
self.cash = float(initial_cash)
|
||||
self.commission_rate = float(commission_rate)
|
||||
self.stamp_duty_rate = float(stamp_duty_rate)
|
||||
self.min_commission = float(min_commission)
|
||||
self.slippage = float(slippage)
|
||||
self.price_getter = price_getter
|
||||
self.on_trade = on_trade
|
||||
self._now = now_provider or datetime.now
|
||||
self._connected = True # 本地柜台永远"在线"
|
||||
# security -> {"amount": int, "avg_cost": float}
|
||||
self.positions: Dict[str, Dict[str, Any]] = {}
|
||||
# T+1:今日买入数量(security -> int),before_open 清零
|
||||
self._today_bought: Dict[str, int] = {}
|
||||
self._today: str = ""
|
||||
self.orders: Dict[str, Dict[str, Any]] = {}
|
||||
self.trades: List[Dict[str, Any]] = []
|
||||
|
||||
# ===== 生命周期 =====
|
||||
def connect(self) -> bool:
|
||||
return True
|
||||
|
||||
def disconnect(self) -> bool:
|
||||
return True
|
||||
|
||||
def is_connected(self) -> bool:
|
||||
return True
|
||||
|
||||
def heartbeat(self) -> None:
|
||||
return None
|
||||
|
||||
def before_open(self) -> None:
|
||||
"""每个交易日开盘前:清 T+1 买入记录(昨日买的今天可卖)。"""
|
||||
self._today_bought = {}
|
||||
self._today = self._now().strftime("%Y-%m-%d")
|
||||
|
||||
def after_close(self) -> None:
|
||||
return None
|
||||
|
||||
# ===== 行情 =====
|
||||
def _ref_price(self, security: str, price: Optional[float]) -> Optional[float]:
|
||||
ref = price if price and price > 0 else None
|
||||
if ref is None and self.price_getter is not None:
|
||||
try:
|
||||
ref = self.price_getter(security)
|
||||
except Exception as exc: # noqa: BLE001 - 行情失败拒单而非崩柜台
|
||||
logger.warning("[shadow] 取价失败 %s: %s", security, exc)
|
||||
ref = None
|
||||
return ref
|
||||
|
||||
# ===== 下单(即时全额成交) =====
|
||||
async def buy(self, security: str, amount: int, price: Optional[float] = None,
|
||||
wait_timeout: Optional[float] = None, remark: Optional[str] = None,
|
||||
*, market: bool = False) -> str:
|
||||
order_id = self._new_order("buy", security, amount, price)
|
||||
ref = self._ref_price(security, price)
|
||||
if ref is None or ref <= 0:
|
||||
return self._reject(order_id, "无参考价")
|
||||
amount = int(amount)
|
||||
if amount <= 0:
|
||||
return self._reject(order_id, "数量非法")
|
||||
amount = amount - amount % LOT # 整手向下取
|
||||
if amount <= 0:
|
||||
return self._reject(order_id, "不足一手(100股)")
|
||||
fill = round(ref * (1 + self.slippage) + 1e-9, 2) # 买入价上浮滑点
|
||||
gross = amount * fill
|
||||
commission = max(gross * self.commission_rate, self.min_commission)
|
||||
if self.cash < gross + commission:
|
||||
return self._reject(order_id, f"资金不足 需{gross + commission:.2f} 有{self.cash:.2f}")
|
||||
self.cash -= gross + commission
|
||||
pos = self.positions.setdefault(security, {"amount": 0, "avg_cost": 0.0})
|
||||
old_amt, old_cost = pos["amount"], pos["avg_cost"]
|
||||
pos["amount"] = old_amt + amount
|
||||
pos["avg_cost"] = (old_amt * old_cost + gross) / pos["amount"]
|
||||
self._today_bought[security] = self._today_bought.get(security, 0) + amount
|
||||
self._fill(order_id, security, "buy", amount, fill, commission, 0.0)
|
||||
return order_id
|
||||
|
||||
async def sell(self, security: str, amount: int, price: Optional[float] = None,
|
||||
wait_timeout: Optional[float] = None, remark: Optional[str] = None,
|
||||
*, market: bool = False) -> str:
|
||||
order_id = self._new_order("sell", security, amount, price)
|
||||
ref = self._ref_price(security, price)
|
||||
if ref is None or ref <= 0:
|
||||
return self._reject(order_id, "无参考价")
|
||||
amount = int(amount)
|
||||
pos = self.positions.get(security)
|
||||
held = int(pos["amount"]) if pos else 0
|
||||
if amount <= 0 or held <= 0:
|
||||
return self._reject(order_id, "无持仓")
|
||||
# T+1:今日买入部分不可卖
|
||||
locked = self._today_bought.get(security, 0)
|
||||
sellable = max(held - locked, 0)
|
||||
if amount > sellable:
|
||||
amount = sellable
|
||||
amount = amount - amount % LOT
|
||||
if amount <= 0:
|
||||
return self._reject(order_id, f"可卖不足(T+1锁定{locked}股)")
|
||||
fill = round(ref * (1 - self.slippage) - 1e-9, 2) # 卖出价下压滑点
|
||||
gross = amount * fill
|
||||
commission = max(gross * self.commission_rate, self.min_commission)
|
||||
stamp_duty = gross * self.stamp_duty_rate
|
||||
self.cash += gross - commission - stamp_duty
|
||||
pos["amount"] = held - amount
|
||||
if pos["amount"] <= 0:
|
||||
self.positions.pop(security, None)
|
||||
self._fill(order_id, security, "sell", amount, fill, commission, stamp_duty)
|
||||
return order_id
|
||||
|
||||
async def cancel_order(self, order_id: str) -> bool:
|
||||
# 即时全额成交,无可撤单
|
||||
return False
|
||||
|
||||
async def get_order_status(self, order_id: str) -> Dict[str, Any]:
|
||||
st = self.orders.get(order_id) or {"order_id": order_id, "status": "not_found"}
|
||||
return dict(st)
|
||||
|
||||
def get_orders(self, order_id=None, security=None, status=None,
|
||||
from_broker: bool = False) -> List[Dict[str, Any]]:
|
||||
rows = [dict(o) for o in self.orders.values()
|
||||
if (order_id is None or o["order_id"] == order_id)
|
||||
and (security is None or o["security"] == security)]
|
||||
return rows
|
||||
|
||||
def get_open_orders(self) -> List[Dict[str, Any]]:
|
||||
return [dict(o) for o in self.orders.values() if o["status"] == "open"]
|
||||
|
||||
def get_trades(self, order_id=None, security=None) -> List[Dict[str, Any]]:
|
||||
return [dict(t) for t in self.trades
|
||||
if (order_id is None or t["order_id"] == order_id)
|
||||
and (security is None or t["security"] == security)]
|
||||
|
||||
# ===== 账户 =====
|
||||
def get_positions(self) -> List[Dict[str, Any]]:
|
||||
out = []
|
||||
for sym, pos in self.positions.items():
|
||||
px = self._ref_price(sym, None) or pos["avg_cost"]
|
||||
out.append({"security": sym, "amount": pos["amount"],
|
||||
"avg_cost": round(pos["avg_cost"], 6),
|
||||
"market_value": pos["amount"] * px,
|
||||
"price": px})
|
||||
return out
|
||||
|
||||
def get_account_info(self) -> Dict[str, Any]:
|
||||
positions = self.get_positions()
|
||||
mv = sum(p["market_value"] for p in positions)
|
||||
return {
|
||||
"total_value": self.cash + mv,
|
||||
"available_cash": self.cash,
|
||||
"positions": positions,
|
||||
"market_value": mv,
|
||||
}
|
||||
|
||||
# ===== 内部 =====
|
||||
def _new_order(self, side: str, security: str, amount: int,
|
||||
price: Optional[float]) -> str:
|
||||
order_id = f"shadow_{uuid.uuid4().hex[:12]}"
|
||||
self.orders[order_id] = {
|
||||
"order_id": order_id, "status": "open", "side": side,
|
||||
"security": security, "amount": int(amount),
|
||||
"price": price, "created_at": self._now().isoformat(timespec="seconds"),
|
||||
}
|
||||
return order_id
|
||||
|
||||
def _reject(self, order_id: str, reason: str) -> str:
|
||||
o = self.orders[order_id]
|
||||
o["status"] = "rejected"
|
||||
o["reject_reason"] = reason
|
||||
logger.info("[shadow] 拒单 %s %s %s: %s", o["side"], o["security"], o["amount"], reason)
|
||||
return order_id
|
||||
|
||||
def _fill(self, order_id: str, security: str, side: str, amount: int,
|
||||
fill: float, commission: float, stamp_duty: float) -> None:
|
||||
o = self.orders[order_id]
|
||||
o.update(status="filled", filled_amount=amount, filled_price=fill)
|
||||
trade = {
|
||||
"order_id": order_id, "security": security, "side": side,
|
||||
"amount": amount, "price": fill, "commission": round(commission, 2),
|
||||
"stamp_duty": round(stamp_duty, 2),
|
||||
"datetime": self._now().strftime("%Y-%m-%d %H:%M:%S"),
|
||||
}
|
||||
self.trades.append(trade)
|
||||
logger.info("[shadow] 成交 %s %s %d股 @%.2f 费%.2f",
|
||||
side, security, amount, fill, commission + stamp_duty)
|
||||
if self.on_trade is not None:
|
||||
try:
|
||||
self.on_trade(trade)
|
||||
except Exception as exc: # noqa: BLE001 - 落库失败不阻断撮合
|
||||
logger.warning("[shadow] on_trade 回调失败: %s", exc)
|
||||
@@ -0,0 +1,165 @@
|
||||
"""影子柜台常驻进程入口(P1,VPS Windows / miniQMT 行情)。
|
||||
|
||||
与组合实盘(``sanguo_portfolio.runner_live``)同一个 bullet_trade LiveEngine,
|
||||
唯一区别:broker_factory 换成 ShadowBroker(本地撮合,订单不出门)。
|
||||
策略/行情/调度完全同款 → 双轨一致性验证(设计 §8)的基础。
|
||||
|
||||
环境变量(复用 live_strategy.py 的 SANGUO_LIVE_* 命名 + 影子专属 SANGUO_SHADOW_*):
|
||||
SANGUO_LIVE_STRATEGY/_MAX_POOL/_BENCHMARK/_CASH 策略配置(live_strategy.py 读)
|
||||
SANGUO_SHADOW_DB / SANGUO_SHADOW_ACCOUNT_ID 落库目标(paper 库)
|
||||
SANGUO_SHADOW_COMMISSION/_STAMP/_MIN_COMM/_SLIPPAGE 费率滑点(对齐实盘券商参数)
|
||||
|
||||
手动用法(VPS 交易日):
|
||||
set SANGUO_LIVE_STRATEGY=all_weather
|
||||
set SANGUO_SHADOW_DB=C:\\sanguo_vnpy_v2\\data\\paper.db
|
||||
python -m sanguo_trader.shadow
|
||||
|
||||
不做多账户轮询:MVP 一进程一账户(与 runner_live 一致),多账户由 supervisor
|
||||
按 paper_accounts(engine='shadow')逐行拉子进程(后续接入)。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
# ENV GUARD 必须早于任何 bullet_trade import
|
||||
import os
|
||||
os.environ.setdefault("DEFAULT_DATA_PROVIDER", "miniqmt")
|
||||
|
||||
import logging
|
||||
import threading
|
||||
import time
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# 策略适配文件与组合实盘共用(读 SANGUO_LIVE_* env)
|
||||
ADAPTER_FILE = Path(__file__).resolve().parents[2] / "sanguo_portfolio" / "live_strategy.py"
|
||||
|
||||
|
||||
def shadow_env() -> Dict[str, str]:
|
||||
"""解析影子柜台 env(独立出来便于单测)。"""
|
||||
return {
|
||||
"db": os.environ.get("SANGUO_SHADOW_DB", ""),
|
||||
"account_id": os.environ.get("SANGUO_SHADOW_ACCOUNT_ID", ""),
|
||||
"commission": os.environ.get("SANGUO_SHADOW_COMMISSION", "0.0003"),
|
||||
"stamp": os.environ.get("SANGUO_SHADOW_STAMP", "0.001"),
|
||||
"min_comm": os.environ.get("SANGUO_SHADOW_MIN_COMM", "5"),
|
||||
"slippage": os.environ.get("SANGUO_SHADOW_SLIPPAGE", "0.001"),
|
||||
"snapshot_sec": os.environ.get("SANGUO_SHADOW_SNAPSHOT_SEC", "30"),
|
||||
}
|
||||
|
||||
|
||||
def build_price_getter(provider: Any) -> Any:
|
||||
"""从数据 provider 取标的最新价(实时/最新收盘)。返回闭包给 ShadowBroker。"""
|
||||
|
||||
def get_price(security: str) -> Optional[float]:
|
||||
from datetime import datetime, timedelta
|
||||
|
||||
end = datetime.now().strftime("%Y-%m-%d %H:%M:%S")
|
||||
start = (datetime.now() - timedelta(days=10)).strftime("%Y-%m-%d")
|
||||
try:
|
||||
df = provider.get_price(
|
||||
security=security, start_date=start, end_date=end,
|
||||
frequency="daily", fields=["close"], fq="pre",
|
||||
)
|
||||
if df is None or len(df) == 0:
|
||||
return None
|
||||
return float(df["close"].iloc[-1])
|
||||
except Exception: # noqa: BLE001 - provider 接口差异兜底
|
||||
cols = [c for c in ("close", "Close") if c in (df.columns if df is not None else [])]
|
||||
if cols:
|
||||
return float(df[cols[0]].iloc[-1])
|
||||
return None
|
||||
|
||||
return get_price
|
||||
|
||||
|
||||
def _paper_on_trade(db: str, account_id: int, strategy_id: str):
|
||||
"""成交回调:落 paper_trades(与组合实走 EOD 同表,前端模拟盘页直接可见)。"""
|
||||
from sanguo_trader.persistence import save_trade
|
||||
|
||||
def hook(trade: Dict[str, Any]) -> None:
|
||||
side = trade["side"]
|
||||
save_trade(db, account_id, {
|
||||
"strategy_id": strategy_id,
|
||||
"datetime": trade["datetime"],
|
||||
"symbol": trade["security"],
|
||||
"direction": "long" if side == "buy" else "short",
|
||||
"offset": "open" if side == "buy" else "close",
|
||||
"match_session": "shadow_realtime",
|
||||
"price": trade["price"],
|
||||
"volume": trade["amount"],
|
||||
"commission": trade["commission"],
|
||||
"stamp_duty": trade["stamp_duty"],
|
||||
"bar_date": trade["datetime"][:10],
|
||||
})
|
||||
|
||||
return hook
|
||||
|
||||
|
||||
def _snapshot_loop(broker: Any, db: str, account_id: int,
|
||||
interval_sec: float = 30.0) -> None:
|
||||
"""后台线程:定期把影子账户快照落 paper_positions/paper_daily_balance。"""
|
||||
from sanguo_trader.persistence import save_daily_balance, save_positions
|
||||
|
||||
while True:
|
||||
time.sleep(interval_sec)
|
||||
try:
|
||||
info = broker.get_account_info()
|
||||
positions = {
|
||||
p["security"]: {"volume": float(p["amount"]), "frozen": 0.0,
|
||||
"avg_price": p["avg_cost"]}
|
||||
for p in info["positions"]
|
||||
}
|
||||
save_positions(db, account_id, "account", positions,
|
||||
date=broker.trades[-1]["datetime"][:10] if broker.trades else "")
|
||||
save_daily_balance(
|
||||
db, account_id, info.get("as_of", ""),
|
||||
cash=info["available_cash"], market_value=info["market_value"],
|
||||
total_equity=info["total_value"],
|
||||
)
|
||||
except Exception as exc: # noqa: BLE001 - 落库失败不中断柜台
|
||||
logger.warning("[shadow-snapshot] 落库失败 (account=%s): %s", account_id, exc)
|
||||
|
||||
|
||||
def run_shadow(provider_config: Optional[Dict[str, Any]] = None) -> None:
|
||||
"""装配 LiveEngine(影子 broker)并 run(阻塞)。"""
|
||||
from bullet_trade.core.live_engine import LiveEngine # type: ignore
|
||||
from bullet_trade.data.api import set_data_provider # type: ignore
|
||||
|
||||
from sanguo_portfolio.runner_live import build_provider
|
||||
from .broker import ShadowBroker
|
||||
|
||||
cfg = shadow_env()
|
||||
from sanguo_portfolio.runner_live import live_env
|
||||
le = live_env()
|
||||
|
||||
provider = build_provider(provider_config)
|
||||
set_data_provider(provider)
|
||||
|
||||
broker = ShadowBroker(
|
||||
initial_cash=float(le["cash"]),
|
||||
commission_rate=float(cfg["commission"]),
|
||||
stamp_duty_rate=float(cfg["stamp"]),
|
||||
min_commission=float(cfg["min_comm"]),
|
||||
slippage=float(cfg["slippage"]),
|
||||
price_getter=build_price_getter(provider),
|
||||
on_trade=_paper_on_trade(cfg["db"], int(cfg["account_id"]), le["strategy"])
|
||||
if cfg["db"] and cfg["account_id"] else None,
|
||||
)
|
||||
logger.info(
|
||||
"影子柜台启动: strategy=%s cash=%s 费率=佣金%s/印花%s/最低%s 滑点%s db=%s",
|
||||
le["strategy"], le["cash"], cfg["commission"], cfg["stamp"],
|
||||
cfg["min_comm"], cfg["slippage"], cfg["db"] or "(不落库)",
|
||||
)
|
||||
|
||||
engine = LiveEngine(ADAPTER_FILE, broker_factory=lambda: broker)
|
||||
|
||||
if cfg["db"] and cfg["account_id"]:
|
||||
t = threading.Thread(
|
||||
target=_snapshot_loop,
|
||||
args=(broker, cfg["db"], int(cfg["account_id"]), float(cfg["snapshot_sec"])),
|
||||
daemon=True, name="shadow-snapshot",
|
||||
)
|
||||
t.start()
|
||||
|
||||
engine.run()
|
||||
@@ -37,6 +37,48 @@ def test_create_paper(tmp_path):
|
||||
assert "account_id" in resp.json()
|
||||
|
||||
|
||||
def test_create_shadow_mode(tmp_path):
|
||||
"""影子=第三种运行模式(CTA/组合都可):engine 由 mode 推导,不被日终 job 结算。"""
|
||||
c, token = _client(tmp_path)
|
||||
# CTA 影子
|
||||
r1 = c.post("/api/v1/paper/create", json={
|
||||
"mode": "shadow",
|
||||
"symbols": ["600000"],
|
||||
"strategies": [{"name": "DoubleMa", "symbol": "600000"}],
|
||||
"start": "2024-01-01", "end": "2024-06-30",
|
||||
}, headers=_auth(token))
|
||||
assert r1.status_code == 200
|
||||
# 组合影子:mode=shadow → engine=shadow
|
||||
r2 = c.post("/api/v1/paper/create", json={
|
||||
"mode": "shadow", "strategy_type": "portfolio",
|
||||
"symbols": ["hs300_subset"],
|
||||
"strategies": [{"name": "all_weather", "symbol": "hs300_subset"}],
|
||||
"start": "2024-01-01", "end": "2024-12-31",
|
||||
"pool": "hs300_subset",
|
||||
}, headers=_auth(token))
|
||||
assert r2.status_code == 200
|
||||
lst = c.get("/api/v1/paper", headers=_auth(token)).json()
|
||||
items = lst if isinstance(lst, list) else lst.get("accounts", lst.get("papers", []))
|
||||
by_id = {a["id"]: a for a in items}
|
||||
cta = by_id[r1.json()["account_id"]]
|
||||
assert cta["mode"] == "shadow"
|
||||
assert cta["engine"] == "shadow"
|
||||
pf = by_id[r2.json()["account_id"]]
|
||||
assert pf["mode"] == "shadow" and pf["engine"] == "shadow"
|
||||
|
||||
|
||||
def test_create_portfolio_rejects_replay(tmp_path):
|
||||
c, token = _client(tmp_path)
|
||||
resp = c.post("/api/v1/paper/create", json={
|
||||
"mode": "replay", "strategy_type": "portfolio",
|
||||
"symbols": ["hs300_subset"],
|
||||
"strategies": [{"name": "all_weather", "symbol": "hs300_subset"}],
|
||||
"start": "2024-01-01", "end": "2024-12-31",
|
||||
}, headers=_auth(token))
|
||||
assert resp.status_code == 400
|
||||
assert "组合回测" in resp.json()["detail"]
|
||||
|
||||
|
||||
def test_get_paper_and_empty_trades(tmp_path):
|
||||
c, token = _client(tmp_path)
|
||||
aid = c.post(
|
||||
|
||||
@@ -27,6 +27,7 @@ def _create_portfolio(db, **kw):
|
||||
strategy_type="portfolio", strategy_class=kw.get("strategy_class", "all_weather"),
|
||||
pool=kw.get("pool", "hs300_subset"), max_pool=kw.get("max_pool", 30),
|
||||
benchmark=kw.get("benchmark", "000300.XSHG"),
|
||||
interval=kw.get("interval", ""),
|
||||
initial_capital=kw.get("initial_capital", 500000),
|
||||
)
|
||||
return rl.create_live(req)["account_id"]
|
||||
@@ -72,10 +73,17 @@ def test_create_portfolio_live(live_db):
|
||||
acc = live_persistence.get_account(live_db, aid)
|
||||
assert acc["strategy_type"] == "portfolio"
|
||||
assert acc["vt_symbol"] == "hs300_subset" # 组合行 vt_symbol=池名
|
||||
assert acc["interval"] == "d"
|
||||
assert acc["interval"] == "d" # 未传周期 → 组合默认日线
|
||||
assert acc["status"] == "stopped"
|
||||
|
||||
|
||||
def test_create_portfolio_interval_passthrough(live_db):
|
||||
"""组合实盘周期由前端下拉传(miniQMT 档位),后端不写死。"""
|
||||
aid = _create_portfolio(live_db, interval="15m")
|
||||
acc = live_persistence.get_account(live_db, aid)
|
||||
assert acc["interval"] == "15m"
|
||||
|
||||
|
||||
def test_create_portfolio_rejects_empty_strategy(live_db):
|
||||
with pytest.raises(Exception):
|
||||
_create_portfolio(live_db, strategy_class="")
|
||||
|
||||
@@ -0,0 +1,86 @@
|
||||
"""通路测试策略(ChannelTestStrategy)单元测试。
|
||||
|
||||
验证轮换目标集 + 卖旧买新调度 + T+1 探针,用 mock broker 记录下单调用。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from sanguo_portfolio.strategies import ChannelTestConfig, ChannelTestStrategy
|
||||
from sanguo_portfolio.strategies.all_weather import BrokerFacade
|
||||
|
||||
|
||||
class _MockBroker(BrokerFacade):
|
||||
def __init__(self) -> None:
|
||||
# 注意:BrokerFacade 是 dataclass,父类 __init__ 会用字段默认值覆盖同名
|
||||
# 实例属性,子类方法重写会被遮蔽 → 必须在 super().__init__() 之后
|
||||
# 用实例属性注入记录函数。
|
||||
self.calls: list[tuple[str, str, float]] = [] # (method, code, value)
|
||||
super().__init__()
|
||||
self.order_target_value = self._record_otv
|
||||
|
||||
def _record_otv(self, code: str, value: float):
|
||||
self.calls.append(("otv", code, value))
|
||||
return None
|
||||
|
||||
|
||||
class _Pos:
|
||||
def __init__(self, value: float) -> None:
|
||||
self.value = value
|
||||
|
||||
|
||||
class _Ctx:
|
||||
def __init__(self, positions: dict, cash: float) -> None:
|
||||
self.portfolio = type("P", (), {"positions": positions,
|
||||
"available_cash": cash,
|
||||
"total_value": cash + sum(p.value for p in positions.values())})()
|
||||
|
||||
|
||||
def test_target_set_rotates_with_day():
|
||||
s = ChannelTestStrategy(provider=None, config=ChannelTestConfig(
|
||||
universe=["A", "B", "C", "D"], hold_n=2, period=1))
|
||||
s._day = 0
|
||||
assert set(s._target_set()) <= {"A", "B", "C", "D"}
|
||||
s._day = 1
|
||||
t1 = s._target_set()
|
||||
s._day = 2
|
||||
t2 = s._target_set()
|
||||
assert len(t1) == 2 and len(t2) == 2
|
||||
assert t1 != t2 # 不同周期目标偏移
|
||||
|
||||
|
||||
def test_rotate_sells_non_target_and_buys_target():
|
||||
broker = _MockBroker()
|
||||
s = ChannelTestStrategy(provider=None, broker=broker, config=ChannelTestConfig(
|
||||
universe=["A", "B", "C", "D"], hold_n=2, period=1, probe_t1=False))
|
||||
# rotate#1 → day=1 → offset=1 → target=[B,C];当前持仓 C,D
|
||||
# → 卖出 D(C 在目标内保留),对 B,C 调仓(买入)
|
||||
ctx = _Ctx({"C": _Pos(1000), "D": _Pos(1000)}, cash=8000)
|
||||
s.rotate(ctx)
|
||||
zero_calls = {c for m, c, v in broker.calls if m == "otv" and v == 0}
|
||||
assert zero_calls == {"D"} # 只卖非目标的 D
|
||||
buys = {c for m, c, v in broker.calls if m == "otv" and v > 0}
|
||||
assert buys == {"B", "C"} # 买新 B + 调仓 C
|
||||
|
||||
|
||||
def test_rotate_t1_probe_fires():
|
||||
broker = _MockBroker()
|
||||
s = ChannelTestStrategy(provider=None, broker=broker, config=ChannelTestConfig(
|
||||
universe=["A", "B"], hold_n=1, period=1, probe_t1=True))
|
||||
ctx = _Ctx({}, cash=10000)
|
||||
s.rotate(ctx)
|
||||
# rotate#1 → target=[B];探针对 target[0]=B 再次 order_target_value(0)(当日卖,预期被 T+1 拒)
|
||||
otv_calls = [c for m, c, v in broker.calls if m == "otv"]
|
||||
assert otv_calls.count("B") >= 2 # 一次买入调仓 + 一次 T+1 探针
|
||||
|
||||
|
||||
def test_period_skips_off_cycle_days():
|
||||
broker = _MockBroker()
|
||||
s = ChannelTestStrategy(provider=None, broker=broker, config=ChannelTestConfig(
|
||||
universe=["A", "B"], hold_n=1, period=3, probe_t1=False))
|
||||
ctx = _Ctx({}, cash=10000)
|
||||
s.rotate(ctx) # day1 → 触发
|
||||
n1 = len(broker.calls)
|
||||
s.rotate(ctx) # day2 → 跳过
|
||||
s.rotate(ctx) # day3 → 跳过
|
||||
assert len(broker.calls) == n1
|
||||
s.rotate(ctx) # day4 → 触发
|
||||
assert len(broker.calls) > n1
|
||||
@@ -0,0 +1,135 @@
|
||||
"""ShadowBroker(影子柜台本地撮合)单元测试。
|
||||
|
||||
纯逻辑测试:不依赖 bullet_trade/xtquant,价格由固定 price_getter 注入。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from datetime import datetime
|
||||
|
||||
import pytest
|
||||
|
||||
from sanguo_trader.shadow.broker import ShadowBroker
|
||||
|
||||
|
||||
def _mk_broker(cash: float = 100_000.0, **kw) -> ShadowBroker:
|
||||
prices = kw.pop("prices", {"600519.SH": 100.0})
|
||||
fixed_now = kw.pop("now", datetime(2026, 8, 14, 10, 0, 0))
|
||||
return ShadowBroker(
|
||||
initial_cash=cash,
|
||||
price_getter=lambda s: prices.get(s),
|
||||
now_provider=lambda: fixed_now,
|
||||
**kw,
|
||||
)
|
||||
|
||||
|
||||
def _buy(b: ShadowBroker, sec: str, amt: int, px: float | None = None):
|
||||
return asyncio.run(b.buy(sec, amt, px))
|
||||
|
||||
|
||||
def _sell(b: ShadowBroker, sec: str, amt: int, px: float | None = None):
|
||||
return asyncio.run(b.sell(sec, amt, px))
|
||||
|
||||
|
||||
def test_buy_fills_with_commission_and_slippage():
|
||||
b = _mk_broker(cash=100_000, slippage=0.001, commission_rate=0.0003, min_commission=5)
|
||||
oid = _buy(b, "600519.SH", 100)
|
||||
assert b.orders[oid]["status"] == "filled"
|
||||
fill = b.orders[oid]["filled_price"]
|
||||
assert fill == pytest.approx(100.0 * 1.001, abs=0.01) # 买入上浮滑点
|
||||
# 现金扣减 = 全额 + 佣金(低于最低佣金取 5 元)
|
||||
commission = max(100 * fill * 0.0003, 5.0)
|
||||
assert b.cash == pytest.approx(100_000 - 100 * fill - commission)
|
||||
assert b.positions["600519.SH"]["amount"] == 100
|
||||
assert b.positions["600519.SH"]["avg_cost"] == pytest.approx(fill)
|
||||
|
||||
|
||||
def test_sell_charges_stamp_duty_and_slippage_down():
|
||||
b = _mk_broker(cash=100_000, slippage=0.001, stamp_duty_rate=0.001)
|
||||
_buy(b, "600519.SH", 200, px=100.0) # 固定委托价,滑点仍生效
|
||||
cash_after_buy = b.cash
|
||||
# T+1:当日买入不可卖 → 先模拟次日(before_open 清锁)
|
||||
b.before_open()
|
||||
oid = _sell(b, "600519.SH", 200, px=100.0)
|
||||
assert b.orders[oid]["status"] == "filled"
|
||||
fill = b.orders[oid]["filled_price"]
|
||||
assert fill == pytest.approx(100.0 * 0.999, abs=0.01) # 卖出下压滑点
|
||||
gross = 200 * fill
|
||||
commission = max(gross * 0.0003, 5.0)
|
||||
stamp = gross * 0.001
|
||||
assert b.cash == pytest.approx(cash_after_buy + gross - commission - stamp)
|
||||
assert "600519.SH" not in b.positions # 清仓移除
|
||||
|
||||
|
||||
def test_t1_blocks_same_day_sell():
|
||||
b = _mk_broker()
|
||||
_buy(b, "600519.SH", 200, px=100.0)
|
||||
oid = _sell(b, "600519.SH", 200, px=100.0) # 当日卖 → 拒
|
||||
assert b.orders[oid]["status"] == "rejected"
|
||||
assert "T+1" in b.orders[oid]["reject_reason"]
|
||||
# 次日可卖
|
||||
b.before_open()
|
||||
oid2 = _sell(b, "600519.SH", 200, px=100.0)
|
||||
assert b.orders[oid2]["status"] == "filled"
|
||||
|
||||
|
||||
def test_insufficient_cash_rejects():
|
||||
b = _mk_broker(cash=5_000)
|
||||
oid = _buy(b, "600519.SH", 100) # 需约 1 万
|
||||
assert b.orders[oid]["status"] == "rejected"
|
||||
assert "资金不足" in b.orders[oid]["reject_reason"]
|
||||
assert b.cash == 5_000 # 拒单不动账
|
||||
|
||||
|
||||
def test_odd_lot_floors_to_100():
|
||||
b = _mk_broker(cash=1_000_000)
|
||||
oid = _buy(b, "600519.SH", 250) # → 200
|
||||
assert b.orders[oid]["status"] == "filled"
|
||||
assert b.orders[oid]["filled_amount"] == 200
|
||||
oid2 = _buy(b, "600519.SH", 50) # 不足一手 → 拒
|
||||
assert b.orders[oid2]["status"] == "rejected"
|
||||
|
||||
|
||||
def test_no_price_rejects():
|
||||
b = _mk_broker(prices={})
|
||||
oid = _buy(b, "600519.SH", 100)
|
||||
assert b.orders[oid]["status"] == "rejected"
|
||||
assert "无参考价" in b.orders[oid]["reject_reason"]
|
||||
|
||||
|
||||
def test_on_trade_callback_receives_fills():
|
||||
seen: list[dict] = []
|
||||
b = ShadowBroker(
|
||||
initial_cash=100_000,
|
||||
price_getter=lambda s: 10.0,
|
||||
on_trade=seen.append,
|
||||
)
|
||||
_buy(b, "600519.SH", 100)
|
||||
assert len(seen) == 1
|
||||
t = seen[0]
|
||||
assert t["side"] == "buy" and t["amount"] == 100 and t["price"] == pytest.approx(10.0)
|
||||
|
||||
|
||||
def test_get_account_info_totals():
|
||||
b = _mk_broker(cash=100_000)
|
||||
_buy(b, "600519.SH", 100, px=100.0)
|
||||
info = b.get_account_info()
|
||||
assert info["available_cash"] == pytest.approx(100_000 - 100 * 100.0 - max(100 * 100 * 0.0003, 5))
|
||||
assert info["market_value"] == pytest.approx(100 * 100.0) # price_getter=100
|
||||
assert info["total_value"] == pytest.approx(info["available_cash"] + info["market_value"])
|
||||
|
||||
|
||||
def test_avg_cost_weighted_on_second_buy():
|
||||
b = _mk_broker(cash=1_000_000, slippage=0.0)
|
||||
_buy(b, "600519.SH", 100, px=100.0)
|
||||
_buy(b, "600519.SH", 100, px=110.0)
|
||||
pos = b.positions["600519.SH"]
|
||||
assert pos["amount"] == 200
|
||||
assert pos["avg_cost"] == pytest.approx(105.0)
|
||||
|
||||
|
||||
def test_cancel_always_false_and_open_orders_empty():
|
||||
b = _mk_broker()
|
||||
_buy(b, "600519.SH", 100, px=100.0)
|
||||
assert asyncio.run(b.cancel_order("whatever")) is False
|
||||
assert b.get_open_orders() == [] # 即时成交,无挂单
|
||||
Reference in New Issue
Block a user