from fastapi import APIRouter, Depends, HTTPException, Query, Body, UploadFile, File, Form from pydantic import BaseModel, ConfigDict, Field from sqlalchemy.orm import Session from sqlalchemy import asc, desc from typing import List, Optional import logging from database import get_db from models.model import Model from response import ApiResponse, PageData from curve_client import fetch_and_validate_curve import requests import os import shutil from pathlib import Path # Configure logging logger = logging.getLogger(__name__) router = APIRouter(prefix="/api/models", tags=["models"]) class PushToSearchBody(BaseModel): """推送 Meilisearch 请求体;字段名 model_ids 需关闭 protected_namespaces 避免 Pydantic 警告。""" model_config = ConfigDict(protected_namespaces=()) model_ids: List[int] = Field(..., min_length=1, description="型号 ID 列表") # Meilisearch 配置(从环境变量读取) MEILISEARCH_URL = os.getenv("MEILISEARCH_URL", "http://localhost:7700") MEILISEARCH_API_KEY = os.getenv("MEILISEARCH_API_KEY", "") MEILISEARCH_INDEX = os.getenv("MEILISEARCH_INDEX", "models") # 文件上传配置 UPLOAD_FOLDER = Path("/data/project/autoeq/measurements") ALLOWED_EXTENSIONS = {'.csv', '.txt', '.json'} # 记录上传路径配置 logger.info(f"UPLOAD_FOLDER configured as: {UPLOAD_FOLDER}") logger.info(f"UPLOAD_FOLDER absolute path: {UPLOAD_FOLDER.absolute()}") @router.get("/", response_model=ApiResponse) def get_models( skip: int = Query(0, ge=0, description="跳过记录数"), limit: int = Query(100, ge=1, le=1000, description="返回记录数"), brand_name: Optional[str] = Query(None, description="按品牌名称模糊查询"), name: Optional[str] = Query(None, description="按型号名称模糊查询"), sort_by: str = Query("id", description="排序字段:id 或 create_at"), sort_order: str = Query("desc", description="排序方向:asc 或 desc"), db: Session = Depends(get_db) ): """获取所有型号列表(支持按品牌名称和型号名称模糊查询;支持按 id、create_at 排序)""" try: logger.info( f"Getting models: skip={skip}, limit={limit}, brand_name={brand_name}, name={name}, " f"sort_by={sort_by}, sort_order={sort_order}" ) query = db.query(Model) if brand_name: query = query.filter(Model.brand_name.like(f"%{brand_name}%")) if name: query = query.filter(Model.name.like(f"%{name}%")) total = query.count() sort_columns = {"id": Model.id, "create_at": Model.create_at} order_col = sort_columns.get(sort_by, Model.id) order_dir = (sort_order or "desc").lower() if order_dir not in ("asc", "desc"): order_dir = "desc" if order_dir == "desc": query = query.order_by(desc(order_col)) else: query = query.order_by(asc(order_col)) models = query.offset(skip).limit(limit).all() logger.info(f"Found {len(models)} models, total={total}") if not models: logger.warning("No models found") return ApiResponse(code=2, msg="empty", data=None) # 转换为字典列表 models_data = [model.to_dict() for model in models] return ApiResponse(code=1, msg="success", data={"items": models_data, "total": total, "skip": skip, "limit": limit}) except Exception as e: logger.error(f"Error getting models: {str(e)}", exc_info=True) return ApiResponse(code=0, msg="error", data=None) @router.get("/{model_id}", response_model=ApiResponse) def get_model(model_id: int, db: Session = Depends(get_db)): """获取单个型号""" try: logger.info(f"Getting model: id={model_id}") model = db.query(Model).filter(Model.id == model_id).first() if not model: logger.warning(f"Model not found: id={model_id}") return ApiResponse(code=2, msg="empty", data=None) logger.info(f"Model found: {model.to_dict()}") return ApiResponse(code=1, msg="success", data=model.to_dict()) except Exception as e: logger.error(f"Error getting model {model_id}: {str(e)}", exc_info=True) return ApiResponse(code=0, msg="error", data=None) @router.post("/", response_model=ApiResponse) def create_model( brand_name: str = Form(...), name: str = Form(...), form: str = Form(None), rig: str = Form(None), source: str = Form(None), eq_key: str = Form(None), measurement_file: UploadFile = File(None), db: Session = Depends(get_db) ): """创建新型号(支持文件上传)""" try: logger.info(f"Creating model: brand_name={brand_name}, name={name}, form={form}, source={source}") # 检查是否已存在 existing = db.query(Model).filter( Model.brand_name == brand_name, Model.name == name ).first() if existing: logger.warning(f"Model already exists: brand_name={brand_name}, name={name}") return ApiResponse(code=0, msg="该品牌下型号名称已存在", data=None) # 处理文件上传 measurement_filename = None if measurement_file and measurement_file.filename: logger.info(f"Uploading measurement file: {measurement_file.filename}") logger.info(f"UPLOAD_FOLDER is: {UPLOAD_FOLDER}") # 验证文件扩展名 file_ext = os.path.splitext(measurement_file.filename)[1].lower() if file_ext not in ALLOWED_EXTENSIONS: logger.error(f"Unsupported file format: {file_ext}") return ApiResponse(code=0, msg=f"不支持的文件格式:{file_ext}", data=None) # 创建保存路径:autoeq/measurements/{source}/data/{form}/{filename} save_dir = UPLOAD_FOLDER / source / "data" / form logger.info(f"Creating directory: {save_dir}") save_dir.mkdir(parents=True, exist_ok=True) # 保存文件(保留原文件名) file_path = save_dir / measurement_file.filename logger.info(f"Saving file to: {file_path}") with open(file_path, "wb") as buffer: shutil.copyfileobj(measurement_file.file, buffer) measurement_filename = measurement_file.filename logger.info(f"File saved: {file_path}") db_model = Model( brand_name=brand_name, name=name, form=form, rig=rig, source=source, eq_key=eq_key ) db.add(db_model) db.commit() db.refresh(db_model) logger.info(f"Model created successfully: id={db_model.id}") return ApiResponse(code=1, msg="success", data=db_model.to_dict()) except Exception as e: logger.error(f"Error creating model: {str(e)}", exc_info=True) db.rollback() return ApiResponse(code=0, msg="error", data=None) @router.put("/{model_id}", response_model=ApiResponse) def update_model( model_id: int, brand_name: Optional[str] = Form(None), name: Optional[str] = Form(None), form: Optional[str] = Form(None), rig: Optional[str] = Form(None), source: Optional[str] = Form(None), eq_key: Optional[str] = Form(None), measurement_file: UploadFile = File(None), db: Session = Depends(get_db), ): """更新型号(与 POST 一致,支持 multipart/form-data,可选上传频响文件)""" try: logger.info( f"Updating model: id={model_id}, brand_name={brand_name}, name={name}, form={form}" ) db_model = db.query(Model).filter(Model.id == model_id).first() if not db_model: logger.warning(f"Model not found: id={model_id}") return ApiResponse(code=2, msg="empty", data=None) new_brand_name = db_model.brand_name if brand_name is None else brand_name new_name = db_model.name if name is None else name if new_brand_name != db_model.brand_name or new_name != db_model.name: existing = ( db.query(Model) .filter(Model.brand_name == new_brand_name, Model.name == new_name) .first() ) if existing: logger.warning( f"Model already exists: brand_name={new_brand_name}, name={new_name}" ) return ApiResponse(code=0, msg="该品牌下型号名称已存在", data=None) eff_source = db_model.source if source is None else source eff_form = db_model.form if form is None else form if measurement_file and measurement_file.filename: logger.info(f"Uploading measurement file: {measurement_file.filename}") file_ext = os.path.splitext(measurement_file.filename)[1].lower() if file_ext not in ALLOWED_EXTENSIONS: logger.error(f"Unsupported file format: {file_ext}") return ApiResponse(code=0, msg=f"不支持的文件格式:{file_ext}", data=None) if not eff_source or not eff_form: return ApiResponse(code=0, msg="上传频响文件需要来源与形式字段", data=None) save_dir = UPLOAD_FOLDER / eff_source / "data" / eff_form save_dir.mkdir(parents=True, exist_ok=True) file_path = save_dir / measurement_file.filename with open(file_path, "wb") as buffer: shutil.copyfileobj(measurement_file.file, buffer) logger.info(f"File saved: {file_path}") if brand_name is not None: db_model.brand_name = brand_name if name is not None: db_model.name = name if form is not None: db_model.form = form if rig is not None: db_model.rig = rig if source is not None: db_model.source = source if eq_key is not None: db_model.eq_key = eq_key db.commit() db.refresh(db_model) logger.info(f"Model updated successfully: id={db_model.id}") return ApiResponse(code=1, msg="success", data=db_model.to_dict()) except Exception as e: logger.error(f"Error updating model {model_id}: {str(e)}", exc_info=True) db.rollback() return ApiResponse(code=0, msg="error", data=None) @router.delete("/{model_id}", response_model=ApiResponse) def delete_model(model_id: int, db: Session = Depends(get_db)): """删除型号""" try: logger.info(f"Deleting model: id={model_id}") db_model = db.query(Model).filter(Model.id == model_id).first() if not db_model: logger.warning(f"Model not found: id={model_id}") return ApiResponse(code=2, msg="empty", data=None) db.delete(db_model) db.commit() logger.info(f"Model deleted successfully: id={model_id}") return ApiResponse(code=1, msg="success", data=None) except Exception as e: logger.error(f"Error deleting model {model_id}: {str(e)}", exc_info=True) db.rollback() return ApiResponse(code=0, msg="error", data=None) def _get_models_by_ids(model_ids: List[int], db: Session): if not model_ids: return None, ApiResponse(code=0, msg="请选择要推送的型号", data=None) models = db.query(Model).filter(Model.id.in_(model_ids)).all() if not models: return None, ApiResponse(code=0, msg="未找到选中的型号数据", data=None) return models, None def _validate_models_curve(models) -> List[dict]: validation_errors = [] for model in models: ok, reason = fetch_and_validate_curve( model.brand_name, model.name, model.form or "", ) if not ok: validation_errors.append( { "id": model.id, "brand_name": model.brand_name, "name": model.name, "form": model.form, "reason": reason, } ) return validation_errors def _build_meilisearch_documents(models) -> List[dict]: push_data = [] for model in models: model_dict = { "id": model.id, "brand_name": model.brand_name, "name": model.name, "rig": model.rig, "form": model.form, "source": model.source, } push_data.append({k: v for k, v in model_dict.items() if v is not None}) return push_data def _push_documents_to_meilisearch(push_data: List[dict]) -> ApiResponse: headers = { "Authorization": f"Bearer {MEILISEARCH_API_KEY}", "Content-Type": "application/json", } try: response = requests.post( f"{MEILISEARCH_URL}/indexes/{MEILISEARCH_INDEX}/documents", json=push_data, headers=headers, timeout=30, ) if response.status_code not in [200, 202]: return ApiResponse(code=0, msg=f"推送到 Meilisearch 失败:{response.text}", data=None) task_info = response.json() return ApiResponse( code=1, msg="success", data={ "pushed_count": len(push_data), "task_uid": task_info.get("taskUid"), "models": push_data, }, ) except requests.exceptions.RequestException as e: return ApiResponse(code=0, msg=f"连接 Meilisearch 失败:{str(e)}", data=None) @router.post("/push-to-search/validate", response_model=ApiResponse) def validate_push_to_search(body: PushToSearchBody, db: Session = Depends(get_db)): """推送前校验:拉取并验证曲线 parametric_eq 数据""" try: models, err = _get_models_by_ids(body.model_ids, db) if err: return err validation_errors = _validate_models_curve(models) if validation_errors: names = "、".join( f"{e['brand_name']} {e['name']}" for e in validation_errors[:5] ) suffix = " 等" if len(validation_errors) > 5 else "" return ApiResponse( code=0, msg=f"曲线数据校验未通过:{names}{suffix}", data={"errors": validation_errors, "validated_count": 0}, ) return ApiResponse( code=1, msg="success", data={"validated_count": len(models)}, ) except Exception as e: logger.error("validate_push_to_search failed: %s", e, exc_info=True) return ApiResponse(code=0, msg=f"校验失败:{str(e)}", data=None) @router.post("/push-to-search", response_model=ApiResponse) def push_to_search(body: PushToSearchBody, db: Session = Depends(get_db)): """推送型号数据到 Meilisearch(需先通过 validate 接口)""" try: models, err = _get_models_by_ids(body.model_ids, db) if err: return err push_data = _build_meilisearch_documents(models) return _push_documents_to_meilisearch(push_data) except Exception as e: logger.error("push_to_search failed: %s", e, exc_info=True) return ApiResponse(code=0, msg=f"推送失败:{str(e)}", data=None)