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 } }