234 lines
5.7 KiB
Go
234 lines
5.7 KiB
Go
package datasource
|
|
|
|
import (
|
|
"bytes"
|
|
"encoding/json"
|
|
"fmt"
|
|
"io"
|
|
"net/http"
|
|
"strconv"
|
|
"time"
|
|
)
|
|
|
|
const (
|
|
tushareBaseURL = "http://api.tushare.pro"
|
|
tushareToken = "76efd8465f9f2591aa42a385268e06acf6b80b7a15be2267ad2281b7"
|
|
defaultTimeout = 60 * time.Second
|
|
maxRetries = 3
|
|
retryBaseDelay = 1 * time.Second
|
|
)
|
|
|
|
// StockBasic 表示股票基础信息。
|
|
type StockBasic struct {
|
|
TsCode string
|
|
Symbol string
|
|
Name string
|
|
Area string
|
|
Industry string
|
|
Fullname string
|
|
Enname string
|
|
Cnspell string
|
|
Market string
|
|
Exchange string
|
|
CurrType string
|
|
ListStatus string
|
|
ListDate string
|
|
DelistDate string
|
|
IsHs string
|
|
ActName string
|
|
ActEntType string
|
|
}
|
|
|
|
const stockBasicFields = "ts_code,symbol,name,area,industry,fullname,enname,cnspell,market,exchange,curr_type,list_status,list_date,delist_date,is_hs,act_name,act_ent_type"
|
|
|
|
// Client 封装 Tushare Pro HTTP API 调用。
|
|
type Client struct {
|
|
token string
|
|
baseURL string
|
|
client *http.Client
|
|
}
|
|
|
|
// NewClient 创建 Tushare 客户端。
|
|
func NewClient() *Client {
|
|
return &Client{
|
|
token: tushareToken,
|
|
baseURL: tushareBaseURL,
|
|
client: &http.Client{Timeout: defaultTimeout},
|
|
}
|
|
}
|
|
|
|
func parseStockBasic(row []any, col map[string]int) StockBasic {
|
|
return StockBasic{
|
|
TsCode: stringAt(row, col["ts_code"]),
|
|
Symbol: stringAt(row, col["symbol"]),
|
|
Name: stringAt(row, col["name"]),
|
|
Area: stringAt(row, col["area"]),
|
|
Industry: stringAt(row, col["industry"]),
|
|
Fullname: stringAt(row, col["fullname"]),
|
|
Enname: stringAt(row, col["enname"]),
|
|
Cnspell: stringAt(row, col["cnspell"]),
|
|
Market: stringAt(row, col["market"]),
|
|
Exchange: stringAt(row, col["exchange"]),
|
|
CurrType: stringAt(row, col["curr_type"]),
|
|
ListStatus: stringAt(row, col["list_status"]),
|
|
ListDate: stringAt(row, col["list_date"]),
|
|
DelistDate: stringAt(row, col["delist_date"]),
|
|
IsHs: stringAt(row, col["is_hs"]),
|
|
ActName: stringAt(row, col["act_name"]),
|
|
ActEntType: stringAt(row, col["act_ent_type"]),
|
|
}
|
|
}
|
|
|
|
// ListStocks 通过 stock_basic 接口获取全部上市股票基础信息。
|
|
func (c *Client) ListStocks() ([]StockBasic, error) {
|
|
params := map[string]any{
|
|
"list_status": "L",
|
|
"fields": stockBasicFields,
|
|
}
|
|
return c.listStockBasic(params)
|
|
}
|
|
|
|
// ListStocksByExchange 按交易所获取股票基础信息。
|
|
func (c *Client) ListStocksByExchange(exchange string) ([]StockBasic, error) {
|
|
params := map[string]any{
|
|
"exchange": exchange,
|
|
"fields": stockBasicFields,
|
|
}
|
|
return c.listStockBasic(params)
|
|
}
|
|
|
|
func (c *Client) listStockBasic(params map[string]any) ([]StockBasic, error) {
|
|
fields, items, err := c.call("stock_basic", params)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("stock_basic: %w", err)
|
|
}
|
|
|
|
col := buildColumnMap(fields)
|
|
stocks := make([]StockBasic, 0, len(items))
|
|
for _, row := range items {
|
|
code := stringAt(row, col["ts_code"])
|
|
if code == "" {
|
|
continue
|
|
}
|
|
stocks = append(stocks, parseStockBasic(row, col))
|
|
}
|
|
return stocks, 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
|
|
}
|
|
}
|