This commit is contained in:
2026-09-09 15:07:58 +08:00
parent d656c05b3d
commit 71a0f6e404
31 changed files with 3657 additions and 9 deletions

291
backend/app/api/user.py Normal file
View File

@@ -0,0 +1,291 @@
"""用户数据路由:偏好 / 自选股 / 交割单(个人实盘买卖点)。"""
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)