mirror of
https://ghfast.top/https://github.com/aeroxw/easy-tdx.git
synced 2026-09-12 15:44:15 +08:00
fix(web): validate market/category input — support lowercase, reject invalid with 400
Root cause: _market_from_str/_market/_category in routers used bare MarketEnum[key]/Market[key] without .upper() or try/except, so lowercase or invalid values (sz, ZZZ) threw uncaught KeyError → 500. Fix: extract shared convert.py with market_from_str/category_from_str that do .upper() + ValueError on invalid input. All 4 routers updated. 4 regression tests added for case-insensitive and invalid input.
This commit is contained in:
@@ -0,0 +1,41 @@
|
||||
"""共享参数转换工具(market/category 字符串 → 枚举)。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from easy_tdx.web.schemas import KlineCategoryEnum, MarketEnum
|
||||
|
||||
|
||||
def market_from_str(s: str) -> Any:
|
||||
"""将字符串转为 Market 枚举,支持大小写,非法值抛 ValueError。
|
||||
|
||||
>>> market_from_str("SZ") # 正常
|
||||
>>> market_from_str("sz") # 也正常(自动转大写)
|
||||
>>> market_from_str("ZZZ") # ValueError
|
||||
"""
|
||||
from easy_tdx.models.enums import Market
|
||||
|
||||
key = s.upper()
|
||||
try:
|
||||
return Market[MarketEnum[key].name]
|
||||
except KeyError:
|
||||
valid = ", ".join(m.name for m in MarketEnum)
|
||||
raise ValueError(f"无效市场代码 '{s}',可选值: {valid}") from None
|
||||
|
||||
|
||||
def category_from_str(s: str) -> Any:
|
||||
"""将字符串转为 KlineCategory 枚举,支持大小写和数字字符串。"""
|
||||
from easy_tdx.models.enums import KlineCategory
|
||||
|
||||
key = s.upper()
|
||||
# 支持纯数字(如 "4" 表示日线)
|
||||
try:
|
||||
return KlineCategory(int(key))
|
||||
except (ValueError, TypeError):
|
||||
pass
|
||||
try:
|
||||
return KlineCategory[KlineCategoryEnum[key].name]
|
||||
except KeyError:
|
||||
valid = ", ".join(c.name for c in KlineCategoryEnum)
|
||||
raise ValueError(f"无效K线周期 '{s}',可选值: {valid}") from None
|
||||
@@ -6,28 +6,13 @@ from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, Query
|
||||
|
||||
from easy_tdx.web.convert import category_from_str, market_from_str
|
||||
from easy_tdx.web.deps import get_client
|
||||
from easy_tdx.web.schemas import DataFrameResponse, KlineCategoryEnum
|
||||
from easy_tdx.web.schemas import DataFrameResponse
|
||||
|
||||
router = APIRouter(tags=["bars"])
|
||||
|
||||
|
||||
def _market(market: str) -> Any:
|
||||
from easy_tdx.models.enums import Market
|
||||
|
||||
return Market[market.upper()]
|
||||
|
||||
|
||||
def _category(category: str) -> Any:
|
||||
from easy_tdx.models.enums import KlineCategory
|
||||
|
||||
# Support both int and string name
|
||||
try:
|
||||
return KlineCategory(int(category))
|
||||
except (ValueError, TypeError):
|
||||
return KlineCategory[KlineCategoryEnum[category.upper()].name]
|
||||
|
||||
|
||||
def _df_resp(df: Any) -> DataFrameResponse:
|
||||
return DataFrameResponse.from_dataframe(df)
|
||||
|
||||
@@ -45,7 +30,9 @@ async def security_bars(
|
||||
client: Any = Depends(get_client),
|
||||
) -> DataFrameResponse:
|
||||
"""获取股票K线数据。"""
|
||||
df = await client.get_security_bars(_market(market), code, _category(category), start, count)
|
||||
df = await client.get_security_bars(
|
||||
market_from_str(market), code, category_from_str(category), start, count
|
||||
)
|
||||
return _df_resp(df)
|
||||
|
||||
|
||||
@@ -59,7 +46,9 @@ async def index_bars(
|
||||
client: Any = Depends(get_client),
|
||||
) -> DataFrameResponse:
|
||||
"""获取指数K线数据。"""
|
||||
df = await client.get_index_bars(_market(market), code, _category(category), start, count)
|
||||
df = await client.get_index_bars(
|
||||
market_from_str(market), code, category_from_str(category), start, count
|
||||
)
|
||||
return _df_resp(df)
|
||||
|
||||
|
||||
@@ -70,7 +59,7 @@ async def minute_time(
|
||||
client: Any = Depends(get_client),
|
||||
) -> DataFrameResponse:
|
||||
"""获取今日分时数据。"""
|
||||
df = await client.get_minute_time_data(_market(market), code)
|
||||
df = await client.get_minute_time_data(market_from_str(market), code)
|
||||
return _df_resp(df)
|
||||
|
||||
|
||||
@@ -82,7 +71,7 @@ async def history_minute_time(
|
||||
client: Any = Depends(get_client),
|
||||
) -> DataFrameResponse:
|
||||
"""获取历史某日分时数据。"""
|
||||
df = await client.get_history_minute_time_data(_market(market), code, date)
|
||||
df = await client.get_history_minute_time_data(market_from_str(market), code, date)
|
||||
return _df_resp(df)
|
||||
|
||||
|
||||
@@ -95,7 +84,7 @@ async def transaction_data(
|
||||
client: Any = Depends(get_client),
|
||||
) -> DataFrameResponse:
|
||||
"""获取当日逐笔成交。"""
|
||||
df = await client.get_transaction_data(_market(market), code, start, count)
|
||||
df = await client.get_transaction_data(market_from_str(market), code, start, count)
|
||||
return _df_resp(df)
|
||||
|
||||
|
||||
@@ -109,5 +98,7 @@ async def history_transaction_data(
|
||||
client: Any = Depends(get_client),
|
||||
) -> DataFrameResponse:
|
||||
"""获取历史逐笔成交。"""
|
||||
df = await client.get_history_transaction_data(_market(market), code, date, start, count)
|
||||
df = await client.get_history_transaction_data(
|
||||
market_from_str(market), code, date, start, count
|
||||
)
|
||||
return _df_resp(df)
|
||||
|
||||
@@ -6,27 +6,13 @@ from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends
|
||||
|
||||
from easy_tdx.web.convert import category_from_str, market_from_str
|
||||
from easy_tdx.web.deps import get_client
|
||||
from easy_tdx.web.schemas import ChanlunRequest
|
||||
|
||||
router = APIRouter(tags=["chanlun"])
|
||||
|
||||
|
||||
def _market(market: str) -> Any:
|
||||
from easy_tdx.models.enums import Market
|
||||
|
||||
return Market[market.upper()]
|
||||
|
||||
|
||||
def _category(category: str) -> Any:
|
||||
from easy_tdx.models.enums import KlineCategory
|
||||
|
||||
try:
|
||||
return KlineCategory(int(category))
|
||||
except (ValueError, TypeError):
|
||||
return KlineCategory[category.upper()]
|
||||
|
||||
|
||||
@router.post("/chanlun/analyze")
|
||||
async def chanlun_analyze(
|
||||
req: ChanlunRequest,
|
||||
@@ -41,7 +27,11 @@ async def chanlun_analyze(
|
||||
|
||||
# 1. Fetch kline data
|
||||
df = await client.get_security_bars(
|
||||
_market(req.market), req.code, _category(req.category), req.start, req.count
|
||||
market_from_str(req.market),
|
||||
req.code,
|
||||
category_from_str(req.category),
|
||||
req.start,
|
||||
req.count,
|
||||
)
|
||||
|
||||
# 2. Run chanlun analysis
|
||||
|
||||
@@ -6,18 +6,13 @@ from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, Query
|
||||
|
||||
from easy_tdx.web.convert import market_from_str
|
||||
from easy_tdx.web.deps import get_client
|
||||
from easy_tdx.web.schemas import DataFrameResponse
|
||||
|
||||
router = APIRouter(tags=["finance"])
|
||||
|
||||
|
||||
def _market(market: str) -> Any:
|
||||
from easy_tdx.models.enums import Market
|
||||
|
||||
return Market[market.upper()]
|
||||
|
||||
|
||||
def _df_resp(df: Any) -> DataFrameResponse:
|
||||
return DataFrameResponse.from_dataframe(df)
|
||||
|
||||
@@ -29,7 +24,7 @@ async def xdxr_info(
|
||||
client: Any = Depends(get_client),
|
||||
) -> DataFrameResponse:
|
||||
"""获取除权除息历史记录。"""
|
||||
df = await client.get_xdxr_info(_market(market), code)
|
||||
df = await client.get_xdxr_info(market_from_str(market), code)
|
||||
return _df_resp(df)
|
||||
|
||||
|
||||
@@ -40,7 +35,7 @@ async def finance_info(
|
||||
client: Any = Depends(get_client),
|
||||
) -> DataFrameResponse:
|
||||
"""获取最新财务数据。"""
|
||||
df = await client.get_finance_info(_market(market), code)
|
||||
df = await client.get_finance_info(market_from_str(market), code)
|
||||
return _df_resp(df)
|
||||
|
||||
|
||||
@@ -51,7 +46,7 @@ async def company_info_category(
|
||||
client: Any = Depends(get_client),
|
||||
) -> DataFrameResponse:
|
||||
"""获取公司信息文件目录。"""
|
||||
df = await client.get_company_info_category(_market(market), code)
|
||||
df = await client.get_company_info_category(market_from_str(market), code)
|
||||
return _df_resp(df)
|
||||
|
||||
|
||||
@@ -65,7 +60,9 @@ async def company_info_content(
|
||||
client: Any = Depends(get_client),
|
||||
) -> dict[str, str]:
|
||||
"""读取公司信息文本。"""
|
||||
content = await client.get_company_info_content(_market(market), code, filename, offset, length)
|
||||
content = await client.get_company_info_content(
|
||||
market_from_str(market), code, filename, offset, length
|
||||
)
|
||||
return {"content": content}
|
||||
|
||||
|
||||
|
||||
@@ -6,24 +6,17 @@ from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, Query
|
||||
|
||||
from easy_tdx.web.convert import market_from_str
|
||||
from easy_tdx.web.deps import get_client
|
||||
from easy_tdx.web.schemas import (
|
||||
CountResponse,
|
||||
DataFrameResponse,
|
||||
MarketEnum,
|
||||
QuoteRequest,
|
||||
)
|
||||
|
||||
router = APIRouter(tags=["market"])
|
||||
|
||||
|
||||
def _market_from_str(s: str) -> Any:
|
||||
"""将字符串转为 Market 枚举。"""
|
||||
from easy_tdx.models.enums import Market
|
||||
|
||||
return Market[MarketEnum[s].name]
|
||||
|
||||
|
||||
def _df_response(df: Any) -> DataFrameResponse:
|
||||
"""将 DataFrame 转为 API 响应。"""
|
||||
return DataFrameResponse.from_dataframe(df)
|
||||
@@ -35,7 +28,7 @@ async def security_count(
|
||||
client: Any = Depends(get_client),
|
||||
) -> CountResponse:
|
||||
"""获取市场证券总数。"""
|
||||
count = await client.get_security_count(_market_from_str(market))
|
||||
count = await client.get_security_count(market_from_str(market))
|
||||
return CountResponse(count=count)
|
||||
|
||||
|
||||
@@ -46,7 +39,7 @@ async def security_list(
|
||||
client: Any = Depends(get_client),
|
||||
) -> DataFrameResponse:
|
||||
"""获取证券列表(每页约1000条)。"""
|
||||
df = await client.get_security_list(_market_from_str(market), start)
|
||||
df = await client.get_security_list(market_from_str(market), start)
|
||||
return _df_response(df)
|
||||
|
||||
|
||||
@@ -66,11 +59,9 @@ async def security_quotes(
|
||||
client: Any = Depends(get_client),
|
||||
) -> DataFrameResponse:
|
||||
"""批量获取实时五档行情(最多80只/次)。"""
|
||||
from easy_tdx.models.enums import Market
|
||||
|
||||
stocks_parsed: list[tuple[Any, str]] = []
|
||||
for s in req.stocks:
|
||||
m = Market[MarketEnum[s.market].name]
|
||||
m = market_from_str(s.market)
|
||||
stocks_parsed.append((m, s.code))
|
||||
df = await client.get_security_quotes(stocks_parsed)
|
||||
return _df_response(df)
|
||||
@@ -92,7 +83,7 @@ async def fund_flow(
|
||||
client: Any = Depends(get_client),
|
||||
) -> DataFrameResponse:
|
||||
"""获取个股当日资金流向。"""
|
||||
df = await client.get_fund_flow(_market_from_str(market), code)
|
||||
df = await client.get_fund_flow(market_from_str(market), code)
|
||||
return _df_response(df)
|
||||
|
||||
|
||||
@@ -105,5 +96,5 @@ async def history_fund_flow(
|
||||
client: Any = Depends(get_client),
|
||||
) -> DataFrameResponse:
|
||||
"""获取个股历史日线资金流向。"""
|
||||
df = await client.get_history_fund_flow(_market_from_str(market), code, start, count)
|
||||
df = await client.get_history_fund_flow(market_from_str(market), code, start, count)
|
||||
return _df_response(df)
|
||||
|
||||
@@ -212,6 +212,49 @@ def test_serve_command_exists():
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Regression: input validation (case-insensitive + invalid → ValueError → 400)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_convert_market_lowercase():
|
||||
"""market_from_str should accept lowercase input."""
|
||||
pytest.importorskip("fastapi")
|
||||
from easy_tdx.models.enums import Market
|
||||
from easy_tdx.web.convert import market_from_str
|
||||
|
||||
assert market_from_str("sz") == Market.SZ
|
||||
assert market_from_str("sh") == Market.SH
|
||||
assert market_from_str("Bj") == Market.BJ
|
||||
|
||||
|
||||
def test_convert_market_invalid_raises_valueerror():
|
||||
"""market_from_str should raise ValueError for invalid market codes."""
|
||||
pytest.importorskip("fastapi")
|
||||
from easy_tdx.web.convert import market_from_str
|
||||
|
||||
with pytest.raises(ValueError, match="无效市场代码"):
|
||||
market_from_str("ZZZ")
|
||||
|
||||
|
||||
def test_convert_category_from_int_string():
|
||||
"""category_from_str should accept numeric string like '4'."""
|
||||
pytest.importorskip("fastapi")
|
||||
from easy_tdx.models.enums import KlineCategory
|
||||
from easy_tdx.web.convert import category_from_str
|
||||
|
||||
assert category_from_str("4") == KlineCategory.DAY
|
||||
|
||||
|
||||
def test_convert_category_invalid_raises_valueerror():
|
||||
"""category_from_str should raise ValueError for invalid period."""
|
||||
pytest.importorskip("fastapi")
|
||||
from easy_tdx.web.convert import category_from_str
|
||||
|
||||
with pytest.raises(ValueError, match="无效K线周期"):
|
||||
category_from_str("INVALID_PERIOD")
|
||||
|
||||
|
||||
def test_full_app_routes_registered():
|
||||
"""All routers should be mounted and accessible."""
|
||||
pytest.importorskip("fastapi")
|
||||
|
||||
Reference in New Issue
Block a user