diff --git a/src/easy_tdx/web/convert.py b/src/easy_tdx/web/convert.py new file mode 100644 index 0000000..da9e4de --- /dev/null +++ b/src/easy_tdx/web/convert.py @@ -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 diff --git a/src/easy_tdx/web/routers/bars.py b/src/easy_tdx/web/routers/bars.py index 5075b74..3f18839 100644 --- a/src/easy_tdx/web/routers/bars.py +++ b/src/easy_tdx/web/routers/bars.py @@ -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) diff --git a/src/easy_tdx/web/routers/chanlun.py b/src/easy_tdx/web/routers/chanlun.py index aa91723..09dd277 100644 --- a/src/easy_tdx/web/routers/chanlun.py +++ b/src/easy_tdx/web/routers/chanlun.py @@ -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 diff --git a/src/easy_tdx/web/routers/finance.py b/src/easy_tdx/web/routers/finance.py index 162e087..6ca9bc7 100644 --- a/src/easy_tdx/web/routers/finance.py +++ b/src/easy_tdx/web/routers/finance.py @@ -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} diff --git a/src/easy_tdx/web/routers/market.py b/src/easy_tdx/web/routers/market.py index 8d6e81c..6e11d4b 100644 --- a/src/easy_tdx/web/routers/market.py +++ b/src/easy_tdx/web/routers/market.py @@ -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) diff --git a/tests/unit/test_web_api.py b/tests/unit/test_web_api.py index c618351..02e1046 100644 --- a/tests/unit/test_web_api.py +++ b/tests/unit/test_web_api.py @@ -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")