229 lines
4.9 KiB
Go
229 lines
4.9 KiB
Go
package logic
|
|
|
|
import (
|
|
"encoding/json"
|
|
"fmt"
|
|
"os"
|
|
"path/filepath"
|
|
"sync"
|
|
|
|
"big-qmt/go-client/config"
|
|
"big-qmt/go-client/sdk"
|
|
)
|
|
|
|
const (
|
|
pendingNone = ""
|
|
pendingBaseOpening = "base_opening"
|
|
pendingAdd = "add"
|
|
pendingSellAdd = "sell_add"
|
|
pendingSellBase = "sell_base"
|
|
)
|
|
|
|
type SymbolState struct {
|
|
Code string `json:"code"`
|
|
BaseQty int `json:"base_qty"`
|
|
BaseCost float64 `json:"base_cost"`
|
|
AddQty int `json:"add_qty"`
|
|
AddCost float64 `json:"add_cost"`
|
|
Pending string `json:"pending"`
|
|
PendingOrderID string `json:"pending_order_id,omitempty"`
|
|
OrderStatus string `json:"order_status,omitempty"`
|
|
}
|
|
|
|
func setPending(item *SymbolState, pending, orderID string) {
|
|
item.Pending = pending
|
|
item.PendingOrderID = orderID
|
|
item.OrderStatus = "submitted"
|
|
}
|
|
|
|
func clearPending(item *SymbolState) {
|
|
item.Pending = pendingNone
|
|
item.PendingOrderID = ""
|
|
item.OrderStatus = ""
|
|
}
|
|
|
|
type filePayload struct {
|
|
Version int `json:"version"`
|
|
Data map[string]any `json:"data"`
|
|
}
|
|
|
|
type ZTState struct {
|
|
path string
|
|
Items map[string]*SymbolState
|
|
LoadError string
|
|
fresh bool
|
|
mu sync.Mutex
|
|
}
|
|
|
|
var (
|
|
statesMu sync.Mutex
|
|
states = map[string]*ZTState{}
|
|
)
|
|
|
|
// BootstrapState 在状态文件首次不存在时,将启动前已有持仓登记为底仓。
|
|
func BootstrapState(positions []sdk.Position) {
|
|
state := getState()
|
|
if !state.Fresh() {
|
|
return
|
|
}
|
|
for _, pos := range positions {
|
|
code := pos.StockCode
|
|
if code == "" || pos.Volume <= 0 || pos.OpenPrice <= 0 {
|
|
continue
|
|
}
|
|
item := state.Ensure(code)
|
|
item.BaseQty, item.BaseCost = pos.Volume, pos.OpenPrice
|
|
logf("WARNING", "[ZT][状态] %s 首次接管为底仓 数量=%d 成本=%.2f", code, pos.Volume, pos.OpenPrice)
|
|
}
|
|
state.completeBootstrap()
|
|
state.Save()
|
|
}
|
|
|
|
func getState() *ZTState {
|
|
statesMu.Lock()
|
|
defer statesMu.Unlock()
|
|
accountID := config.Account.AccountID
|
|
if s, ok := states[accountID]; ok {
|
|
return s
|
|
}
|
|
s := loadZTState(config.Global.QMTDataDir, accountID)
|
|
states[accountID] = s
|
|
return s
|
|
}
|
|
|
|
func loadZTState(dataDir, accountID string) *ZTState {
|
|
st := &ZTState{
|
|
path: filepath.Join(dataDir, fmt.Sprintf("zt_%s_state.json", accountID)),
|
|
Items: map[string]*SymbolState{},
|
|
}
|
|
raw, err := os.ReadFile(st.path)
|
|
if err != nil {
|
|
if os.IsNotExist(err) {
|
|
st.fresh = true
|
|
return st
|
|
}
|
|
st.LoadError = err.Error()
|
|
logf("ERROR", "[ZT][状态] 读取状态文件失败: %v", err)
|
|
return st
|
|
}
|
|
var payload filePayload
|
|
if err := json.Unmarshal(raw, &payload); err != nil || payload.Version != 1 {
|
|
st.LoadError = "状态文件版本无效"
|
|
logf("ERROR", "[ZT][状态] %s", st.LoadError)
|
|
return st
|
|
}
|
|
data := payload.Data
|
|
if data == nil {
|
|
st.LoadError = "状态文件内容无效"
|
|
logf("ERROR", "[ZT][状态] %s", st.LoadError)
|
|
return st
|
|
}
|
|
symbolsAny, _ := data["symbols"]
|
|
symbols, _ := symbolsAny.(map[string]any)
|
|
if symbols == nil {
|
|
if _, ok := data["code"]; ok {
|
|
symbols = map[string]any{}
|
|
} else {
|
|
symbols = data
|
|
}
|
|
}
|
|
for code, value := range symbols {
|
|
m, ok := value.(map[string]any)
|
|
if !ok {
|
|
continue
|
|
}
|
|
item := &SymbolState{Code: code}
|
|
b, _ := json.Marshal(m)
|
|
_ = json.Unmarshal(b, item)
|
|
item.Code = code
|
|
st.Items[code] = item
|
|
}
|
|
return st
|
|
}
|
|
|
|
func (s *ZTState) Fresh() bool {
|
|
s.mu.Lock()
|
|
defer s.mu.Unlock()
|
|
return s.fresh
|
|
}
|
|
|
|
func (s *ZTState) completeBootstrap() {
|
|
s.mu.Lock()
|
|
defer s.mu.Unlock()
|
|
s.fresh = false
|
|
}
|
|
|
|
func (s *ZTState) Get(code string) *SymbolState {
|
|
s.mu.Lock()
|
|
defer s.mu.Unlock()
|
|
return s.Items[code]
|
|
}
|
|
|
|
func (s *ZTState) Ensure(code string) *SymbolState {
|
|
s.mu.Lock()
|
|
defer s.mu.Unlock()
|
|
if item, ok := s.Items[code]; ok {
|
|
return item
|
|
}
|
|
item := &SymbolState{Code: code}
|
|
s.Items[code] = item
|
|
return item
|
|
}
|
|
|
|
func (s *ZTState) Remove(code string) {
|
|
s.mu.Lock()
|
|
defer s.mu.Unlock()
|
|
delete(s.Items, code)
|
|
}
|
|
|
|
func (s *ZTState) Codes() []string {
|
|
s.mu.Lock()
|
|
defer s.mu.Unlock()
|
|
out := make([]string, 0, len(s.Items))
|
|
for code := range s.Items {
|
|
out = append(out, code)
|
|
}
|
|
return out
|
|
}
|
|
|
|
func (s *ZTState) Save() {
|
|
s.mu.Lock()
|
|
defer s.mu.Unlock()
|
|
if s.LoadError != "" {
|
|
return
|
|
}
|
|
if err := s.saveUnlocked(); err != nil {
|
|
s.LoadError = err.Error()
|
|
logf("ERROR", "[ZT][状态] 保存失败: %v", err)
|
|
return
|
|
}
|
|
}
|
|
|
|
func (s *ZTState) saveUnlocked() error {
|
|
symbols := map[string]any{}
|
|
for code, item := range s.Items {
|
|
symbols[code] = item
|
|
}
|
|
payload := filePayload{Version: 1, Data: map[string]any{"symbols": symbols}}
|
|
raw, err := json.Marshal(payload)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
if err := os.MkdirAll(filepath.Dir(s.path), 0o755); err != nil {
|
|
return err
|
|
}
|
|
tmp := s.path + ".tmp"
|
|
if err := os.WriteFile(tmp, raw, 0o644); err != nil {
|
|
return err
|
|
}
|
|
return replaceFile(tmp, s.path)
|
|
}
|
|
|
|
func replaceFile(tmp, dest string) error {
|
|
if err := os.Rename(tmp, dest); err == nil {
|
|
return nil
|
|
}
|
|
_ = os.Remove(dest)
|
|
return os.Rename(tmp, dest)
|
|
}
|