看股功能更新

This commit is contained in:
2026-08-16 00:05:26 +08:00
parent 9cce670b74
commit fc86fe0674
28 changed files with 3823 additions and 96 deletions

View File

@@ -15,11 +15,13 @@ import json
from datetime import datetime
import pandas as pd
from fastapi import APIRouter, Depends, HTTPException
from sqlalchemy import select, text
from fastapi import APIRouter, Depends, File, HTTPException, UploadFile
from sqlalchemy import delete, select, text
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.sql.elements import TextClause
from .backtest.engine import BacktestConfig, run_backtest
from . import cache
from .auth import require_user
from .backtest.events import EventEngineError, run_event_backtest
from .backtest.strategies import build_strategy
@@ -30,6 +32,7 @@ from .data.symbols import plain_code
from .db import get_session
from .domain import Bar
from . import indicators as ind
from .trades import parse_statement
from .models import (
AdjFactor,
BacktestRun,
@@ -38,6 +41,7 @@ from .models import (
ScreenerQuery,
StockBasic,
UserPreference,
UserTrade,
WatchlistItem,
)
from .schemas import (
@@ -66,6 +70,9 @@ from .schemas import (
FacetItemOut,
SyncRequest,
SyncResponse,
TradesClearResponse,
TradesImportResponse,
UserTradeOut,
WatchlistOp,
)
from .screener import engine, market_sync
@@ -163,37 +170,63 @@ async def sync_data(req: SyncRequest, session: AsyncSession = Depends(get_sessio
# ---------- 股票列表(全市场浏览) ----------
_STOCKS_SQL = text("""
SELECT sb.ts_code, sb.symbol, sb.name, sb.industry, sb.market,
c.close AS close, p.close AS prev_close, c.ts AS last_ts, cnt.n AS bar_count,
CASE WHEN c.close IS NOT NULL AND p.close IS NOT NULL AND p.close <> 0
THEN round(((c.close / p.close - 1) * 100)::numeric, 2) END AS pct_chg,
(w.id IS NOT NULL) AS watched
FROM stock_basic sb
# 过滤/排序/分页在 stock_basic+watchlist+daily_snapshot 上完成(快照按最新交易日
# 走唯一索引 join便宜再对「本页」≤limit 只股票补最新价/昨收LATERAL 扫
# candles——旧写法对全市场 ~5000 只逐个算,每页都白算 50 倍的行情量。
# 排序列白名单键→CTE 内表达式order_by 由白名单拼接进模板,不接收用户原文。
_STOCKS_SORTS = {
"symbol": "sb.symbol",
"total_mv": "snap.total_mv",
"circ_mv": "snap.circ_mv",
"pe_ttm": "snap.pe_ttm",
"pb": "snap.pb",
"turnover_rate": "snap.turnover_rate",
}
_STOCKS_SQL_TMPL = """
WITH page AS (
SELECT sb.ts_code, sb.symbol, sb.name, sb.industry, sb.market,
(w.id IS NOT NULL) AS watched,
snap.turnover_rate, snap.pe_ttm, snap.pb, snap.total_mv, snap.circ_mv
FROM stock_basic sb
LEFT JOIN watchlist_items w ON w.ts_code = sb.ts_code AND w.user_id = :uid
LEFT JOIN daily_snapshot snap ON snap.ts_code = sb.ts_code
AND snap.trade_date = (SELECT max(trade_date) FROM daily_snapshot)
WHERE sb.list_status = 'L'
AND (:search = '' OR sb.symbol LIKE :psearch OR sb.name LIKE :psearch)
AND (:market = '' OR sb.market = :market)
AND (:industry = '' OR sb.industry = :industry)
AND (:area = '' OR sb.area = :area)
AND (:watched_only = false OR w.id IS NOT NULL)
ORDER BY {order_by}
LIMIT :limit OFFSET :offset
)
SELECT p.ts_code, p.symbol, p.name, p.industry, p.market, p.watched,
c.close AS close, prev.close AS prev_close, c.ts AS last_ts,
CASE WHEN c.close IS NOT NULL AND prev.close IS NOT NULL AND prev.close <> 0
THEN round(((c.close / prev.close - 1) * 100)::numeric, 2) END AS pct_chg,
p.turnover_rate, p.pe_ttm, p.pb,
round((p.total_mv / 10000.0)::numeric, 2) AS total_mv,
round((p.circ_mv / 10000.0)::numeric, 2) AS circ_mv
FROM page p
LEFT JOIN LATERAL (
SELECT close, ts FROM candles
WHERE symbol = sb.symbol AND timeframe = '1d'
WHERE symbol = p.symbol AND timeframe = '1d'
ORDER BY ts DESC LIMIT 1
) c ON true
LEFT JOIN LATERAL (
SELECT close FROM candles
WHERE symbol = sb.symbol AND timeframe = '1d' AND ts < c.ts
WHERE symbol = p.symbol AND timeframe = '1d' AND ts < c.ts
ORDER BY ts DESC LIMIT 1
) p ON c.ts IS NOT NULL
LEFT JOIN LATERAL (
SELECT count(*) AS n FROM candles
WHERE symbol = sb.symbol AND timeframe = '1d'
) cnt ON true
LEFT JOIN watchlist_items w ON w.ts_code = sb.ts_code AND w.user_id = :uid
WHERE sb.list_status = 'L'
AND (:search = '' OR sb.symbol LIKE :psearch OR sb.name LIKE :psearch)
AND (:market = '' OR sb.market = :market)
AND (:industry = '' OR sb.industry = :industry)
AND (:area = '' OR sb.area = :area)
AND (:watched_only = false OR w.id IS NOT NULL)
ORDER BY w.id DESC NULLS LAST, sb.symbol
LIMIT :limit OFFSET :offset
""")
) prev ON c.ts IS NOT NULL
"""
def _stocks_sql(sort: str, order: str) -> TextClause:
col = _STOCKS_SORTS.get(sort, _STOCKS_SORTS["symbol"])
direction = "DESC" if order == "desc" else "ASC"
nulls = " NULLS LAST" if col != "sb.symbol" else "" # 快照缺失/亏损无 PE 的排最后
return text(_STOCKS_SQL_TMPL.format(order_by=f"{col} {direction}{nulls}"))
_STOCKS_COUNT_SQL = text("""
SELECT count(*) FROM stock_basic sb
@@ -214,16 +247,32 @@ async def list_stocks(
industry: str = "",
area: str = "",
watched_only: bool = False,
sort: str = "symbol",
order: str = "asc",
limit: int = 100,
offset: int = 0,
session: AsyncSession = Depends(get_session),
user=Depends(require_user),
) -> StockListResponse:
"""全市场股票列表stock_basic 基本信息 + candles 最新行情(本地缓存,无缓存则行情列为空)。
自选股watchlist_items排最前watched_only=true 只看自选。"""
"""全市场股票列表stock_basic 基本信息 + candles 最新行情 + daily_snapshot 估值指标
(换手率/PE-TTM/PB/市值,无快照则这些列为空)。
watched_only=true 只看自选(自选有独立的「自选」分类入口,列表不再把自选排最前)。
sort ∈ {symbol,total_mv,circ_mv,pe_ttm,pb,turnover_rate}(白名单,其他值回落 symbol
order ∈ asc/desc快照列排序时缺失值无快照/亏损无 PE恒排末尾。
Redis 缓存:按「用户自选版本 + 查询参数(含排序)」缓存整页(含 total自选增删即时失效。"""
search = search.strip()
sort = sort if sort in _STOCKS_SORTS else "symbol"
order = "desc" if order.lower() == "desc" else "asc"
limit = max(1, min(limit, 500))
offset = max(0, offset)
key = (
f"stocks:u{user.id}"
f":v{await cache.get_version(f'watchlist:{user.id}')}"
f":{cache.digest(search, market, industry, area, watched_only, sort, order, limit, offset)}"
)
cached = await cache.cache_get(key)
if cached is not None:
return StockListResponse(**cached)
params = {
"search": search,
"psearch": f"%{search}%",
@@ -236,13 +285,18 @@ async def list_stocks(
"offset": offset,
}
total = (await session.execute(_STOCKS_COUNT_SQL, params)).scalar_one()
rows = (await session.execute(_STOCKS_SQL, params)).mappings().all()
return StockListResponse(total=total, items=[StockListItemOut(**r) for r in rows])
rows = (await session.execute(_stocks_sql(sort, order), params)).mappings().all()
resp = StockListResponse(total=total, items=[StockListItemOut(**r) for r in rows])
await cache.cache_set(key, resp.model_dump(mode="json"), settings.stocks_cache_ttl)
return resp
@router.get("/stocks/facets", response_model=StockFacetsResponse)
async def stock_facets(session: AsyncSession = Depends(get_session)) -> StockFacetsResponse:
"""看股页筛选项:行业 / 地域(含数量,按数量降序)。"""
"""看股页筛选项:行业 / 地域(含数量,按数量降序)。stock_basic 很少变,长缓存。"""
cached = await cache.cache_get("facets:stocks")
if cached is not None:
return StockFacetsResponse(**cached)
industries = (
await session.execute(text("""
SELECT industry AS name, count(*) AS n FROM stock_basic
@@ -257,10 +311,12 @@ async def stock_facets(session: AsyncSession = Depends(get_session)) -> StockFac
GROUP BY area ORDER BY n DESC
"""))
).mappings().all()
return StockFacetsResponse(
resp = StockFacetsResponse(
industries=[FacetItemOut(name=r["name"], count=r["n"]) for r in industries],
areas=[FacetItemOut(name=r["name"], count=r["n"]) for r in areas],
)
await cache.cache_set("facets:stocks", resp.model_dump(mode="json"), settings.facets_cache_ttl)
return resp
@router.post("/backtest", response_model=BacktestResponse)
@@ -551,6 +607,7 @@ async def add_watchlist(
if exists is None:
session.add(WatchlistItem(user_id=user.id, ts_code=req.ts_code))
await session.commit()
await cache.bump_version(f"watchlist:{user.id}") # 作废该用户的股票列表缓存
return await get_watchlist(session=session, user=user)
@@ -565,9 +622,125 @@ async def remove_watchlist(
{"u": user.id, "c": ts_code},
)
await session.commit()
await cache.bump_version(f"watchlist:{user.id}") # 作废该用户的股票列表缓存
return await get_watchlist(session=session, user=user)
# ---------- 交割单(个人实盘买卖点) ----------
def _trade_out(r: UserTrade) -> UserTradeOut:
return UserTradeOut(
id=r.id, ts_code=r.ts_code, name=r.name, trade_date=r.trade_date,
direction=r.direction, price=r.price, qty=r.qty, amount=r.amount, fee=r.fee,
)
@router.get("/trades", response_model=list[UserTradeOut])
async def list_trades(
ts_code: str | None = None,
session: AsyncSession = Depends(get_session),
user=Depends(require_user),
) -> list[UserTradeOut]:
"""当前用户导入的实盘成交(可选 ts_code 过滤按日期升序K线买卖点数据源"""
q = (
select(UserTrade)
.where(UserTrade.user_id == user.id)
.order_by(UserTrade.trade_date, UserTrade.id)
)
if ts_code:
q = q.where(UserTrade.ts_code == ts_code)
rows = (await session.execute(q)).scalars().all()
return [_trade_out(r) for r in rows]
@router.post("/trades/import", response_model=TradesImportResponse)
async def import_trades(
file: UploadFile = File(...),
session: AsyncSession = Depends(get_session),
user=Depends(require_user),
) -> TradesImportResponse:
"""上传券商交割单CSV/Excel/HTML 表格均可,自动识别列名),解析出买卖成交入库。
同一笔成交(同日同股同向同价同量)重复上传会跳过,重复导出幂等。
"""
data = await file.read()
if not data:
raise HTTPException(status_code=422, detail="文件是空的")
if len(data) > 20 * 1024 * 1024:
raise HTTPException(status_code=413, detail="文件超过 20MB请分时间段导出")
parsed = parse_statement(data, file.filename or "")
# 无证券代码列的导出(招商式):按证券名称反查 stock_basic 补 ts_code同名多码或查不到则弃行
unnamed = {t.name for t in parsed.trades if not t.ts_code and t.name}
if unnamed:
name_map: dict[str, str] = {}
for ts_code, name in (await session.execute(
select(StockBasic.ts_code, StockBasic.name).where(StockBasic.name.in_(unnamed))
)).all():
name_map[name] = "" if name in name_map else ts_code
for t in parsed.trades:
if not t.ts_code and t.name:
tc = name_map.get(t.name, "")
if tc:
t.ts_code, t.code = tc, tc.split(".")[0]
else:
parsed.skipped_bad.append(f"{t.trade_date} {t.name} 名称无法唯一对应代码,未入库")
def _key(t) -> tuple:
return (t.trade_date, t.ts_code, t.direction, None if t.price is None else round(t.price, 4), round(t.qty, 4))
# Python 侧去重兜底(唯一约束对 NULL price 不生效)
existing = {
(r.trade_date, r.ts_code, r.direction, None if r.price is None else round(r.price, 4), round(r.qty, 4))
for r in (
await session.execute(
select(UserTrade.trade_date, UserTrade.ts_code, UserTrade.direction, UserTrade.price, UserTrade.qty)
.where(UserTrade.user_id == user.id, UserTrade.ts_code.in_({t.ts_code for t in parsed.trades}))
)
).all()
}
inserted: list[UserTrade] = []
seen: set[tuple] = set()
skipped_dup = 0
for t in parsed.trades:
if not t.ts_code:
continue # 名称反查失败的行,已在 bad 里说明
k = _key(t)
if k in existing or k in seen:
skipped_dup += 1
continue
seen.add(k)
inserted.append(UserTrade(
user_id=user.id, ts_code=t.ts_code, code=t.code, name=t.name or None,
trade_date=t.trade_date, direction=t.direction, price=t.price,
qty=t.qty, amount=t.amount, fee=t.fee,
raw_json=json.dumps(t.raw, ensure_ascii=False, default=str),
))
if inserted:
session.add_all(inserted)
await session.commit()
return TradesImportResponse(
inserted=len(inserted),
skipped_dup=skipped_dup,
skipped_other=parsed.skipped_other,
stocks=len({t.ts_code for t in parsed.trades}),
bad=parsed.skipped_bad[:5],
sample=[_trade_out(r) for r in inserted[:5]],
)
@router.delete("/trades", response_model=TradesClearResponse)
async def clear_trades(
session: AsyncSession = Depends(get_session),
user=Depends(require_user),
) -> TradesClearResponse:
"""清空当前用户导入的全部成交(重新导入前用)。"""
res = await session.execute(delete(UserTrade).where(UserTrade.user_id == user.id))
await session.commit()
return TradesClearResponse(deleted=res.rowcount or 0)
@router.post("/screener/sync", response_model=ScreenerSyncStatus)
async def screener_sync_start(
req: ScreenerSyncRequest, session: AsyncSession = Depends(get_session)