292 lines
10 KiB
Python
292 lines
10 KiB
Python
"""用户数据路由:偏好 / 自选股 / 交割单(个人实盘买卖点)。"""
|
||
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)
|