Files
stock/backend/internal/services/stock_sync.go
T

321 lines
9.1 KiB
Go

package services
import (
"context"
"fmt"
"log"
"math"
"sync"
"time"
"stock-user-system/internal/config"
"stock-user-system/internal/datasource"
"stock-user-system/internal/models"
"gorm.io/gorm"
"gorm.io/gorm/clause"
)
var (
coreIndexSymbols = []string{"000001.SH", "399001.SZ", "399006.SZ", "000688.SH"}
chinaLoc, _ = time.LoadLocation("Asia/Shanghai")
)
// StockSyncService 负责盘后行情同步。
type StockSyncService struct {
cfg *config.Config
db *gorm.DB
client *datasource.Client
mu sync.Mutex
}
// NewStockSyncService 创建同步服务。
func NewStockSyncService(cfg *config.Config, db *gorm.DB) *StockSyncService {
return &StockSyncService{
cfg: cfg,
db: db,
client: datasource.NewClient(cfg),
}
}
// Sync 执行一次盘后同步,返回写入的记录数。
func (s *StockSyncService) Sync(ctx context.Context, triggerBy string) (int, error) {
return s.runSyncJob(ctx, triggerBy, func() (int, error) {
return s.doSync(ctx)
})
}
// SyncHistory 同步近 N 个月的历史日 K 数据,返回写入的记录数。
func (s *StockSyncService) SyncHistory(ctx context.Context, triggerBy string, months int) (int, error) {
return s.runSyncJob(ctx, triggerBy, func() (int, error) {
return s.doSyncHistory(ctx, months)
})
}
// StartSync 异步启动一次盘后同步,返回已创建的 job。
func (s *StockSyncService) StartSync(ctx context.Context, triggerBy string) (*models.StockSyncJob, error) {
job, err := s.createJob(ctx, triggerBy)
if err != nil {
return nil, err
}
go s.runSyncJobByID(context.Background(), job.ID, func(ctx context.Context) (int, error) {
return s.doSync(ctx)
})
return job, nil
}
// StartSyncHistory 异步启动近 N 个月历史同步,返回已创建的 job。
func (s *StockSyncService) StartSyncHistory(ctx context.Context, triggerBy string, months int) (*models.StockSyncJob, error) {
job, err := s.createJob(ctx, triggerBy)
if err != nil {
return nil, err
}
go s.runSyncJobByID(context.Background(), job.ID, func(ctx context.Context) (int, error) {
return s.doSyncHistory(ctx, months)
})
return job, nil
}
func (s *StockSyncService) createJob(ctx context.Context, triggerBy string) (*models.StockSyncJob, error) {
job := models.StockSyncJob{
JobDate: today(),
Status: "running",
StartedAt: time.Now().UTC(),
TriggerBy: triggerBy,
}
if err := s.db.WithContext(ctx).Create(&job).Error; err != nil {
return nil, fmt.Errorf("create sync job: %w", err)
}
return &job, nil
}
func (s *StockSyncService) runSyncJob(ctx context.Context, triggerBy string, fn func() (int, error)) (int, error) {
job, err := s.createJob(ctx, triggerBy)
if err != nil {
return 0, err
}
recordsCount, err := fn()
s.finishJob(ctx, job.ID, recordsCount, err)
return recordsCount, err
}
func (s *StockSyncService) runSyncJobByID(ctx context.Context, jobID string, fn func(context.Context) (int, error)) {
s.mu.Lock()
defer s.mu.Unlock()
recordsCount, err := fn(ctx)
s.finishJob(ctx, jobID, recordsCount, err)
}
func (s *StockSyncService) finishJob(ctx context.Context, jobID string, recordsCount int, jobErr error) {
now := time.Now().UTC()
updates := map[string]any{
"records_count": recordsCount,
"finished_at": &now,
}
if jobErr != nil {
updates["status"] = "failed"
updates["error_message"] = jobErr.Error()
log.Printf("[stock sync] job %s failed: %v", jobID, jobErr)
} else {
updates["status"] = "success"
log.Printf("[stock sync] job %s success, records=%d", jobID, recordsCount)
}
if err := s.db.WithContext(ctx).Model(&models.StockSyncJob{}).Where("id = ?", jobID).Updates(updates).Error; err != nil {
log.Printf("[stock sync] update job %s failed: %v", jobID, err)
}
}
func (s *StockSyncService) doSync(ctx context.Context) (int, error) {
// 1. 拉取全 A 股行情快照
stockQuotes, err := s.client.GetQuotesByUniverses([]string{"CN_Equity_A"})
if err != nil {
return 0, fmt.Errorf("fetch stock quotes: %w", err)
}
// 2. 拉取核心指数行情
indexQuotes, err := s.client.GetQuotesBySymbols(coreIndexSymbols)
if err != nil {
return 0, fmt.Errorf("fetch index quotes: %w", err)
}
quotes := append(stockQuotes, indexQuotes...)
if len(quotes) == 0 {
return 0, fmt.Errorf("no quotes returned")
}
log.Printf("[stock sync] fetched %d stock quotes from CN_Equity_A, %d index quotes, total %d",
len(stockQuotes), len(indexQuotes), len(quotes))
// 3. 归一化并写入 DB
records := make([]models.StockDailyQuote, 0, len(quotes))
for _, q := range quotes {
tradeDate := msToDate(q.Timestamp)
open := round2(q.Open)
high := round2(q.High)
low := round2(q.Low)
close := round2(q.Close)
prevClose := round2(q.PrevClose)
var changePct float64
if prevClose != 0 {
changePct = (close - prevClose) / prevClose
}
records = append(records, models.StockDailyQuote{
Symbol: q.Symbol,
Name: q.DisplayName(),
TradeDate: tradeDate,
Open: open,
High: high,
Low: low,
Close: close,
PrevClose: prevClose,
Volume: q.Volume,
Amount: q.Amount,
ChangePct: changePct,
TurnoverRate: q.TurnoverRate(),
Region: q.Region,
})
}
if err := s.db.WithContext(ctx).Clauses(clause.OnConflict{
Columns: []clause.Column{{Name: "symbol"}, {Name: "trade_date"}},
UpdateAll: true,
}).CreateInBatches(records, 500).Error; err != nil {
return 0, fmt.Errorf("upsert daily quotes: %w", err)
}
return len(records), nil
}
func (s *StockSyncService) LatestJob(ctx context.Context) (*models.StockSyncJob, error) {
var job models.StockSyncJob
if err := s.db.WithContext(ctx).
Where("started_at > ?", time.Time{}).
Order("started_at DESC").
First(&job).Error; err != nil {
if err == gorm.ErrRecordNotFound {
return nil, nil
}
return nil, err
}
return &job, nil
}
func (s *StockSyncService) doSyncHistory(ctx context.Context, months int) (int, error) {
// 1. 获取全 A 股代码列表
universe, err := s.client.GetUniverse("CN_Equity_A")
if err != nil {
return 0, fmt.Errorf("fetch universe: %w", err)
}
if universe == nil || len(universe.Symbols) == 0 {
return 0, fmt.Errorf("empty universe")
}
// 2. 获取核心指数列表
allSymbols := append([]string{}, universe.Symbols...)
allSymbols = append(allSymbols, coreIndexSymbols...)
// 3. 计算时间范围(近 N 个月)
endMs := time.Now().In(chinaLoc).UnixMilli()
startMs := time.Now().In(chinaLoc).AddDate(0, -months, 0).UnixMilli()
log.Printf("[stock sync] start history sync for %d months, symbols=%d, range=%d-%d",
months, len(allSymbols), startMs, endMs)
// 4. 分批拉取并写入
const batchSize = 50
totalRecords := 0
for i := 0; i < len(allSymbols); i += batchSize {
batch := allSymbols[i:min(i+batchSize, len(allSymbols))]
klinesMap, err := s.client.GetKlinesBatch(batch, "1d", startMs, endMs)
if err != nil {
log.Printf("[stock sync] klines batch %d-%d failed: %v", i+1, min(i+batchSize, len(allSymbols)), err)
continue
}
records := make([]models.StockDailyQuote, 0)
for sym, data := range klinesMap {
name := sym
records = append(records, klineDataToRecords(sym, name, data)...)
}
if len(records) == 0 {
continue
}
if err := s.db.WithContext(ctx).Clauses(clause.OnConflict{
Columns: []clause.Column{{Name: "symbol"}, {Name: "trade_date"}},
UpdateAll: true,
}).CreateInBatches(records, 500).Error; err != nil {
return totalRecords, fmt.Errorf("upsert history batch %d: %w", i/batchSize, err)
}
totalRecords += len(records)
log.Printf("[stock sync] history batch %d/%d done, records=%d, total=%d",
(i/batchSize)+1, (len(allSymbols)+batchSize-1)/batchSize, len(records), totalRecords)
}
return totalRecords, nil
}
func klineDataToRecords(symbol, name string, data datasource.KlineData) []models.StockDailyQuote {
n := len(data.Timestamp)
if n == 0 {
return nil
}
if len(data.Open) < n { n = len(data.Open) }
if len(data.High) < n { n = len(data.High) }
if len(data.Low) < n { n = len(data.Low) }
if len(data.Close) < n { n = len(data.Close) }
if len(data.Volume) < n { n = len(data.Volume) }
if len(data.Amount) < n { n = len(data.Amount) }
records := make([]models.StockDailyQuote, 0, n)
for i := 0; i < n; i++ {
open := round2(data.Open[i])
high := round2(data.High[i])
low := round2(data.Low[i])
close := round2(data.Close[i])
var prevClose float64
if i > 0 {
prevClose = round2(data.Close[i-1])
}
var changePct float64
if prevClose != 0 {
changePct = (close - prevClose) / prevClose
}
records = append(records, models.StockDailyQuote{
Symbol: symbol,
Name: name,
TradeDate: msToDate(data.Timestamp[i]),
Open: open,
High: high,
Low: low,
Close: close,
PrevClose: prevClose,
Volume: data.Volume[i],
Amount: data.Amount[i],
ChangePct: changePct,
})
}
return records
}
func today() time.Time {
return time.Now().In(chinaLoc).Truncate(24 * time.Hour)
}
func msToDate(ms int64) time.Time {
return time.UnixMilli(ms).In(chinaLoc).Truncate(24 * time.Hour)
}
func round2(v float64) float64 {
return math.Round(v*100) / 100
}
func min(a, b int) int {
if a < b {
return a
}
return b
}