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:
Justin Gu
2026-06-12 03:26:29 +08:00
parent 9d7a9161f7
commit 0e74752701
6 changed files with 117 additions and 64 deletions
+41
View File
@@ -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
+14 -23
View File
@@ -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 -16
View File
@@ -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
+7 -10
View File
@@ -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 -15
View File
@@ -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)
+43
View File
@@ -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")