Files
easy_tdx_max/src/easy_tdx/web/routers/bars.py
T
GitHub 5a3ad15477 feat(web): /bars 迁移到 MacClient + 支持复权(issue #43)
Web 获取 K 线此前用 AsyncTdxClient.get_security_bars(标准协议不支持复权),
导致 REST API 无法取前复权/后复权数据。将 /bars(个股 K 线)迁移到
AsyncMacClient.get_stock_kline(MAC 协议,支持 NONE/QFQ/HFQ + QFQ 负价兜底),
保持旧输出契约不变(日线 date 列、分钟线 datetime 列、OHLC 顺序、无 float_shares),
新增 adjust 参数(默认 QFQ),MAC 主机不可用时自动回退标准 TdxClient。

⚠️ 半破坏性变更:/bars 默认复权从"不复权"改为 QFQ(前复权)。老调用方若需
不复权请显式传 ?adjust=NONE。输出 DataFrame 列名/顺序/字段与旧版完全一致
(_normalize_mac_df 规整),仅价格数值因复权变化。

改动:
- bars.py:/bars 改走 mac_client.get_stock_kline(adjust=...),MAC 不可用回退
  get_security_bars(无复权 + warning);新增 _normalize_mac_df 规整输出契约;
  新增 adjust 查询参数(默认 QFQ)。
- convert.py:period_times_from_category(KlineCategory→Period 映射,显式处理
  YEAR 9→YEARLY 11、SEASON→QUARTERLY)+ adjust_from_str。
- deps.py:get_mac_client_optional(未连接返回 None,供回退判断)。
- schemas.py:AdjustEnum(OpenAPI 文档用)。

测试:新增 7 个(period_times 映射全表 + 不可映射值、adjust 转换、_normalize_mac_df
日线/分钟线/空 df)。全套 996 passed;ruff format/check + mypy 改动文件零错误。

不在本次范围:/bars/index(指数 K 线,MAC 另一套接口)、/minute、/transaction*
(MacClient tick_chart 语义不同,暂不迁移)。
2026-08-05 16:37:10 +08:00

185 lines
7.1 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""K线 / 分时 / 逐笔成交路由。"""
from __future__ import annotations
import logging
from typing import Any
import pandas as pd
from fastapi import APIRouter, Depends, Query
from easy_tdx.models.enums import KlineCategory
from easy_tdx.web.convert import (
adjust_from_str,
category_from_str,
market_from_str,
market_value_from_str,
period_times_from_category,
)
from easy_tdx.web.deps import get_client, get_mac_client_optional
from easy_tdx.web.schemas import DataFrameResponse
_logger = logging.getLogger(__name__)
router = APIRouter(tags=["bars"])
# 规整后保持的列顺序(匹配旧 SecurityBar 输出契约)
_NORMAL_COLS = ["open", "close", "high", "low", "vol", "amount"]
def _df_resp(df: Any) -> DataFrameResponse:
return DataFrameResponse.from_dataframe(df)
def _normalize_mac_df(df: pd.DataFrame, daily_plus: bool) -> pd.DataFrame:
"""规整 MacClient.get_stock_kline 的输出以匹配旧 /bars 契约。
MacClient 返回 ``datetime`` 列(含时分秒)+ ``float_shares`` 列,OHLC 顺序为
open/high/low/close。旧 /barsSecurityBar 路径)日线返回 ``date`` 列(仅日期)、
分钟线返回 ``datetime`` 列,无 float_sharesOHLC 顺序为 open/close/high/low。
本函数做对齐,保证迁移后调用方输出契约不变。
Args:
df: MacClient 返回的 DataFrame(可能为空)。
daily_plus: True=日线及以上周期(datetime→date),False=分钟线(保留 datetime)。
"""
if df.empty:
return df
out = df.copy()
if "float_shares" in out.columns:
out = out.drop(columns=["float_shares"])
time_col = "date" if daily_plus else "datetime"
if "datetime" in out.columns:
if daily_plus:
# 截断为仅日期(00:00:00),与旧 _merge_bar_datetime 的 date 列语义一致
out["datetime"] = pd.to_datetime(out["datetime"]).dt.normalize()
out = out.rename(columns={"datetime": time_col})
# 重排列顺序:时间列在前,OHLC 顺序 open/close/high/low,再 vol/amount
cols = [c for c in [time_col, *_NORMAL_COLS] if c in out.columns]
# 兜底:保留未列出的列(理论上不应有),追加到末尾
cols += [c for c in out.columns if c not in cols]
return out[cols]
@router.get("/bars", response_model=DataFrameResponse)
async def security_bars(
market: str = Query(..., description="市场: SZ, SH, BJ"),
code: str = Query(..., min_length=6, max_length=6),
category: str = Query(
"DAY",
description="K线周期: MIN_1, MIN_5, MIN_15, MIN_30, MIN_60, DAY, WEEK, MONTH, YEAR",
),
start: int = Query(0, ge=0),
count: int = Query(800, ge=1, le=800),
bar_time: str = Query(
"start", description="时间戳: start=bar开始时间(默认) / end=bar结束时间(对齐Tushare)"
),
adjust: str = Query(
"QFQ", description="复权: NONE=不复权 / QFQ=前复权(默认) / HFQ=后复权(需 MAC 客户端)"
),
mac_client: Any = Depends(get_mac_client_optional),
client: Any = Depends(get_client),
) -> DataFrameResponse:
"""获取股票K线数据(MAC 协议,支持复权)。
优先走 AsyncMacClient.get_stock_kline(支持 NONE/QFQ/HFQ 复权 + QFQ 负价兜底);
MAC 主机未连接时自动回退 AsyncTdxClient.get_security_bars(无复权,adjust 参数忽略)。
输出契约与旧版一致:日线返回 ``date`` 列,分钟线返回 ``datetime`` 列。
"""
cat = category_from_str(category)
if mac_client is not None:
period, times = period_times_from_category(cat)
df = await mac_client.get_stock_kline(
market_value_from_str(market),
code,
period,
start,
count,
times,
adjust=adjust_from_str(adjust),
bar_time=bar_time,
)
# daily_plus:日线及以上周期(DAY=4 及以上)datetime→date
df = _normalize_mac_df(df, daily_plus=int(cat) >= int(KlineCategory.DAY))
else:
# MAC 不可用:回退标准 TdxClient(无复权),adjust 参数忽略
_logger.warning(
"/bars MAC 客户端未连接,回退标准 TdxClient(不支持复权,adjust=%s 被忽略)",
adjust,
)
df = await client.get_security_bars(
market_from_str(market), code, cat, start, count, bar_time=bar_time
)
return _df_resp(df)
@router.get("/bars/index", response_model=DataFrameResponse)
async def index_bars(
market: str = Query(..., description="市场: SZ, SH"),
code: str = Query(..., min_length=6, max_length=6),
category: str = Query("DAY", description="K线周期"),
start: int = Query(0, ge=0),
count: int = Query(800, ge=1, le=800),
bar_time: str = Query(
"start", description="时间戳: start=bar开始时间(默认) / end=bar结束时间(对齐Tushare)"
),
client: Any = Depends(get_client),
) -> DataFrameResponse:
"""获取指数K线数据。"""
df = await client.get_index_bars(
market_from_str(market), code, category_from_str(category), start, count, bar_time=bar_time
)
return _df_resp(df)
@router.get("/minute", response_model=DataFrameResponse)
async def minute_time(
market: str = Query(..., description="市场: SZ, SH"),
code: str = Query(..., min_length=6, max_length=6),
client: Any = Depends(get_client),
) -> DataFrameResponse:
"""获取今日分时数据。"""
df = await client.get_minute_time_data(market_from_str(market), code)
return _df_resp(df)
@router.get("/minute/history", response_model=DataFrameResponse)
async def history_minute_time(
market: str = Query(..., description="市场: SZ, SH"),
code: str = Query(..., min_length=6, max_length=6),
date: int = Query(..., description="日期 YYYYMMDD"),
client: Any = Depends(get_client),
) -> DataFrameResponse:
"""获取历史某日分时数据。"""
df = await client.get_history_minute_time_data(market_from_str(market), code, date)
return _df_resp(df)
@router.get("/transaction", response_model=DataFrameResponse)
async def transaction_data(
market: str = Query(..., description="市场: SZ, SH"),
code: str = Query(..., min_length=6, max_length=6),
start: int = Query(0, ge=0),
count: int = Query(800, ge=1, le=800),
client: Any = Depends(get_client),
) -> DataFrameResponse:
"""获取当日逐笔成交。"""
df = await client.get_transaction_data(market_from_str(market), code, start, count)
return _df_resp(df)
@router.get("/transaction/history", response_model=DataFrameResponse)
async def history_transaction_data(
market: str = Query(..., description="市场: SZ, SH"),
code: str = Query(..., min_length=6, max_length=6),
date: int = Query(..., description="日期 YYYYMMDD"),
start: int = Query(0, ge=0),
count: int = Query(800, ge=1, le=800),
client: Any = Depends(get_client),
) -> DataFrameResponse:
"""获取历史逐笔成交。"""
df = await client.get_history_transaction_data(
market_from_str(market), code, date, start, count
)
return _df_resp(df)