65 lines
2.7 KiB
Python
65 lines
2.7 KiB
Python
from sqlalchemy import select
|
|
from sqlalchemy.orm import Session
|
|
|
|
from backend.app.models.business import Customer
|
|
|
|
|
|
class CustomerRepository:
|
|
"""客户数据访问层。"""
|
|
|
|
def find_by_name_and_mobile(self, session: Session, customer_name: str, mobile: str) -> Customer | None:
|
|
stmt = select(Customer).where(
|
|
Customer.customer_name == customer_name,
|
|
Customer.mobile == mobile,
|
|
Customer.deleted == 0,
|
|
)
|
|
return session.execute(stmt).scalar_one_or_none()
|
|
|
|
def list_customers_by_filters(self, session: Session, filters: dict) -> list[Customer]:
|
|
stmt = select(Customer).where(Customer.deleted == 0)
|
|
|
|
if filters.get("customer_name"):
|
|
stmt = stmt.where(Customer.customer_name.contains(filters["customer_name"]))
|
|
if filters.get("mobile"):
|
|
stmt = stmt.where(Customer.mobile.contains(filters["mobile"]))
|
|
if filters.get("customer_type"):
|
|
stmt = stmt.where(Customer.customer_type == filters["customer_type"])
|
|
if filters.get("settlement_type"):
|
|
stmt = stmt.where(Customer.settlement_type == filters["settlement_type"])
|
|
if filters.get("salesman_id") is not None:
|
|
stmt = stmt.where(Customer.salesman_id == filters["salesman_id"])
|
|
|
|
stmt = stmt.order_by(Customer.id.desc())
|
|
return list(session.execute(stmt).scalars())
|
|
|
|
def get_customer(self, session: Session, customer_id: int) -> Customer | None:
|
|
stmt = select(Customer).where(Customer.id == customer_id, Customer.deleted == 0)
|
|
return session.execute(stmt).scalar_one_or_none()
|
|
|
|
def create_customer(self, session: Session, payload: dict) -> Customer:
|
|
customer = Customer(**payload)
|
|
session.add(customer)
|
|
session.flush()
|
|
return customer
|
|
|
|
def update_customer(self, session: Session, customer: Customer, payload: dict) -> Customer:
|
|
customer.customer_name = payload["customer_name"]
|
|
customer.mobile = payload["mobile"]
|
|
customer.address = payload.get("address")
|
|
customer.settlement_type = payload.get("settlement_type")
|
|
customer.settlement_days = payload.get("settlement_days", 0)
|
|
customer.customer_type = payload.get("customer_type")
|
|
customer.salesman_id = payload.get("salesman_id")
|
|
customer.credit_limit = payload.get("credit_limit", 0)
|
|
customer.remark = payload.get("remark")
|
|
session.add(customer)
|
|
session.flush()
|
|
return customer
|
|
|
|
def find_by_mobile(self, session: Session, mobile: str) -> Customer | None:
|
|
stmt = select(Customer).where(
|
|
Customer.mobile == mobile,
|
|
Customer.deleted == 0,
|
|
)
|
|
return session.execute(stmt).scalar_one_or_none()
|