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)
|