提交
This commit is contained in:
291
backend/app/api/user.py
Normal file
291
backend/app/api/user.py
Normal 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)
|
||||
Reference in New Issue
Block a user