8/12/2025

并发拉取k线数据,并同步到数据库

package main

import (
	"encoding/json"
	"fmt"
	"log"
	"net/http"
	"strconv"
	"time"

	"golang.org/x/sync/errgroup"
	"gorm.io/driver/sqlite"
	"gorm.io/gorm"
)

// ================= 数据模型 =================
type Kline struct {
	ID        uint   `gorm:"primaryKey"`
	Symbol    string `gorm:"index:idx_symbol_open_time"`
	OpenTime  int64  `gorm:"index:idx_symbol_open_time"`
	Open      float64
	High      float64
	Low       float64
	Close     float64
	Volume    float64
	CloseTime int64
}

func (Kline) TableName() string {
	return "kline"
}

// ================= 币安 API 拉取 =================
func fetchBinanceKlines(symbol string, interval string, startTime, endTime int64, limit int) ([]Kline, error) {
	url := fmt.Sprintf(
		"https://fapi.binance.com/fapi/v1/klines?symbol=%s&interval=%s&limit=%d",
		symbol, interval, limit,
	)
	if startTime > 0 {
		url += fmt.Sprintf("&startTime=%d", startTime)
	}
	if endTime > 0 {
		url += fmt.Sprintf("&endTime=%d", endTime)
	}

	resp, err := http.Get(url)
	if err != nil {
		return nil, err
	}
	defer resp.Body.Close()

	var raw [][]interface{}
	if err := json.NewDecoder(resp.Body).Decode(&raw); err != nil {
		return nil, err
	}

	klines := make([]Kline, 0, len(raw))
	for _, item := range raw {
		openTime := int64(item[0].(float64))
		open, _ := strconv.ParseFloat(item[1].(string), 64)
		high, _ := strconv.ParseFloat(item[2].(string), 64)
		low, _ := strconv.ParseFloat(item[3].(string), 64)
		closePrice, _ := strconv.ParseFloat(item[4].(string), 64)
		volume, _ := strconv.ParseFloat(item[5].(string), 64)
		closeTime := int64(item[6].(float64))

		klines = append(klines, Kline{
			Symbol:    symbol,
			OpenTime:  openTime,
			Open:      open,
			High:      high,
			Low:       low,
			Close:     closePrice,
			Volume:    volume,
			CloseTime: closeTime,
		})
	}
	return klines, nil
}

// ================= 数据更新逻辑 =================
func updateKlines(db *gorm.DB, symbol string) error {
	var last Kline
	res := db.Where("symbol = ?", symbol).Order("open_time DESC").Limit(1).Find(&last)

	var startTime int64
	if res.RowsAffected == 0 {
		// 没有数据,从当前时间回溯 99*5m
		startTime = time.Now().Add(-time.Duration(99*5) * time.Minute).UnixMilli()
	} else {
		// 有数据,从最新一条的时间开始拉取
		if time.Since(time.UnixMilli(last.OpenTime)) >= 5*time.Minute {
			startTime = last.OpenTime
		} else {
			return nil // 最新数据足够
		}
	}

	klines, err := fetchBinanceKlines(symbol, "5m", startTime, 0, 99)
	if err != nil {
		return err
	}

	for _, k := range klines {
		var existing Kline
		if err := db.Where("symbol = ? AND open_time = ?", k.Symbol, k.OpenTime).First(&existing).Error; err == nil {
			// 更新(收盘时间可能未完成)
			db.Model(&existing).Updates(k)
		} else {
			// 新增
			db.Create(&k)
		}
	}
	return nil
}

// ================= 动态窗口聚合查询 =================
func queryAggregatedKlines(db *gorm.DB, symbol string, interval string) ([][]interface{}, error) {
	var bucketMs int64
	switch interval {
	case "15m":
		bucketMs = 15 * 60 * 1000
	case "1h":
		bucketMs = 60 * 60 * 1000
	case "4h":
		bucketMs = 4 * 60 * 60 * 1000
	case "1d":
		bucketMs = 24 * 60 * 60 * 1000
	default:
		return nil, fmt.Errorf("unsupported interval: %s", interval)
	}

	query := fmt.Sprintf(`
		WITH base AS (
		SELECT symbol, open_time, open, high, low, close, volume, close_time, CAST(open_time / %[1]d AS INTEGER) * %[1]d AS bucket_start
		FROM kline WHERE symbol = ?
		),
		agg AS (
		SELECT
			symbol,
			bucket_start,
			FIRST_VALUE(open) OVER (PARTITION BY bucket_start ORDER BY open_time ASC) AS open,
			MAX(high)   OVER (PARTITION BY bucket_start) AS high,
			MIN(low)    OVER (PARTITION BY bucket_start) AS low,
			FIRST_VALUE(close) OVER (PARTITION BY bucket_start ORDER BY open_time DESC) AS close,
			SUM(volume) OVER (PARTITION BY bucket_start) AS volume,
			MAX(close_time) OVER (PARTITION BY bucket_start) AS close_time,
			ROW_NUMBER() OVER (PARTITION BY bucket_start ORDER BY open_time ASC) AS rn
		FROM base
		)
		SELECT symbol, bucket_start AS open_time, open, high, low, close, volume, close_time FROM agg WHERE rn = 1 ORDER BY open_time;
	`, bucketMs)

	rows, err := db.Raw(query, symbol).Rows()
	if err != nil {
		return nil, err
	}
	defer rows.Close()

	// var result []Kline
	// 按币安 API 返回格式组装(二维数组)
	resp := make([][]interface{}, 0)
	for rows.Next() {
		var k Kline
		if err := rows.Scan(&k.Symbol, &k.OpenTime, &k.Open, &k.High, &k.Low, &k.Close, &k.Volume, &k.CloseTime); err != nil {
			return nil, err
		}
		// result = append(result, k)

		resp = append(resp, []interface{}{
			k.OpenTime,                    // 开盘时间 (ms)
			fmt.Sprintf("%.8f", k.Open),   // 开盘价
			fmt.Sprintf("%.8f", k.High),   // 最高价
			fmt.Sprintf("%.8f", k.Low),    // 最低价
			fmt.Sprintf("%.8f", k.Close),  // 收盘价
			fmt.Sprintf("%.8f", k.Volume), // 成交量
			k.CloseTime,                   // 收盘时间 (ms)
			"0",                           // Quote asset volume
			0,                             // Number of trades
			"0",                           // Taker buy base asset volume
			"0",                           // Taker buy quote asset volume
			"0",                           // Ignore
		})
	}

	return resp, nil
}

