切换股票数据源为 Tushare
This commit is contained in:
@@ -10,6 +10,8 @@ import (
|
||||
"github.com/joho/godotenv"
|
||||
)
|
||||
|
||||
const defaultTushareToken = "76efd8465f9f2591aa42a385268e06acf6b80b7a15be2267ad2281b7"
|
||||
|
||||
type Config struct {
|
||||
DatabaseURL string
|
||||
JWTSecret string
|
||||
@@ -18,8 +20,7 @@ type Config struct {
|
||||
AllowedOrigins []string
|
||||
|
||||
// 股票数据源
|
||||
TickFlowAPIKey string
|
||||
TickFlowBaseURL string
|
||||
TushareToken string
|
||||
DataSyncEnabled bool
|
||||
DataSyncTime string // HH:MM
|
||||
DataSyncWeekdays bool
|
||||
@@ -57,14 +58,9 @@ func Load() (*Config, error) {
|
||||
}
|
||||
}
|
||||
|
||||
tickFlowAPIKey := os.Getenv("TICKFLOW_API_KEY")
|
||||
if tickFlowAPIKey == "" {
|
||||
tickFlowAPIKey = "tk_94a20304993f45b5b0e376b9767597cc"
|
||||
}
|
||||
|
||||
tickFlowBaseURL := os.Getenv("TICKFLOW_BASE_URL")
|
||||
if tickFlowBaseURL == "" {
|
||||
tickFlowBaseURL = "https://api.tickflow.org"
|
||||
tushareToken := os.Getenv("TUSHARE_TOKEN")
|
||||
if tushareToken == "" {
|
||||
tushareToken = defaultTushareToken
|
||||
}
|
||||
|
||||
dataSyncEnabled := strings.ToLower(os.Getenv("DATA_SYNC_ENABLED")) != "false"
|
||||
@@ -80,8 +76,7 @@ func Load() (*Config, error) {
|
||||
JWTExpirationHours: expHours,
|
||||
Port: port,
|
||||
AllowedOrigins: allowedOrigins,
|
||||
TickFlowAPIKey: tickFlowAPIKey,
|
||||
TickFlowBaseURL: tickFlowBaseURL,
|
||||
TushareToken: tushareToken,
|
||||
DataSyncEnabled: dataSyncEnabled,
|
||||
DataSyncTime: dataSyncTime,
|
||||
DataSyncWeekdays: dataSyncWeekdays,
|
||||
|
||||
@@ -1,241 +0,0 @@
|
||||
package datasource
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"stock-user-system/internal/config"
|
||||
)
|
||||
|
||||
const (
|
||||
defaultTimeout = 60 * time.Second
|
||||
maxRetries = 3
|
||||
retryBaseDelay = 1 * time.Second
|
||||
)
|
||||
|
||||
// Quote 对应 tickflow /v1/quotes 返回的单条行情数据。
|
||||
type Quote struct {
|
||||
Symbol string `json:"symbol"`
|
||||
Name string `json:"name"`
|
||||
Open float64 `json:"open"`
|
||||
High float64 `json:"high"`
|
||||
Low float64 `json:"low"`
|
||||
Close float64 `json:"last_price"`
|
||||
PrevClose float64 `json:"prev_close"`
|
||||
Volume int64 `json:"volume"`
|
||||
Amount float64 `json:"amount"`
|
||||
Timestamp int64 `json:"timestamp"`
|
||||
Region string `json:"region"`
|
||||
Ext struct {
|
||||
ChangePct float64 `json:"change_pct"`
|
||||
TurnoverRate float64 `json:"turnover_rate"`
|
||||
Name string `json:"name"`
|
||||
} `json:"ext"`
|
||||
}
|
||||
|
||||
// ChangePct 优先使用 ext.change_pct(tickflow 返回小数,如 0.30 = 30%),
|
||||
// 否则根据 close/prev_close 计算并统一返回小数形式。
|
||||
func (q *Quote) ChangePct() float64 {
|
||||
if q.Ext.ChangePct != 0 {
|
||||
return q.Ext.ChangePct
|
||||
}
|
||||
if q.PrevClose != 0 {
|
||||
return (q.Close - q.PrevClose) / q.PrevClose
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
// TurnoverRate 返回 ext.turnover_rate(仅 A 股有效)。
|
||||
func (q *Quote) TurnoverRate() float64 {
|
||||
return q.Ext.TurnoverRate
|
||||
}
|
||||
|
||||
// DisplayName 优先使用 ext.name,否则使用 symbol。
|
||||
func (q *Quote) DisplayName() string {
|
||||
if q.Ext.Name != "" {
|
||||
return q.Ext.Name
|
||||
}
|
||||
if q.Name != "" {
|
||||
return q.Name
|
||||
}
|
||||
return q.Symbol
|
||||
}
|
||||
|
||||
// UniverseDetail 对应 /v1/universes/:id 返回的标的池详情。
|
||||
type UniverseDetail struct {
|
||||
ID string `json:"id"`
|
||||
Name string `json:"name"`
|
||||
SymbolCount int `json:"symbol_count"`
|
||||
Symbols []string `json:"symbols"`
|
||||
}
|
||||
|
||||
// KlineData 对应 /v1/klines/batch 返回的单标的历史 K 线(紧凑列式)。
|
||||
type KlineData struct {
|
||||
Symbol string `json:"-"`
|
||||
Name string `json:"-"`
|
||||
Timestamp []int64 `json:"timestamp"`
|
||||
Open []float64 `json:"open"`
|
||||
High []float64 `json:"high"`
|
||||
Low []float64 `json:"low"`
|
||||
Close []float64 `json:"close"`
|
||||
Volume []int64 `json:"volume"`
|
||||
Amount []float64 `json:"amount"`
|
||||
}
|
||||
|
||||
// Client 封装 tickflow HTTP API 调用。
|
||||
type Client struct {
|
||||
baseURL string
|
||||
apiKey string
|
||||
client *http.Client
|
||||
}
|
||||
|
||||
// NewClient 从配置创建 tickflow 客户端。
|
||||
func NewClient(cfg *config.Config) *Client {
|
||||
return &Client{
|
||||
baseURL: strings.TrimRight(cfg.TickFlowBaseURL, "/"),
|
||||
apiKey: cfg.TickFlowAPIKey,
|
||||
client: &http.Client{Timeout: defaultTimeout},
|
||||
}
|
||||
}
|
||||
|
||||
// GetUniverse 获取指定标的池的完整代码列表。
|
||||
func (c *Client) GetUniverse(id string) (*UniverseDetail, error) {
|
||||
path := fmt.Sprintf("/v1/universes/%s", id)
|
||||
var detail UniverseDetail
|
||||
if err := c.request("GET", path, nil, nil, &detail); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &detail, nil
|
||||
}
|
||||
|
||||
// GetQuotesByUniverses 按标的池批量获取行情快照。
|
||||
func (c *Client) GetQuotesByUniverses(universes []string) ([]Quote, error) {
|
||||
body := map[string]any{"universes": universes}
|
||||
var quotes []Quote
|
||||
if err := c.request("POST", "/v1/quotes", nil, body, "es); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return quotes, nil
|
||||
}
|
||||
|
||||
// GetQuotesBySymbols 按代码列表批量获取行情快照。
|
||||
func (c *Client) GetQuotesBySymbols(symbols []string) ([]Quote, error) {
|
||||
body := map[string]any{"symbols": symbols}
|
||||
var quotes []Quote
|
||||
if err := c.request("POST", "/v1/quotes", nil, body, "es); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return quotes, nil
|
||||
}
|
||||
|
||||
// GetKlinesBatch 批量获取多只股票的历史日 K 线。
|
||||
// symbols 最多 100 只(tickflow 限制)。
|
||||
func (c *Client) GetKlinesBatch(symbols []string, period string, startMs, endMs int64) (map[string]KlineData, error) {
|
||||
if len(symbols) == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
if len(symbols) > 100 {
|
||||
return nil, fmt.Errorf("GetKlinesBatch supports up to 100 symbols, got %d", len(symbols))
|
||||
}
|
||||
params := map[string]any{
|
||||
"symbols": strings.Join(symbols, ","),
|
||||
"period": period,
|
||||
"adjust": "none",
|
||||
"start_time": startMs,
|
||||
"end_time": endMs,
|
||||
"count": 10000,
|
||||
}
|
||||
var result map[string]KlineData
|
||||
if err := c.request("GET", "/v1/klines/batch", params, nil, &result); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
for sym, data := range result {
|
||||
data.Symbol = sym
|
||||
result[sym] = data
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func (c *Client) request(method, path string, params, body map[string]any, result any) error {
|
||||
url := c.baseURL + path
|
||||
|
||||
var lastErr error
|
||||
for attempt := 0; attempt <= maxRetries; attempt++ {
|
||||
if attempt > 0 {
|
||||
time.Sleep(retryBaseDelay * time.Duration(1<<(attempt-1)))
|
||||
}
|
||||
|
||||
var bodyReader io.Reader
|
||||
if len(body) > 0 && (method == "POST" || method == "PUT") {
|
||||
data, err := json.Marshal(body)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
bodyReader = bytes.NewReader(data)
|
||||
}
|
||||
|
||||
req, err := http.NewRequest(method, url, bodyReader)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if c.apiKey != "" {
|
||||
req.Header.Set("x-api-key", c.apiKey)
|
||||
}
|
||||
req.Header.Set("Accept", "application/json")
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
|
||||
if len(params) > 0 {
|
||||
q := req.URL.Query()
|
||||
for k, v := range params {
|
||||
q.Set(k, fmt.Sprintf("%v", v))
|
||||
}
|
||||
req.URL.RawQuery = q.Encode()
|
||||
}
|
||||
|
||||
resp, err := c.client.Do(req)
|
||||
if err != nil {
|
||||
lastErr = err
|
||||
continue
|
||||
}
|
||||
|
||||
respBody, err := io.ReadAll(resp.Body)
|
||||
resp.Body.Close()
|
||||
if err != nil {
|
||||
lastErr = err
|
||||
continue
|
||||
}
|
||||
|
||||
if resp.StatusCode >= 500 || resp.StatusCode == 429 {
|
||||
lastErr = fmt.Errorf("tickflow %s %s returned %d: %s", method, path, resp.StatusCode, string(respBody))
|
||||
continue
|
||||
}
|
||||
|
||||
if resp.StatusCode >= 400 {
|
||||
return fmt.Errorf("tickflow %s %s returned %d: %s", method, path, resp.StatusCode, string(respBody))
|
||||
}
|
||||
|
||||
var wrapper struct {
|
||||
Data json.RawMessage `json:"data"`
|
||||
}
|
||||
if err := json.Unmarshal(respBody, &wrapper); err != nil {
|
||||
return fmt.Errorf("decode tickflow response: %w", err)
|
||||
}
|
||||
if result != nil {
|
||||
if err := json.Unmarshal(wrapper.Data, result); err != nil {
|
||||
return fmt.Errorf("decode tickflow data: %w", err)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
if lastErr != nil {
|
||||
return fmt.Errorf("tickflow request failed after %d retries: %w", maxRetries, lastErr)
|
||||
}
|
||||
return fmt.Errorf("tickflow request failed")
|
||||
}
|
||||
@@ -0,0 +1,594 @@
|
||||
package datasource
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"sort"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"stock-user-system/internal/config"
|
||||
)
|
||||
|
||||
const (
|
||||
tushareBaseURL = "http://api.tushare.pro"
|
||||
chinaTZ = "Asia/Shanghai"
|
||||
defaultTimeout = 60 * time.Second
|
||||
maxRetries = 3
|
||||
retryBaseDelay = 1 * time.Second
|
||||
)
|
||||
|
||||
var chinaLoc, _ = time.LoadLocation(chinaTZ)
|
||||
|
||||
// Quote 对应行情接口返回的单条行情数据。
|
||||
type Quote struct {
|
||||
Symbol string `json:"symbol"`
|
||||
Name string `json:"name"`
|
||||
Open float64 `json:"open"`
|
||||
High float64 `json:"high"`
|
||||
Low float64 `json:"low"`
|
||||
Close float64 `json:"last_price"`
|
||||
PrevClose float64 `json:"prev_close"`
|
||||
Volume int64 `json:"volume"`
|
||||
Amount float64 `json:"amount"`
|
||||
Timestamp int64 `json:"timestamp"`
|
||||
Region string `json:"region"`
|
||||
Ext struct {
|
||||
ChangePct float64 `json:"change_pct"`
|
||||
TurnoverRate float64 `json:"turnover_rate"`
|
||||
Name string `json:"name"`
|
||||
} `json:"ext"`
|
||||
}
|
||||
|
||||
// ChangePct 优先使用 ext.change_pct(统一返回小数形式)。
|
||||
func (q *Quote) ChangePct() float64 {
|
||||
if q.Ext.ChangePct != 0 {
|
||||
return q.Ext.ChangePct
|
||||
}
|
||||
if q.PrevClose != 0 {
|
||||
return (q.Close - q.PrevClose) / q.PrevClose
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
// TurnoverRate 返回 ext.turnover_rate(Tushare 当前未单独获取,可能为 0)。
|
||||
func (q *Quote) TurnoverRate() float64 {
|
||||
return q.Ext.TurnoverRate
|
||||
}
|
||||
|
||||
// DisplayName 优先使用 ext.name,否则使用 symbol。
|
||||
func (q *Quote) DisplayName() string {
|
||||
if q.Ext.Name != "" {
|
||||
return q.Ext.Name
|
||||
}
|
||||
if q.Name != "" {
|
||||
return q.Name
|
||||
}
|
||||
return q.Symbol
|
||||
}
|
||||
|
||||
// UniverseDetail 对应标的池详情。
|
||||
type UniverseDetail struct {
|
||||
ID string `json:"id"`
|
||||
Name string `json:"name"`
|
||||
SymbolCount int `json:"symbol_count"`
|
||||
Symbols []string `json:"symbols"`
|
||||
}
|
||||
|
||||
// KlineData 对应单标的历史日 K 线(紧凑列式)。
|
||||
type KlineData struct {
|
||||
Symbol string `json:"-"`
|
||||
Name string `json:"-"`
|
||||
Timestamp []int64 `json:"timestamp"`
|
||||
Open []float64 `json:"open"`
|
||||
High []float64 `json:"high"`
|
||||
Low []float64 `json:"low"`
|
||||
Close []float64 `json:"close"`
|
||||
Volume []int64 `json:"volume"`
|
||||
Amount []float64 `json:"amount"`
|
||||
}
|
||||
|
||||
// Tushare 股票代码后缀到 region 的映射。
|
||||
var regionBySuffix = map[string]string{
|
||||
".SH": "CN_SH",
|
||||
".SZ": "CN_SZ",
|
||||
".BJ": "CN_BJ",
|
||||
}
|
||||
|
||||
// knownIndexSymbols 已知的核心指数代码,历史 K 线同步时需要使用 index_daily 接口。
|
||||
var knownIndexSymbols = map[string]bool{
|
||||
"000001.SH": true,
|
||||
"399001.SZ": true,
|
||||
"399006.SZ": true,
|
||||
"000688.SH": true,
|
||||
}
|
||||
|
||||
// indexNames 已知核心指数的中文名称。
|
||||
var indexNames = map[string]string{
|
||||
"000001.SH": "上证指数",
|
||||
"399001.SZ": "深证成指",
|
||||
"399006.SZ": "创业板指",
|
||||
"000688.SH": "科创50",
|
||||
}
|
||||
|
||||
// Client 封装 Tushare Pro HTTP API 调用,接口语义与旧 tickflow 客户端保持一致,
|
||||
// 方便 stock_sync 服务直接切换数据源。
|
||||
type Client struct {
|
||||
token string
|
||||
baseURL string
|
||||
client *http.Client
|
||||
nameCache map[string]string
|
||||
cacheMu sync.RWMutex
|
||||
}
|
||||
|
||||
// NewClient 从配置创建 Tushare 客户端。
|
||||
func NewClient(cfg *config.Config) *Client {
|
||||
return &Client{
|
||||
token: cfg.TushareToken,
|
||||
baseURL: tushareBaseURL,
|
||||
client: &http.Client{Timeout: defaultTimeout},
|
||||
nameCache: make(map[string]string),
|
||||
}
|
||||
}
|
||||
|
||||
// GetUniverse 获取指定标的池的完整代码列表。Tushare 没有标的池概念,
|
||||
// 这里固定返回全部 A 股(list_status=L),忽略传入的 id。
|
||||
func (c *Client) GetUniverse(id string) (*UniverseDetail, error) {
|
||||
params := map[string]any{
|
||||
"list_status": "L",
|
||||
"fields": "ts_code,name",
|
||||
}
|
||||
fields, items, err := c.call("stock_basic", params)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("stock_basic: %w", err)
|
||||
}
|
||||
|
||||
col := buildColumnMap(fields)
|
||||
idxCode := col["ts_code"]
|
||||
idxName := col["name"]
|
||||
|
||||
symbols := make([]string, 0, len(items))
|
||||
c.cacheMu.Lock()
|
||||
for _, row := range items {
|
||||
code := stringAt(row, idxCode)
|
||||
name := stringAt(row, idxName)
|
||||
if code == "" {
|
||||
continue
|
||||
}
|
||||
symbols = append(symbols, code)
|
||||
if name != "" {
|
||||
c.nameCache[code] = name
|
||||
}
|
||||
}
|
||||
c.cacheMu.Unlock()
|
||||
|
||||
return &UniverseDetail{
|
||||
ID: id,
|
||||
Name: "A股全市场",
|
||||
SymbolCount: len(symbols),
|
||||
Symbols: symbols,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// GetQuotesByUniverses 按标的池批量获取行情快照。Tushare 不支持按标的池查询,
|
||||
// 这里返回最新交易日的全部 A 股日 K 数据。
|
||||
func (c *Client) GetQuotesByUniverses(universes []string) ([]Quote, error) {
|
||||
_ = universes
|
||||
// 名称缓存失败不影响行情同步,忽略错误
|
||||
_ = c.ensureNameCache()
|
||||
|
||||
latestDate, err := c.latestTradeDate()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("get latest trade date: %w", err)
|
||||
}
|
||||
|
||||
params := map[string]any{
|
||||
"trade_date": latestDate,
|
||||
}
|
||||
fields, items, err := c.call("daily", params)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("daily: %w", err)
|
||||
}
|
||||
|
||||
return c.rowsToQuotes(fields, items, latestDate, false)
|
||||
}
|
||||
|
||||
// GetQuotesBySymbols 按代码列表批量获取行情快照。用于获取核心指数日 K。
|
||||
func (c *Client) GetQuotesBySymbols(symbols []string) ([]Quote, error) {
|
||||
if len(symbols) == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
latestDate, err := c.latestTradeDate()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("get latest trade date: %w", err)
|
||||
}
|
||||
|
||||
params := map[string]any{
|
||||
"ts_code": strings.Join(symbols, ","),
|
||||
"trade_date": latestDate,
|
||||
}
|
||||
fields, items, err := c.call("index_daily", params)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("index_daily: %w", err)
|
||||
}
|
||||
|
||||
return c.rowsToQuotes(fields, items, latestDate, true)
|
||||
}
|
||||
|
||||
// GetKlinesBatch 批量获取多只股票的历史日 K 线。
|
||||
// Tushare 的 daily 接口按单代码查询历史数据效率更高,因此内部串行/并发调用。
|
||||
func (c *Client) GetKlinesBatch(symbols []string, period string, startMs, endMs int64) (map[string]KlineData, error) {
|
||||
if len(symbols) == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
if period != "" && period != "1d" {
|
||||
return nil, fmt.Errorf("tushare only supports 1d period, got %q", period)
|
||||
}
|
||||
|
||||
startDate := msToTushareDate(startMs)
|
||||
endDate := msToTushareDate(endMs)
|
||||
|
||||
result := make(map[string]KlineData, len(symbols))
|
||||
var mu sync.Mutex
|
||||
var wg sync.WaitGroup
|
||||
|
||||
// Tushare 免费/付费账号有频率限制,默认 4 并发,避免触发限流。
|
||||
sem := make(chan struct{}, 4)
|
||||
var errMu sync.Mutex
|
||||
var firstErr error
|
||||
|
||||
for _, sym := range symbols {
|
||||
wg.Add(1)
|
||||
sem <- struct{}{}
|
||||
go func(symbol string) {
|
||||
defer wg.Done()
|
||||
defer func() { <-sem }()
|
||||
|
||||
data, err := c.getKlines(symbol, startDate, endDate)
|
||||
if err != nil {
|
||||
errMu.Lock()
|
||||
if firstErr == nil {
|
||||
firstErr = err
|
||||
}
|
||||
errMu.Unlock()
|
||||
return
|
||||
}
|
||||
mu.Lock()
|
||||
result[symbol] = data
|
||||
mu.Unlock()
|
||||
}(sym)
|
||||
}
|
||||
wg.Wait()
|
||||
|
||||
if firstErr != nil {
|
||||
return result, firstErr
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func (c *Client) getKlines(symbol, startDate, endDate string) (KlineData, error) {
|
||||
apiName := "daily"
|
||||
if knownIndexSymbols[symbol] {
|
||||
apiName = "index_daily"
|
||||
}
|
||||
|
||||
params := map[string]any{
|
||||
"ts_code": symbol,
|
||||
"start_date": startDate,
|
||||
"end_date": endDate,
|
||||
}
|
||||
fields, items, err := c.call(apiName, params)
|
||||
if err != nil {
|
||||
return KlineData{}, err
|
||||
}
|
||||
|
||||
col := buildColumnMap(fields)
|
||||
data := KlineData{
|
||||
Symbol: symbol,
|
||||
Timestamp: make([]int64, 0, len(items)),
|
||||
Open: make([]float64, 0, len(items)),
|
||||
High: make([]float64, 0, len(items)),
|
||||
Low: make([]float64, 0, len(items)),
|
||||
Close: make([]float64, 0, len(items)),
|
||||
Volume: make([]int64, 0, len(items)),
|
||||
Amount: make([]float64, 0, len(items)),
|
||||
}
|
||||
|
||||
for _, row := range items {
|
||||
tradeDate := stringAt(row, col["trade_date"])
|
||||
ts := tushareDateToMs(tradeDate)
|
||||
if ts == 0 {
|
||||
continue
|
||||
}
|
||||
data.Timestamp = append(data.Timestamp, ts)
|
||||
data.Open = append(data.Open, floatAt(row, col["open"]))
|
||||
data.High = append(data.High, floatAt(row, col["high"]))
|
||||
data.Low = append(data.Low, floatAt(row, col["low"]))
|
||||
data.Close = append(data.Close, floatAt(row, col["close"]))
|
||||
// Tushare 成交量单位为"手",成交金额单位为"千元"
|
||||
data.Volume = append(data.Volume, int64(floatAt(row, col["vol"])*100))
|
||||
data.Amount = append(data.Amount, floatAt(row, col["amount"])*1000)
|
||||
}
|
||||
|
||||
// 按时间升序排列所有列
|
||||
indices := make([]int, len(data.Timestamp))
|
||||
for i := range indices {
|
||||
indices[i] = i
|
||||
}
|
||||
sort.SliceStable(indices, func(i, j int) bool {
|
||||
return data.Timestamp[indices[i]] < data.Timestamp[indices[j]]
|
||||
})
|
||||
|
||||
reorder := func(src []int64) []int64 {
|
||||
dst := make([]int64, len(src))
|
||||
for i, idx := range indices {
|
||||
dst[i] = src[idx]
|
||||
}
|
||||
return dst
|
||||
}
|
||||
reorderF := func(src []float64) []float64 {
|
||||
dst := make([]float64, len(src))
|
||||
for i, idx := range indices {
|
||||
dst[i] = src[idx]
|
||||
}
|
||||
return dst
|
||||
}
|
||||
|
||||
data.Timestamp = reorder(data.Timestamp)
|
||||
data.Open = reorderF(data.Open)
|
||||
data.High = reorderF(data.High)
|
||||
data.Low = reorderF(data.Low)
|
||||
data.Close = reorderF(data.Close)
|
||||
data.Volume = reorder(data.Volume)
|
||||
data.Amount = reorderF(data.Amount)
|
||||
|
||||
return data, nil
|
||||
}
|
||||
|
||||
// ensureNameCache 如果名称缓存为空,则通过 stock_basic 加载一次。
|
||||
func (c *Client) ensureNameCache() error {
|
||||
c.cacheMu.RLock()
|
||||
hasData := len(c.nameCache) > 0
|
||||
c.cacheMu.RUnlock()
|
||||
if hasData {
|
||||
return nil
|
||||
}
|
||||
|
||||
_, err := c.GetUniverse("CN_Equity_A")
|
||||
return err
|
||||
}
|
||||
|
||||
// displayName 返回股票名称;索引使用内置名称,股票从缓存读取,缺失则返回 symbol。
|
||||
func (c *Client) displayName(symbol string) string {
|
||||
if name, ok := indexNames[symbol]; ok {
|
||||
return name
|
||||
}
|
||||
c.cacheMu.RLock()
|
||||
name, ok := c.nameCache[symbol]
|
||||
c.cacheMu.RUnlock()
|
||||
if ok && name != "" {
|
||||
return name
|
||||
}
|
||||
return symbol
|
||||
}
|
||||
|
||||
func (c *Client) rowsToQuotes(fields []string, items [][]any, tradeDate string, isIndex bool) ([]Quote, error) {
|
||||
col := buildColumnMap(fields)
|
||||
quotes := make([]Quote, 0, len(items))
|
||||
|
||||
for _, row := range items {
|
||||
symbol := stringAt(row, col["ts_code"])
|
||||
if symbol == "" {
|
||||
continue
|
||||
}
|
||||
|
||||
vol := floatAt(row, col["vol"])
|
||||
amount := floatAt(row, col["amount"])
|
||||
changePct := floatAt(row, col["pct_chg"]) / 100
|
||||
|
||||
quote := Quote{
|
||||
Symbol: symbol,
|
||||
Name: c.displayName(symbol),
|
||||
Open: floatAt(row, col["open"]),
|
||||
High: floatAt(row, col["high"]),
|
||||
Low: floatAt(row, col["low"]),
|
||||
Close: floatAt(row, col["close"]),
|
||||
PrevClose: floatAt(row, col["pre_close"]),
|
||||
Volume: int64(vol * 100),
|
||||
Amount: amount * 1000,
|
||||
Timestamp: tushareDateToMs(tradeDate),
|
||||
Region: regionFromSymbol(symbol),
|
||||
Ext: struct {
|
||||
ChangePct float64 `json:"change_pct"`
|
||||
TurnoverRate float64 `json:"turnover_rate"`
|
||||
Name string `json:"name"`
|
||||
}{
|
||||
ChangePct: changePct,
|
||||
Name: c.displayName(symbol),
|
||||
},
|
||||
}
|
||||
quotes = append(quotes, quote)
|
||||
}
|
||||
|
||||
return quotes, nil
|
||||
}
|
||||
|
||||
// latestTradeDate 返回最近一个交易日的日期字符串(YYYYMMDD)。
|
||||
func (c *Client) latestTradeDate() (string, error) {
|
||||
today := time.Now().In(chinaLoc).Format("20060102")
|
||||
params := map[string]any{
|
||||
"exchange": "SSE",
|
||||
"start_date": "20200101",
|
||||
"end_date": today,
|
||||
"is_open": "1",
|
||||
}
|
||||
fields, items, err := c.call("trade_cal", params)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
col := buildColumnMap(fields)
|
||||
idxDate, ok := col["cal_date"]
|
||||
if !ok {
|
||||
return "", fmt.Errorf("trade_cal response missing cal_date field")
|
||||
}
|
||||
|
||||
latest := ""
|
||||
for _, row := range items {
|
||||
date := stringAt(row, idxDate)
|
||||
if date > latest {
|
||||
latest = date
|
||||
}
|
||||
}
|
||||
if latest == "" {
|
||||
return "", fmt.Errorf("no trade date found")
|
||||
}
|
||||
return latest, nil
|
||||
}
|
||||
|
||||
// call 调用 Tushare Pro API,返回字段名与数据行。
|
||||
func (c *Client) call(apiName string, params map[string]any) ([]string, [][]any, error) {
|
||||
reqBody := map[string]any{
|
||||
"api_name": apiName,
|
||||
"token": c.token,
|
||||
"params": params,
|
||||
"fields": "",
|
||||
}
|
||||
data, err := json.Marshal(reqBody)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
|
||||
var lastErr error
|
||||
for attempt := 0; attempt <= maxRetries; attempt++ {
|
||||
if attempt > 0 {
|
||||
time.Sleep(retryBaseDelay * time.Duration(1<<(attempt-1)))
|
||||
}
|
||||
|
||||
req, err := http.NewRequest("POST", c.baseURL, bytes.NewReader(data))
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("Accept", "application/json")
|
||||
|
||||
resp, err := c.client.Do(req)
|
||||
if err != nil {
|
||||
lastErr = err
|
||||
continue
|
||||
}
|
||||
|
||||
respBody, err := io.ReadAll(resp.Body)
|
||||
resp.Body.Close()
|
||||
if err != nil {
|
||||
lastErr = err
|
||||
continue
|
||||
}
|
||||
|
||||
if resp.StatusCode >= 500 || resp.StatusCode == 429 {
|
||||
lastErr = fmt.Errorf("tushare %s returned %d: %s", apiName, resp.StatusCode, string(respBody))
|
||||
continue
|
||||
}
|
||||
if resp.StatusCode >= 400 {
|
||||
return nil, nil, fmt.Errorf("tushare %s returned %d: %s", apiName, resp.StatusCode, string(respBody))
|
||||
}
|
||||
|
||||
var wrapper struct {
|
||||
Code int `json:"code"`
|
||||
Msg string `json:"msg"`
|
||||
Data *struct {
|
||||
Fields []string `json:"fields"`
|
||||
Items [][]any `json:"items"`
|
||||
} `json:"data"`
|
||||
}
|
||||
if err := json.Unmarshal(respBody, &wrapper); err != nil {
|
||||
return nil, nil, fmt.Errorf("decode tushare response: %w", err)
|
||||
}
|
||||
if wrapper.Code != 0 {
|
||||
return nil, nil, fmt.Errorf("tushare %s error %d: %s", apiName, wrapper.Code, wrapper.Msg)
|
||||
}
|
||||
if wrapper.Data == nil {
|
||||
return nil, nil, nil
|
||||
}
|
||||
return wrapper.Data.Fields, wrapper.Data.Items, nil
|
||||
}
|
||||
|
||||
if lastErr != nil {
|
||||
return nil, nil, fmt.Errorf("tushare %s failed after %d retries: %w", apiName, maxRetries, lastErr)
|
||||
}
|
||||
return nil, nil, fmt.Errorf("tushare %s request failed", apiName)
|
||||
}
|
||||
|
||||
func buildColumnMap(fields []string) map[string]int {
|
||||
m := make(map[string]int, len(fields))
|
||||
for i, f := range fields {
|
||||
m[f] = i
|
||||
}
|
||||
return m
|
||||
}
|
||||
|
||||
func stringAt(row []any, idx int) string {
|
||||
if idx < 0 || idx >= len(row) || row[idx] == nil {
|
||||
return ""
|
||||
}
|
||||
switch v := row[idx].(type) {
|
||||
case string:
|
||||
return v
|
||||
case []byte:
|
||||
return string(v)
|
||||
default:
|
||||
return fmt.Sprintf("%v", v)
|
||||
}
|
||||
}
|
||||
|
||||
func floatAt(row []any, idx int) float64 {
|
||||
if idx < 0 || idx >= len(row) || row[idx] == nil {
|
||||
return 0
|
||||
}
|
||||
switch v := row[idx].(type) {
|
||||
case float64:
|
||||
return v
|
||||
case float32:
|
||||
return float64(v)
|
||||
case int:
|
||||
return float64(v)
|
||||
case int64:
|
||||
return float64(v)
|
||||
case string:
|
||||
f, _ := strconv.ParseFloat(v, 64)
|
||||
return f
|
||||
default:
|
||||
f, _ := strconv.ParseFloat(fmt.Sprintf("%v", v), 64)
|
||||
return f
|
||||
}
|
||||
}
|
||||
|
||||
func regionFromSymbol(symbol string) string {
|
||||
for suffix, region := range regionBySuffix {
|
||||
if strings.HasSuffix(symbol, suffix) {
|
||||
return region
|
||||
}
|
||||
}
|
||||
return "CN"
|
||||
}
|
||||
|
||||
func msToTushareDate(ms int64) string {
|
||||
return time.UnixMilli(ms).In(chinaLoc).Format("20060102")
|
||||
}
|
||||
|
||||
func tushareDateToMs(date string) int64 {
|
||||
if len(date) != 8 {
|
||||
return 0
|
||||
}
|
||||
t, err := time.ParseInLocation("20060102", date, chinaLoc)
|
||||
if err != nil {
|
||||
return 0
|
||||
}
|
||||
return t.UnixMilli()
|
||||
}
|
||||
Reference in New Issue
Block a user