import csv from pathlib import Path from sqlalchemy.exc import SQLAlchemyError from sqlalchemy.orm import Session from backend.app.core.error_codes import ErrorCode from backend.app.core.exceptions import AppException from backend.app.repositories.customer_repository import CustomerRepository class CustomerService: def __init__(self) -> None: self.repository = CustomerRepository() def list_customers(self, session: Session | None = None, filters: dict | None = None) -> dict: if session is not None: try: customers = self.repository.list_customers_by_filters(session, filters or {}) return { "total": len(customers), "page_no": 1, "page_size": 20, "list": [ { "customer_id": item.id, "customer_name": item.customer_name, "mobile": item.mobile, "address": item.address, "settlement_type": item.settlement_type, "settlement_days": item.settlement_days, "customer_type": item.customer_type, "salesman_id": item.salesman_id, "salesman_name": "", "arrears_amount": 0, "status": 1, } for item in customers ], } except SQLAlchemyError: pass return { "total": 1, "page_no": 1, "page_size": 20, "list": [ { "customer_id": 3001, "customer_name": "演示客户", "mobile": "13900000000", "address": "杭州市西湖区演示地址 1 号", "settlement_type": "monthly", "settlement_days": 30, "customer_type": "channel", "salesman_id": 1, "salesman_name": "演示业务员", "arrears_amount": 0, "status": 1, } ], } def create_customer( self, session: Session | None = None, payload: dict | None = None, current_user: dict | None = None, ) -> dict: if session is not None and payload is not None: try: if not payload["customer_name"].strip(): raise AppException(code=ErrorCode.PARAM_ERROR, message="客户姓名不能为空", status_code=400) if not payload["mobile"].strip(): raise AppException(code=ErrorCode.PARAM_ERROR, message="客户手机号不能为空", status_code=400) existed = self.repository.find_by_name_and_mobile( session, payload["customer_name"], payload["mobile"], ) if existed is not None: raise AppException(code=ErrorCode.DUPLICATE, message="客户已存在", status_code=400) customer = self.repository.create_customer( session, { "customer_name": payload["customer_name"], "mobile": payload["mobile"], "address": payload.get("address"), "settlement_type": payload.get("settlement_type"), "settlement_days": payload.get("settlement_days", 0), "customer_type": payload.get("customer_type"), "salesman_id": self._resolve_salesman_id(payload, current_user), "credit_limit": payload.get("credit_limit", 0), "remark": payload.get("remark"), }, ) session.commit() return { "customer_id": customer.id, "customer_name": customer.customer_name, "mobile": customer.mobile, } except AppException: session.rollback() raise except SQLAlchemyError: session.rollback() return { "customer_id": 3002, "customer_name": "新建演示客户", "mobile": "13800000009", } def get_customer( self, customer_id: int, session: Session | None = None, current_user: dict | None = None, ) -> dict: if session is not None: try: customer = self.repository.get_customer(session, customer_id) if customer is not None: self._ensure_customer_access(customer, current_user) return { "customer_id": customer.id, "customer_name": customer.customer_name, "mobile": customer.mobile, "address": customer.address, "settlement_type": customer.settlement_type, "settlement_days": customer.settlement_days, "customer_type": customer.customer_type, "salesman_id": customer.salesman_id, "salesman_name": "", "arrears_amount": 0, "status": 1, "credit_limit": float(customer.credit_limit or 0), "remark": customer.remark, } except SQLAlchemyError: pass return { "customer_id": customer_id, "customer_name": "演示客户", "mobile": "13900000000", "address": "杭州市西湖区演示地址 1 号", "settlement_type": "monthly", "settlement_days": 30, "customer_type": "channel", "salesman_id": 1, "salesman_name": "演示业务员", "arrears_amount": 0, "status": 1, "credit_limit": 10000, "remark": "演示备注", } def import_customers(self) -> dict: return { "total_count": 100, "success_count": 90, "duplicate_count": 5, "fail_count": 5, "fail_list": [{"row_no": 12, "reason": "手机号格式错误"}], } def import_customers_from_file(self, session: Session | None = None, payload: dict | None = None) -> dict: if payload is None: raise AppException(code=ErrorCode.PARAM_ERROR, message="导入参数不能为空", status_code=400) file_url = (payload.get("file_url") or "").strip() import_mode = (payload.get("import_mode") or "skip_duplicate").strip() if not file_url: raise AppException(code=ErrorCode.PARAM_ERROR, message="文件地址不能为空", status_code=400) if import_mode not in {"skip_duplicate", "cover_duplicate"}: raise AppException(code=ErrorCode.PARAM_ERROR, message="导入模式不正确", status_code=400) if session is not None: try: rows = self._load_import_rows(file_url) result = self._import_customer_rows(session, rows, import_mode, file_url) session.commit() return result except AppException: session.rollback() raise except SQLAlchemyError: session.rollback() return { "total_count": 100, "success_count": 90, "duplicate_count": 5, "fail_count": 5, "fail_list": [{"row_no": 12, "reason": "手机号为空"}], } def _load_import_rows(self, file_url: str) -> list[dict]: path = self._resolve_import_path(file_url) if not path.exists() or not path.is_file(): raise AppException(code=ErrorCode.PARAM_ERROR, message="导入文件不存在", status_code=400) suffix = path.suffix.lower() if suffix == ".csv": with path.open("r", encoding="utf-8-sig", newline="") as file: return list(csv.DictReader(file)) if suffix in {".xlsx", ".xls"}: try: from openpyxl import load_workbook except ImportError as exc: raise AppException(code=ErrorCode.SYSTEM_ERROR, message="缺少 Excel 解析依赖 openpyxl", status_code=500) from exc workbook = load_workbook(path, read_only=True, data_only=True) sheet = workbook.active rows = list(sheet.iter_rows(values_only=True)) if not rows: return [] headers = [str(item).strip() if item is not None else "" for item in rows[0]] return [ {headers[index]: value for index, value in enumerate(row) if index < len(headers)} for row in rows[1:] ] raise AppException(code=ErrorCode.PARAM_ERROR, message="仅支持导入 csv、xlsx、xls 文件", status_code=400) def _resolve_import_path(self, file_url: str) -> Path: normalized = file_url.replace("\\", "/") if normalized.startswith(("http://", "https://")): raise AppException(code=ErrorCode.PARAM_ERROR, message="当前阶段仅支持本地文件路径导入", status_code=400) path = Path(file_url) if not path.is_absolute(): path = Path.cwd() / path return path def _import_customer_rows(self, session: Session, rows: list[dict], import_mode: str, file_url: str) -> dict: fail_list: list[dict] = [] success_count = 0 duplicate_count = 0 for index, raw_row in enumerate(rows, start=2): try: payload = self._normalize_import_row(raw_row, file_url) existed = self.repository.find_by_mobile(session, payload["mobile"]) if existed is not None: duplicate_count += 1 if import_mode == "skip_duplicate": continue self.repository.update_customer(session, existed, payload) success_count += 1 continue self.repository.create_customer(session, payload) success_count += 1 except AppException as exc: fail_list.append({"row_no": index, "reason": exc.message}) return { "total_count": len(rows), "success_count": success_count, "duplicate_count": duplicate_count, "fail_count": len(fail_list), "fail_list": fail_list, } def _normalize_import_row(self, raw_row: dict, file_url: str) -> dict: customer_name = self._pick_import_value(raw_row, ["customer_name", "客户姓名", "姓名"]) mobile = self._pick_import_value(raw_row, ["mobile", "手机号", "手机号码"]) if not customer_name: raise AppException(code=ErrorCode.PARAM_ERROR, message="客户姓名不能为空", status_code=400) if not mobile: raise AppException(code=ErrorCode.PARAM_ERROR, message="手机号不能为空", status_code=400) settlement_days = self._parse_int(self._pick_import_value(raw_row, ["settlement_days", "账期天数"]), 0) salesman_id = self._parse_int(self._pick_import_value(raw_row, ["salesman_id", "业务员ID"]), None) credit_limit = self._parse_float(self._pick_import_value(raw_row, ["credit_limit", "信用额度"]), 0) return { "customer_name": customer_name, "mobile": mobile, "address": self._pick_import_value(raw_row, ["address", "客户地址", "地址"]) or None, "settlement_type": self._pick_import_value(raw_row, ["settlement_type", "结算方式"]) or "monthly", "settlement_days": settlement_days, "customer_type": self._pick_import_value(raw_row, ["customer_type", "客户类型"]) or "channel", "salesman_id": salesman_id, "credit_limit": credit_limit, "remark": self._pick_import_value(raw_row, ["remark", "备注"]) or f"导入文件:{file_url}", } def _pick_import_value(self, raw_row: dict, keys: list[str]) -> str: for key in keys: value = raw_row.get(key) if value is None: continue text = str(value).strip() if text: return text return "" def _parse_int(self, value: str, default: int | None) -> int | None: if value == "": return default try: return int(float(value)) except ValueError as exc: raise AppException(code=ErrorCode.PARAM_ERROR, message="数字字段格式不正确", status_code=400) from exc def _parse_float(self, value: str, default: float) -> float: if value == "": return default try: return float(value) except ValueError as exc: raise AppException(code=ErrorCode.PARAM_ERROR, message="金额字段格式不正确", status_code=400) from exc def normalize_list_filters(self, filters: dict | None, current_user: dict | None = None) -> dict: normalized = dict(filters or {}) if current_user and current_user.get("role_code") == "salesman": normalized["salesman_id"] = current_user.get("user_id") return normalized def _resolve_salesman_id(self, payload: dict, current_user: dict | None) -> int | None: if current_user and current_user.get("role_code") == "salesman": return current_user.get("user_id") return payload.get("salesman_id") def _ensure_customer_access(self, customer: object, current_user: dict | None) -> None: if not current_user: return role_code = current_user.get("role_code") if role_code in {"admin", "manager"}: return if role_code == "salesman": if customer.salesman_id != current_user.get("user_id"): raise AppException(code=ErrorCode.FORBIDDEN, message="无权限访问他人客户", status_code=403) return raise AppException(code=ErrorCode.FORBIDDEN, message="无权限访问", status_code=403) customer_service = CustomerService()