// ================= 主程序 =================
func main() {
	loc, err := time.LoadLocation("Asia/Shanghai")
	if err != nil {
		log.Fatal(err)
	}

	db, err := gorm.Open(sqlite.Open("klines.db"), &gorm.Config{
		NowFunc: func() time.Time {
			return time.Now().In(loc)
		},
	})
	if err != nil {
		log.Fatal(err)
	}

	// 确认当前时间
	var now time.Time
	db.Raw("SELECT CURRENT_TIMESTAMP").Scan(&now)
	log.Println("当前时间(东八区):", now)

	db.AutoMigrate(&Kline{})

	// 从 symbols.json 读取 symbols
	symbols, err := loadSymbolsFromFile("symbols.json")
	if err != nil {
		log.Fatalf("读取 symbols.json 失败: %v", err)
	}

	// 启动 HTTP 服务
	go func() {
		http.HandleFunc("/klines", handleKlineQuery(db))
		log.Println("HTTP server started on :8080")
		if err := http.ListenAndServe(":9090", nil); err != nil {
			log.Fatal(err)
		}
	}()
	// 定时任务:每分钟更新一次
	ticker := time.NewTicker(1 * time.Minute)
	go func() {
		for range ticker.C {
			if err := processSymbols(symbols, db); err != nil {
				log.Println("部分任务失败:", err)
			}
		}
	}()

	select {}
}

func processSymbols(symbols []string, db *gorm.DB) error {
	var g errgroup.Group
	sem := make(chan struct{}, 2) // 限制并行 4 个

	for _, sym := range symbols {
		sym := sym // 避免闭包变量问题

		g.Go(func() error {
			sem <- struct{}{} // 占用一个并发槽
			defer func() { <-sem }()

			if err := updateKlines(db, sym); err != nil {
				log.Println("update error:", sym, err)
				return err
			}
			log.Println("updated", sym)
			return nil
		})
	}

	return g.Wait()
}

// ================= HTTP 接口 =================
func handleKlineQuery(db *gorm.DB) http.HandlerFunc {
	return func(w http.ResponseWriter, r *http.Request) {
		// 允许跨域
		w.Header().Set("Access-Control-Allow-Origin", "*")
		w.Header().Set("Access-Control-Allow-Methods", "GET, OPTIONS")
		w.Header().Set("Access-Control-Allow-Headers", "Content-Type")

		if r.Method == http.MethodOptions {
			return // 处理预检请求
		}

		symbol := r.URL.Query().Get("symbol")
		interval := r.URL.Query().Get("interval")

		if symbol == "" || interval == "" {
			http.Error(w, "missing symbol or interval", http.StatusBadRequest)
			return
		}
		time1 := time.Now()
		data, err := queryAggregatedKlines(db, symbol, interval)
		if err != nil {
			http.Error(w, fmt.Sprintf("query error: %v", err), http.StatusInternalServerError)
			return
		}
		fmt.Println("统计", time.Since(time1).Milliseconds())

		// 判断是否支持 gzip
		if strings.Contains(r.Header.Get("Accept-Encoding"), "gzip") {
			w.Header().Set("Content-Encoding", "gzip")
			gz := gzip.NewWriter(w)
			defer gz.Close()

			w.Header().Set("Content-Type", "application/json")
			json.NewEncoder(gz).Encode(data)
			return
		}

		// 不支持 gzip,直接返回
		w.Header().Set("Content-Type", "application/json")
		json.NewEncoder(w).Encode(data)
	}
}
func loadSymbolsFromFile(filename string) ([]string, error) {
	data, err := os.ReadFile(filename)
	if err != nil {
		return nil, err
	}

	var symbols []string
	if err := json.Unmarshal(data, &symbols); err != nil {
		return nil, err
	}
	return symbols, nil
}