39 lines
1.5 KiB
Python
39 lines
1.5 KiB
Python
"""trade_order 仓储:按单号/客户查申请单 + 状态流转(条件更新,不 commit)。"""
|
|||
|
|
from __future__ import annotations
|
||
|
|
|
||
|
|
from sqlalchemy import select, update
|
||
|
|
|
||
|
|
from model.trade_order import TradeOrder
|
||
|
|
from repositories.base import BaseRepository
|
||
|
|
|
||
|
|
|
||
|
|
class TradeOrderRepo(BaseRepository):
|
||
|
|
model = TradeOrder
|
||
|
|
|
||
|
|
async def get_by_order_no(self, order_no: str) -> TradeOrder | None:
|
||
|
|
return await self.db.scalar(
|
||
|
|
select(TradeOrder).where(TradeOrder.order_no == order_no)
|
||
|
|
)
|
||
|
|
|
||
|
|
async def list_by_customer(
|
||
|
|
self, customer_id: int, status: str | None = None
|
||
|
|
) -> list[TradeOrder]:
|
||
|
|
"""按客户查申请单,可按状态过滤,按创建时间倒序。"""
|
||
|
|
stmt = select(TradeOrder).where(TradeOrder.customer_id == customer_id)
|
||
|
|
if status is not None:
|
||
|
|
stmt = stmt.where(TradeOrder.status == status)
|
||
|
|
stmt = stmt.order_by(TradeOrder.id.desc())
|
||
|
|
return list((await self.db.scalars(stmt)).all())
|
||
|
|
|
||
|
|
async def update_status(self, order_id: int, *, status: str, **fields) -> int:
|
||
|
|
"""更新订单状态(可附带 risk_alert_id / confirm_time / cancel_reason 等)。
|
||
|
|
|
||
|
|
不 commit,由 service 层事务统一提交;返回 rowcount(0 表示订单不存在)。
|
||
|
|
"""
|
||
|
|
result = await self.db.execute(
|
||
|
|
update(TradeOrder)
|
||
|
|
.where(TradeOrder.id == order_id)
|
||
|
|
.values(status=status, **fields)
|
||
|
|
)
|
||
|
|
return result.rowcount
|