"""用户数据路由:偏好 / 自选股 / 交割单(个人实盘买卖点)。""" from __future__ import annotations import json from fastapi import APIRouter, Depends, File, HTTPException, UploadFile from sqlalchemy import delete, select, text from sqlalchemy.ext.asyncio import AsyncSession from .. import cache from ..auth import require_user from ..db import get_session from ..models import HoldingItem, StockBasic, UserPreference, UserTrade, WatchlistItem from ..schemas import ( HoldingOp, PreferencesOut, PreferencesUpdate, TradesClearResponse, TradesImportResponse, UserTradeOut, WatchlistOp, ) from ..trades import parse_statement router = APIRouter() # ---------- 用户偏好 ---------- @router.get("/preferences", response_model=PreferencesOut) async def get_preferences( session: AsyncSession = Depends(get_session), user=Depends(require_user) ) -> PreferencesOut: prefs: dict[str, object] = {} rows = ( await session.execute(select(UserPreference).where(UserPreference.user_id == user.id)) ).scalars().all() for r in rows: try: prefs[r.key] = json.loads(r.value_json) except Exception: # noqa: BLE001 prefs[r.key] = None return PreferencesOut(prefs=prefs) @router.put("/preferences", response_model=PreferencesOut) async def put_preferences( req: PreferencesUpdate, session: AsyncSession = Depends(get_session), user=Depends(require_user), ) -> PreferencesOut: """部分更新:只覆盖出现的 key;值为 null 表示删除该 key。返回更新后的全量。""" for key, value in req.prefs.items(): if not key or len(key) > 64: continue if value is None: await session.execute( text("DELETE FROM user_preferences WHERE user_id = :u AND key = :k"), {"u": user.id, "k": key}, ) continue existing = ( await session.execute( select(UserPreference).where( UserPreference.user_id == user.id, UserPreference.key == key ) ) ).scalars().first() vj = json.dumps(value, ensure_ascii=False) if existing: existing.value_json = vj else: session.add(UserPreference(user_id=user.id, key=key, value_json=vj)) await session.commit() return await get_preferences(session=session, user=user) # ---------- 自选股 ---------- @router.get("/watchlist", response_model=list[str]) async def get_watchlist( session: AsyncSession = Depends(get_session), user=Depends(require_user) ) -> list[str]: """当前用户自选股 ts_code 列表(加入时间倒序)。""" rows = ( await session.execute( select(WatchlistItem.ts_code) .where(WatchlistItem.user_id == user.id) .order_by(WatchlistItem.created_at.desc(), WatchlistItem.id.desc()) ) ).scalars().all() return list(rows) @router.post("/watchlist", response_model=list[str]) async def add_watchlist( req: WatchlistOp, session: AsyncSession = Depends(get_session), user=Depends(require_user), ) -> list[str]: exists = ( await session.execute( select(WatchlistItem.id).where( WatchlistItem.user_id == user.id, WatchlistItem.ts_code == req.ts_code ) ) ).scalar_one_or_none() 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) @router.delete("/watchlist/{ts_code}", response_model=list[str]) async def remove_watchlist( ts_code: str, session: AsyncSession = Depends(get_session), user=Depends(require_user), ) -> list[str]: await session.execute( text("DELETE FROM watchlist_items WHERE user_id = :u AND ts_code = :c"), {"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) # ---------- 持仓股 ---------- @router.get("/holdings", response_model=list[str]) async def get_holdings( session: AsyncSession = Depends(get_session), user=Depends(require_user) ) -> list[str]: """当前用户持仓股 ts_code 列表(加入时间倒序)。""" rows = ( await session.execute( select(HoldingItem.ts_code) .where(HoldingItem.user_id == user.id) .order_by(HoldingItem.created_at.desc(), HoldingItem.id.desc()) ) ).scalars().all() return list(rows) @router.post("/holdings", response_model=list[str]) async def add_holding( req: HoldingOp, session: AsyncSession = Depends(get_session), user=Depends(require_user), ) -> list[str]: exists = ( await session.execute( select(HoldingItem.id).where( HoldingItem.user_id == user.id, HoldingItem.ts_code == req.ts_code ) ) ).scalar_one_or_none() if exists is None: session.add(HoldingItem(user_id=user.id, ts_code=req.ts_code)) await session.commit() await cache.bump_version(f"holding:{user.id}") # 作废该用户的股票列表缓存 return await get_holdings(session=session, user=user) @router.delete("/holdings/{ts_code}", response_model=list[str]) async def remove_holding( ts_code: str, session: AsyncSession = Depends(get_session), user=Depends(require_user), ) -> list[str]: await session.execute( text("DELETE FROM holding_items WHERE user_id = :u AND ts_code = :c"), {"u": user.id, "c": ts_code}, ) await session.commit() await cache.bump_version(f"holding:{user.id}") # 作废该用户的股票列表缓存 return await get_holdings(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)