xxx
This commit is contained in:
@@ -19,6 +19,7 @@ def process_stp_task(self, task_id: str, file_path: str, stp_file_id: int,
|
||||
# - RustFS(Minio) 为同步客户端,不绑定循环,连一次后跨任务复用。
|
||||
from shared.services.redis_task_manager import redis_task_manager
|
||||
from moldinsight.storage.rustfs_storage import rustfs_manager
|
||||
from shared.database.database import db_manager
|
||||
|
||||
await redis_task_manager.reconnect()
|
||||
if not rustfs_manager.is_connected:
|
||||
@@ -28,6 +29,9 @@ def process_stp_task(self, task_id: str, file_path: str, stp_file_id: int,
|
||||
secret_key=settings.RUSTFS_SECRET_KEY,
|
||||
timeout=settings.RUSTFS_TIMEOUT,
|
||||
)
|
||||
# Celery worker 使用独立连接池配置(较小池)
|
||||
if not db_manager.is_connected:
|
||||
await db_manager.connect(role="celery")
|
||||
|
||||
await processing_service.process_file_with_storage(
|
||||
task_id, file_path, stp_file_id, process_params
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
"""
|
||||
inventory 入口 — 使用 shared.app_factory.create_app() 构建
|
||||
"""
|
||||
import os, sys
|
||||
from pathlib import Path
|
||||
|
||||
@@ -6,73 +9,21 @@ sys.path.insert(0, str(src_root))
|
||||
|
||||
os.chdir(Path(__file__).parent.parent.parent)
|
||||
|
||||
from fastapi import FastAPI, Request
|
||||
from fastapi.staticfiles import StaticFiles
|
||||
from fastapi.middleware.cors import CORSMiddleware
|
||||
from fastapi.responses import JSONResponse, FileResponse
|
||||
import asyncio, time
|
||||
from shared.app_factory import create_app
|
||||
|
||||
from shared.config.settings import settings
|
||||
from shared.services.auth_routes import router as auth_router
|
||||
from shared.utils.logger import setup_logging, get_logger
|
||||
from shared.database.init_db import init_database
|
||||
|
||||
setup_logging()
|
||||
logger = get_logger(__name__)
|
||||
|
||||
app = FastAPI(title="Gemold - 进销存管理系统", version="4.0.0")
|
||||
|
||||
app.add_middleware(CORSMiddleware, allow_origins=["*"], allow_credentials=True, allow_methods=["*"], allow_headers=["*"])
|
||||
|
||||
@app.middleware("http")
|
||||
async def log_requests(request: Request, call_next):
|
||||
start_time = time.time()
|
||||
response = await call_next(request)
|
||||
duration = time.time() - start_time
|
||||
if response.status_code >= 400:
|
||||
logger.warning(f"[HTTP] {request.method} {request.url.path} -> {response.status_code} ({duration:.2f}s)")
|
||||
return response
|
||||
|
||||
@app.on_event("startup")
|
||||
async def startup_event():
|
||||
success = await init_database(keep_connected=True)
|
||||
print(f"[{'OK' if success else 'FAIL'}] 数据库初始化")
|
||||
def _register_routers(app):
|
||||
"""注册 inventory 业务路由"""
|
||||
try:
|
||||
from shared.services.redis_task_manager import redis_task_manager
|
||||
await redis_task_manager.connect()
|
||||
print(f"[{'OK' if redis_task_manager.is_connected else 'WARN'}] Redis")
|
||||
from inventory.api import inventory_router
|
||||
app.include_router(inventory_router)
|
||||
except Exception as e:
|
||||
print(f"[WARN] Redis异常: {e}")
|
||||
print(f"[WARN] Inventory 路由: {e}")
|
||||
|
||||
@app.on_event("shutdown")
|
||||
async def shutdown_event():
|
||||
try:
|
||||
from shared.services.redis_task_manager import redis_task_manager
|
||||
await redis_task_manager.disconnect()
|
||||
except: pass
|
||||
|
||||
app.mount("/static", StaticFiles(directory=os.path.join(os.getcwd(), "static")), name="static")
|
||||
|
||||
app.include_router(auth_router)
|
||||
|
||||
try:
|
||||
from inventory.api import inventory_router
|
||||
app.include_router(inventory_router)
|
||||
except Exception as e:
|
||||
print(f"[WARN] Inventory 路由: {e}")
|
||||
|
||||
@app.get("/health")
|
||||
async def health():
|
||||
from shared.database.database import db_manager
|
||||
from sqlalchemy import text
|
||||
db_ok = False; db_error = None
|
||||
try:
|
||||
if not db_manager.is_connected: await db_manager.connect()
|
||||
async with db_manager.engine.begin() as conn: await conn.execute(text("SELECT 1"))
|
||||
db_ok = True
|
||||
except Exception as e: db_error = str(e)
|
||||
return {"status": "healthy", "service": "inventory", "version": "4.0.0", "database_connected": db_ok, "database_error": db_error}
|
||||
|
||||
@app.get("/{full_path:path}")
|
||||
async def spa_fallback(full_path: str):
|
||||
return FileResponse(os.path.join(os.getcwd(), "static", "index.html"))
|
||||
app = create_app(
|
||||
title="Gemold - 进销存管理系统",
|
||||
service_name="inventory",
|
||||
mount_html=False,
|
||||
register_routers=_register_routers,
|
||||
)
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
"""
|
||||
moldinsight 入口 — 使用 shared.app_factory.create_app() 构建
|
||||
"""
|
||||
import os, sys
|
||||
from pathlib import Path
|
||||
|
||||
@@ -6,83 +9,21 @@ sys.path.insert(0, str(src_root))
|
||||
|
||||
os.chdir(Path(__file__).parent.parent.parent)
|
||||
|
||||
from fastapi import FastAPI, Request
|
||||
from fastapi.staticfiles import StaticFiles
|
||||
from fastapi.middleware.cors import CORSMiddleware
|
||||
from fastapi.responses import JSONResponse, FileResponse
|
||||
import asyncio, time
|
||||
from shared.app_factory import create_app
|
||||
|
||||
from shared.config.settings import settings
|
||||
from shared.services.auth_routes import router as auth_router
|
||||
from shared.utils.logger import setup_logging, get_logger
|
||||
from shared.database.init_db import init_database
|
||||
|
||||
setup_logging()
|
||||
logger = get_logger(__name__)
|
||||
|
||||
app = FastAPI(title="Gemold - 模具分析引擎", version="4.0.0")
|
||||
|
||||
app.add_middleware(CORSMiddleware, allow_origins=["*"], allow_credentials=True, allow_methods=["*"], allow_headers=["*"])
|
||||
|
||||
@app.middleware("http")
|
||||
async def log_requests(request: Request, call_next):
|
||||
start_time = time.time()
|
||||
response = await call_next(request)
|
||||
duration = time.time() - start_time
|
||||
if response.status_code >= 400:
|
||||
logger.warning(f"[HTTP] {request.method} {request.url.path} -> {response.status_code} ({duration:.2f}s)")
|
||||
return response
|
||||
|
||||
@app.on_event("startup")
|
||||
async def startup_event():
|
||||
success = await init_database(keep_connected=True)
|
||||
print(f"[{'OK' if success else 'FAIL'}] 数据库初始化")
|
||||
def _register_routers(app):
|
||||
"""注册 moldinsight 业务路由"""
|
||||
try:
|
||||
from moldinsight.storage.rustfs_storage import rustfs_manager
|
||||
await rustfs_manager.connect(endpoint=settings.RUSTFS_ENDPOINT, access_key=settings.RUSTFS_ACCESS_KEY, secret_key=settings.RUSTFS_SECRET_KEY, timeout=settings.RUSTFS_TIMEOUT)
|
||||
print("[OK] RustFS连接成功")
|
||||
from moldinsight.api import router as moldinsight_router
|
||||
app.include_router(moldinsight_router, prefix="/api")
|
||||
except Exception as e:
|
||||
print(f"[WARN] RustFS连接失败: {e}")
|
||||
try:
|
||||
from shared.services.redis_task_manager import redis_task_manager
|
||||
await redis_task_manager.connect()
|
||||
print(f"[{'OK' if redis_task_manager.is_connected else 'WARN'}] Redis")
|
||||
except Exception as e:
|
||||
print(f"[WARN] Redis异常: {e}")
|
||||
print(f"[WARN] MoldInsight 路由: {e}")
|
||||
|
||||
@app.on_event("shutdown")
|
||||
async def shutdown_event():
|
||||
try:
|
||||
from shared.services.redis_task_manager import redis_task_manager
|
||||
await redis_task_manager.disconnect()
|
||||
except: pass
|
||||
|
||||
UPLOAD_DIR = Path("uploads"); UPLOAD_DIR.mkdir(exist_ok=True)
|
||||
Path("html_output").mkdir(exist_ok=True)
|
||||
|
||||
app.mount("/static", StaticFiles(directory=os.path.join(os.getcwd(), "static")), name="static")
|
||||
app.mount("/html", StaticFiles(directory=os.path.join(os.getcwd(), "html_output")), name="html")
|
||||
|
||||
app.include_router(auth_router)
|
||||
|
||||
try:
|
||||
from moldinsight.api import router as moldinsight_router
|
||||
app.include_router(moldinsight_router, prefix="/api")
|
||||
except Exception as e:
|
||||
print(f"[WARN] MoldInsight 路由: {e}")
|
||||
|
||||
@app.get("/health")
|
||||
async def health():
|
||||
from shared.database.database import db_manager
|
||||
from sqlalchemy import text
|
||||
db_ok = False; db_error = None
|
||||
try:
|
||||
if not db_manager.is_connected: await db_manager.connect()
|
||||
async with db_manager.engine.begin() as conn: await conn.execute(text("SELECT 1"))
|
||||
db_ok = True
|
||||
except Exception as e: db_error = str(e)
|
||||
return {"status": "healthy", "service": "moldinsight", "version": "4.0.0", "database_connected": db_ok, "database_error": db_error}
|
||||
|
||||
@app.get("/{full_path:path}")
|
||||
async def spa_fallback(full_path: str):
|
||||
return FileResponse(os.path.join(os.getcwd(), "static", "index.html"))
|
||||
app = create_app(
|
||||
title="Gemold - 模具分析引擎",
|
||||
service_name="moldinsight",
|
||||
mount_html=True,
|
||||
register_routers=_register_routers,
|
||||
)
|
||||
|
||||
@@ -26,6 +26,7 @@ from .sales_order_routes import router as sales_order_router
|
||||
from .dashboard_routes import router as dashboard_router
|
||||
from .finance_routes import router as finance_router
|
||||
from .material_routes import router as material_router
|
||||
from .purchase_demand_routes import router as purchase_demand_router
|
||||
|
||||
inventory_router = APIRouter(prefix="/api", tags=["进销存"])
|
||||
|
||||
@@ -38,6 +39,7 @@ inventory_router.include_router(stock_movement_router)
|
||||
inventory_router.include_router(purchase_order_router)
|
||||
inventory_router.include_router(sales_order_router)
|
||||
inventory_router.include_router(material_router)
|
||||
inventory_router.include_router(purchase_demand_router)
|
||||
inventory_router.include_router(dashboard_router)
|
||||
inventory_router.include_router(finance_router)
|
||||
|
||||
|
||||
@@ -51,7 +51,7 @@ async def create_customer(
|
||||
|
||||
customer = Customer(**data)
|
||||
db_session.add(customer)
|
||||
await db_session.commit()
|
||||
await db_session.flush()
|
||||
await db_session.refresh(customer)
|
||||
return CustomerResponse.from_orm(customer)
|
||||
|
||||
@@ -71,7 +71,7 @@ async def update_customer(
|
||||
for key, value in customer_data.dict().items():
|
||||
setattr(customer, key, value)
|
||||
|
||||
await db_session.commit()
|
||||
await db_session.flush()
|
||||
await db_session.refresh(customer)
|
||||
return CustomerResponse.from_orm(customer)
|
||||
|
||||
@@ -88,5 +88,5 @@ async def delete_customer(
|
||||
raise HTTPException(status_code=404, detail="客户不存在")
|
||||
|
||||
customer.is_active = False
|
||||
await db_session.commit()
|
||||
await db_session.flush()
|
||||
return {"message": "客户已删除"}
|
||||
|
||||
@@ -63,12 +63,12 @@ async def add_material_price_history(
|
||||
remark=price_data.remark
|
||||
)
|
||||
db_session.add(price_history)
|
||||
await db_session.commit()
|
||||
await db_session.flush()
|
||||
await db_session.refresh(price_history)
|
||||
|
||||
# 更新产品的成本价格为最新价格
|
||||
product.cost_price = price_data.price
|
||||
await db_session.commit()
|
||||
await db_session.flush()
|
||||
|
||||
return MaterialPriceHistoryResponse(
|
||||
id=price_history.id,
|
||||
@@ -212,7 +212,7 @@ async def add_material_supplier(
|
||||
min_order_quantity=supplier_data.min_order_quantity
|
||||
)
|
||||
db_session.add(material_supplier)
|
||||
await db_session.commit()
|
||||
await db_session.flush()
|
||||
await db_session.refresh(material_supplier)
|
||||
|
||||
return MaterialSupplierResponse(
|
||||
@@ -281,7 +281,7 @@ async def remove_material_supplier(
|
||||
raise HTTPException(status_code=404, detail="物料供应商关联不存在")
|
||||
|
||||
await db_session.delete(material_supplier)
|
||||
await db_session.commit()
|
||||
await db_session.flush()
|
||||
|
||||
return {"message": "物料供应商关联已删除"}
|
||||
|
||||
|
||||
@@ -114,7 +114,7 @@ async def create_product(
|
||||
product_dict["max_stock"] = 0
|
||||
product = Product(**product_dict)
|
||||
db_session.add(product)
|
||||
await db_session.commit()
|
||||
await db_session.flush()
|
||||
await db_session.refresh(product)
|
||||
return _build_product_response(product, 0)
|
||||
|
||||
@@ -178,7 +178,7 @@ async def create_product_from_task(
|
||||
db_session.add(product)
|
||||
await db_session.flush()
|
||||
stp_file.product_id = product.id
|
||||
await db_session.commit()
|
||||
await db_session.flush()
|
||||
await db_session.refresh(product)
|
||||
return _build_product_response(product, 0)
|
||||
|
||||
@@ -205,7 +205,7 @@ async def update_product(
|
||||
for key, value in product_dict.items():
|
||||
setattr(product, key, value)
|
||||
|
||||
await db_session.commit()
|
||||
await db_session.flush()
|
||||
await db_session.refresh(product)
|
||||
material_cost_map = await _calculate_material_cost_map(db_session, [product.id])
|
||||
return _build_product_response(product, material_cost_map.get(product.id, 0))
|
||||
@@ -223,7 +223,7 @@ async def delete_product(
|
||||
raise HTTPException(status_code=404, detail="产品不存在")
|
||||
|
||||
product.is_active = False
|
||||
await db_session.commit()
|
||||
await db_session.flush()
|
||||
return {"message": "产品已删除"}
|
||||
|
||||
|
||||
@@ -321,5 +321,5 @@ async def replace_product_bom(
|
||||
)
|
||||
)
|
||||
|
||||
await db_session.commit()
|
||||
await db_session.flush()
|
||||
return await get_product_bom(product_id, db_session, current_user)
|
||||
|
||||
@@ -0,0 +1,28 @@
|
||||
"""采购需求推导路由层 - 薄路由
|
||||
|
||||
业务逻辑下沉至 inventory.services.purchase_demand_service,路由只做参数校验与响应组装。
|
||||
路由前缀: /api/purchase-demands
|
||||
"""
|
||||
from fastapi import APIRouter, Depends
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from shared.database.database import get_db_session
|
||||
from shared.services.auth_service import get_current_active_user
|
||||
from shared.models.database import User
|
||||
from ..schemas import (
|
||||
PurchaseDemandCalculateRequest,
|
||||
PurchaseDemandResponse,
|
||||
)
|
||||
from ..services.purchase_demand_service import purchase_demand_service
|
||||
|
||||
router = APIRouter(prefix="/purchase-demands", tags=["采购需求推导"])
|
||||
|
||||
|
||||
@router.post("/calculate", response_model=PurchaseDemandResponse)
|
||||
async def calculate_purchase_demands(
|
||||
payload: PurchaseDemandCalculateRequest,
|
||||
db_session: AsyncSession = Depends(get_db_session),
|
||||
current_user: User = Depends(get_current_active_user),
|
||||
):
|
||||
"""根据销售订单 ID 列表,自动推导采购需求(BOM 展开 → 库存对比 → 供应商推荐)"""
|
||||
return await purchase_demand_service.calculate_demands(db_session, payload.sales_order_ids)
|
||||
@@ -51,7 +51,7 @@ async def create_supplier(
|
||||
|
||||
supplier = Supplier(**data)
|
||||
db_session.add(supplier)
|
||||
await db_session.commit()
|
||||
await db_session.flush()
|
||||
await db_session.refresh(supplier)
|
||||
return SupplierResponse.from_orm(supplier)
|
||||
|
||||
@@ -71,7 +71,7 @@ async def update_supplier(
|
||||
for key, value in supplier_data.dict().items():
|
||||
setattr(supplier, key, value)
|
||||
|
||||
await db_session.commit()
|
||||
await db_session.flush()
|
||||
await db_session.refresh(supplier)
|
||||
return SupplierResponse.from_orm(supplier)
|
||||
|
||||
@@ -88,5 +88,5 @@ async def delete_supplier(
|
||||
raise HTTPException(status_code=404, detail="供应商不存在")
|
||||
|
||||
supplier.is_active = False
|
||||
await db_session.commit()
|
||||
await db_session.flush()
|
||||
return {"message": "供应商已删除"}
|
||||
|
||||
@@ -44,6 +44,6 @@ async def create_warehouse(
|
||||
|
||||
warehouse = Warehouse(**data)
|
||||
db_session.add(warehouse)
|
||||
await db_session.commit()
|
||||
await db_session.flush()
|
||||
await db_session.refresh(warehouse)
|
||||
return WarehouseResponse.from_orm(warehouse)
|
||||
|
||||
@@ -56,6 +56,11 @@ from .material_schemas import (
|
||||
MaterialSupplierResponse,
|
||||
MaterialPriceTrendResponse
|
||||
)
|
||||
from .purchase_demand_schemas import (
|
||||
PurchaseDemandCalculateRequest,
|
||||
PurchaseDemandItemResponse,
|
||||
PurchaseDemandResponse
|
||||
)
|
||||
|
||||
from .common_schemas import PaginatedResponse
|
||||
|
||||
@@ -81,4 +86,5 @@ __all__ = [
|
||||
"PartnerProductStatementItemResponse", "FinancePartnerProductStatementResponse",
|
||||
"MaterialPriceHistoryCreate", "MaterialPriceHistoryResponse",
|
||||
"MaterialSupplierCreate", "MaterialSupplierResponse", "MaterialPriceTrendResponse",
|
||||
"PurchaseDemandCalculateRequest", "PurchaseDemandItemResponse", "PurchaseDemandResponse",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,36 @@
|
||||
"""采购需求推导相关数据模型
|
||||
|
||||
销售订单 → BOM 展开 → 物料需求 → 对比库存 → 生成采购建议
|
||||
"""
|
||||
from pydantic import BaseModel, Field
|
||||
from decimal import Decimal
|
||||
from typing import Optional, List
|
||||
|
||||
|
||||
class PurchaseDemandCalculateRequest(BaseModel):
|
||||
"""计算采购需求的请求体"""
|
||||
sales_order_ids: List[int] = Field(..., min_length=1, description="销售订单ID列表")
|
||||
|
||||
|
||||
class PurchaseDemandItemResponse(BaseModel):
|
||||
"""单个物料的采购建议"""
|
||||
material_id: int
|
||||
material_sku: str
|
||||
material_name: str
|
||||
required_quantity: Decimal = Field(..., description="BOM 需求量")
|
||||
available_quantity: Decimal = Field(..., description="当前库存量")
|
||||
shortage_quantity: Decimal = Field(..., description="缺口数量 = required - available")
|
||||
unit_cost: Decimal = Field(..., description="物料单价")
|
||||
estimated_cost: Decimal = Field(..., description="预计采购金额 = shortage × unit_cost")
|
||||
suggested_supplier_id: Optional[int] = Field(None, description="建议供应商ID")
|
||||
suggested_supplier_name: Optional[str] = Field(None, description="建议供应商名称")
|
||||
supplier_lead_time: Optional[int] = Field(None, description="供应商交货周期(天)")
|
||||
|
||||
|
||||
class PurchaseDemandResponse(BaseModel):
|
||||
"""采购需求计算结果"""
|
||||
items: List[PurchaseDemandItemResponse] = Field(default_factory=list)
|
||||
total_estimated_cost: Decimal = Field(default=Decimal("0"), description="预计采购总金额")
|
||||
shortage_count: int = Field(default=0, description="缺货物料种类数")
|
||||
source_order_ids: List[int] = Field(default_factory=list, description="来源销售订单ID")
|
||||
source_order_nos: List[str] = Field(default_factory=list, description="来源销售订单编号")
|
||||
@@ -0,0 +1,183 @@
|
||||
"""采购需求自动推导服务
|
||||
|
||||
销售订单确认 → 按 BOM 展开物料需求 → 对比当前库存 → 自动生成采购建议(缺多少、建议供应商、预计金额)
|
||||
"""
|
||||
from math import ceil
|
||||
from decimal import Decimal
|
||||
from typing import List
|
||||
|
||||
from fastapi import HTTPException
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from sqlalchemy import select, func
|
||||
|
||||
from shared.models.database import (
|
||||
Product,
|
||||
ProductMaterial,
|
||||
SalesOrder,
|
||||
SalesOrderItem,
|
||||
Inventory,
|
||||
MaterialSupplier,
|
||||
Supplier,
|
||||
)
|
||||
from ..schemas.purchase_demand_schemas import (
|
||||
PurchaseDemandItemResponse,
|
||||
PurchaseDemandResponse,
|
||||
)
|
||||
|
||||
|
||||
class PurchaseDemandService:
|
||||
"""采购需求推导服务"""
|
||||
|
||||
@staticmethod
|
||||
async def calculate_demands(
|
||||
db_session: AsyncSession,
|
||||
sales_order_ids: List[int],
|
||||
) -> PurchaseDemandResponse:
|
||||
"""
|
||||
核心算法:
|
||||
1. 批量查询销售订单 + 明细项
|
||||
2. 按 BOM 展开所有成品所需的物料(含损耗率)
|
||||
3. 聚合跨订单的同一物料需求量
|
||||
4. 对比当前库存,计算缺口
|
||||
5. 查询 MaterialSupplier 推荐主供应商
|
||||
"""
|
||||
# ── 1. 查询销售订单 ──
|
||||
order_result = await db_session.execute(
|
||||
select(SalesOrder).where(SalesOrder.id.in_(sales_order_ids))
|
||||
)
|
||||
orders = order_result.scalars().all()
|
||||
if not orders:
|
||||
raise HTTPException(status_code=404, detail="未找到有效的销售订单")
|
||||
|
||||
order_ids_found = [o.id for o in orders]
|
||||
order_nos = [o.order_no for o in orders]
|
||||
|
||||
# ── 2. 查询订单明细(成品列表) ──
|
||||
item_result = await db_session.execute(
|
||||
select(SalesOrderItem).where(SalesOrderItem.order_id.in_(order_ids_found))
|
||||
)
|
||||
order_items = item_result.scalars().all()
|
||||
if not order_items:
|
||||
return PurchaseDemandResponse(
|
||||
source_order_ids=order_ids_found,
|
||||
source_order_nos=order_nos,
|
||||
)
|
||||
|
||||
# ── 3. 按 BOM 展开物料需求 ──
|
||||
finished_ids = list({int(i.product_id) for i in order_items})
|
||||
bom_result = await db_session.execute(
|
||||
select(ProductMaterial, Product)
|
||||
.join(Product, ProductMaterial.material_product_id == Product.id)
|
||||
.where(ProductMaterial.finished_product_id.in_(finished_ids))
|
||||
.where(Product.is_active == True)
|
||||
.where(Product.item_type == "material")
|
||||
)
|
||||
bom_rows = bom_result.all()
|
||||
if not bom_rows:
|
||||
return PurchaseDemandResponse(
|
||||
source_order_ids=order_ids_found,
|
||||
source_order_nos=order_nos,
|
||||
)
|
||||
|
||||
# 按 finished_product_id 分组 BOM
|
||||
bom_by_finished: dict = {}
|
||||
for bom, material in bom_rows:
|
||||
bom_by_finished.setdefault(int(bom.finished_product_id), []).append((bom, material))
|
||||
|
||||
# 聚合需求量:material_id → { material, required_qty }
|
||||
required_qty_map: dict = {}
|
||||
for order_item in order_items:
|
||||
bom_items = bom_by_finished.get(int(order_item.product_id)) or []
|
||||
for bom, material in bom_items:
|
||||
qty = (
|
||||
Decimal(str(order_item.quantity))
|
||||
* Decimal(str(bom.quantity or 0))
|
||||
* (1 + Decimal(str(bom.loss_rate or 0)))
|
||||
)
|
||||
entry = required_qty_map.setdefault(
|
||||
material.id,
|
||||
{"material": material, "required_qty": Decimal("0")},
|
||||
)
|
||||
entry["required_qty"] += qty
|
||||
|
||||
if not required_qty_map:
|
||||
return PurchaseDemandResponse(
|
||||
source_order_ids=order_ids_found,
|
||||
source_order_nos=order_nos,
|
||||
)
|
||||
|
||||
# ── 4. 对比当前库存 ──
|
||||
material_ids = list(required_qty_map.keys())
|
||||
stock_result = await db_session.execute(
|
||||
select(Inventory.product_id, func.coalesce(func.sum(Inventory.quantity), 0))
|
||||
.where(Inventory.product_id.in_(material_ids))
|
||||
.group_by(Inventory.product_id)
|
||||
)
|
||||
stock_map = {row[0]: Decimal(str(row[1] or 0)) for row in stock_result.all()}
|
||||
|
||||
# ── 5. 查询物料-供应商关联(推荐主供应商) ──
|
||||
ms_result = await db_session.execute(
|
||||
select(MaterialSupplier, Supplier)
|
||||
.join(Supplier, MaterialSupplier.supplier_id == Supplier.id)
|
||||
.where(MaterialSupplier.product_id.in_(material_ids))
|
||||
.where(Supplier.is_active == True)
|
||||
.order_by(MaterialSupplier.is_primary.desc(), MaterialSupplier.id.asc())
|
||||
)
|
||||
ms_rows = ms_result.all()
|
||||
# 每个物料取第一个(优先 is_primary=True)
|
||||
supplier_map: dict = {}
|
||||
for ms, supplier in ms_rows:
|
||||
if ms.product_id not in supplier_map:
|
||||
supplier_map[ms.product_id] = {
|
||||
"supplier_id": supplier.id,
|
||||
"supplier_name": supplier.name,
|
||||
"lead_time": ms.lead_time,
|
||||
}
|
||||
|
||||
# ── 6. 组装响应 ──
|
||||
items: List[PurchaseDemandItemResponse] = []
|
||||
total_estimated_cost = Decimal("0")
|
||||
shortage_count = 0
|
||||
|
||||
for material_id, entry in required_qty_map.items():
|
||||
material = entry["material"]
|
||||
required_qty = int(ceil(entry["required_qty"]))
|
||||
available_qty = stock_map.get(material_id, Decimal("0"))
|
||||
shortage_qty = max(required_qty - int(available_qty), 0)
|
||||
unit_cost = Decimal(str(material.cost_price or 0))
|
||||
estimated_cost = Decimal(str(shortage_qty)) * unit_cost
|
||||
total_estimated_cost += estimated_cost
|
||||
|
||||
if shortage_qty > 0:
|
||||
shortage_count += 1
|
||||
|
||||
suggested = supplier_map.get(material_id)
|
||||
items.append(
|
||||
PurchaseDemandItemResponse(
|
||||
material_id=material.id,
|
||||
material_sku=material.sku,
|
||||
material_name=material.name,
|
||||
required_quantity=Decimal(str(required_qty)),
|
||||
available_quantity=available_qty,
|
||||
shortage_quantity=Decimal(str(shortage_qty)),
|
||||
unit_cost=unit_cost,
|
||||
estimated_cost=estimated_cost,
|
||||
suggested_supplier_id=suggested["supplier_id"] if suggested else None,
|
||||
suggested_supplier_name=suggested["supplier_name"] if suggested else None,
|
||||
supplier_lead_time=suggested["lead_time"] if suggested else None,
|
||||
)
|
||||
)
|
||||
|
||||
# 按缺口数量降序排列(最缺的排最前)
|
||||
items.sort(key=lambda x: (x.shortage_quantity, x.estimated_cost), reverse=True)
|
||||
|
||||
return PurchaseDemandResponse(
|
||||
items=items,
|
||||
total_estimated_cost=total_estimated_cost,
|
||||
shortage_count=shortage_count,
|
||||
source_order_ids=order_ids_found,
|
||||
source_order_nos=order_nos,
|
||||
)
|
||||
|
||||
|
||||
purchase_demand_service = PurchaseDemandService()
|
||||
@@ -476,6 +476,8 @@ class SalesOrderService:
|
||||
current_user: User,
|
||||
) -> SalesOrderResponse:
|
||||
order, customer = await _get_sales_order_with_customer(db_session, order_id)
|
||||
if order.status == "delivered":
|
||||
raise HTTPException(status_code=400, detail="已交付的销售订单禁止修改")
|
||||
if order.status == "paid":
|
||||
raise HTTPException(status_code=400, detail="已收款的销售订单禁止修改")
|
||||
try:
|
||||
|
||||
+5
-1
@@ -1,4 +1,8 @@
|
||||
# main.py
|
||||
# main.py — 已废弃,保留向后兼容
|
||||
# 推荐使用入口:
|
||||
# - src/entrypoints/moldinsight.py (模具分析服务)
|
||||
# - src/entrypoints/inventory.py (进销存服务)
|
||||
# 两者均基于 shared.app_factory.create_app() 构建,消除重复代码。
|
||||
import sys
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
@@ -21,6 +21,7 @@ def _safe_include(module_path: str, label: str):
|
||||
|
||||
_safe_include("moldinsight.api.health_router", "健康检查")
|
||||
_safe_include("moldinsight.api.upload_router", "上传")
|
||||
_safe_include("moldinsight.api.batch_router", "批量")
|
||||
_safe_include("moldinsight.api.task_router", "任务")
|
||||
_safe_include("moldinsight.api.history_router", "历史")
|
||||
_safe_include("moldinsight.api.debug_router", "调试")
|
||||
|
||||
@@ -307,7 +307,7 @@ async def estimate_cost(
|
||||
request: Request,
|
||||
current_user: User = Depends(get_current_active_user),
|
||||
):
|
||||
"""LLM 模具成本估算(P2-2:真 AI 落地,需启用 LLM)"""
|
||||
"""模具成本估算:优先使用 LLM,未启用时降级为规则式估算"""
|
||||
body = await request.json()
|
||||
task_id = body.get("task_id")
|
||||
if not task_id:
|
||||
@@ -323,11 +323,17 @@ async def estimate_cost(
|
||||
"geometry_data": task_data.get("geometry_data", {}),
|
||||
"metadata": {"selected_material": task_data.get("material")},
|
||||
}
|
||||
# 优先使用 LLM
|
||||
from moldinsight.services.llm_service import llm_service
|
||||
result = await llm_service.estimate_cost(analysis_result, detailed_context)
|
||||
if result is None:
|
||||
raise HTTPException(503, "成本估算不可用(LLM 未启用或生成失败)")
|
||||
return {"status": "success", "data": result}
|
||||
if result is not None:
|
||||
result["source"] = "ai"
|
||||
return {"status": "success", "data": result}
|
||||
|
||||
# LLM 未启用或失败,降级为规则估算
|
||||
from moldinsight.services.cost_estimate_service import estimate_cost_by_rules
|
||||
rules_result = estimate_cost_by_rules(analysis_result, detailed_context)
|
||||
return {"status": "success", "data": rules_result}
|
||||
|
||||
|
||||
@router.post("/design-cam")
|
||||
|
||||
@@ -0,0 +1,226 @@
|
||||
"""
|
||||
moldinsight/api/batch_router.py — 批量分析端点
|
||||
|
||||
- POST /api/batch-upload 批量上传多文件,返回 batch_id + 各 task_id
|
||||
- GET /api/batch/{batch_id} 聚合查询批量任务进度
|
||||
"""
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
from typing import List, Dict, Any
|
||||
|
||||
from fastapi import APIRouter, UploadFile, File, Form, HTTPException, Depends
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from shared.database.database import get_db_session
|
||||
from shared.services.auth_service import get_current_active_user
|
||||
from shared.models.database import User
|
||||
from shared.models.schemas import ProcessingStatus, create_task_info
|
||||
from shared.services.redis_task_manager import redis_task_manager
|
||||
from shared.utils.file_handler import FileHandler
|
||||
from shared.utils.logger import get_logger
|
||||
from moldinsight.services.storage_integration_rustfs import StorageIntegrationService
|
||||
|
||||
try:
|
||||
from celery_tasks import process_stp_task
|
||||
_use_celery = True
|
||||
except ImportError:
|
||||
process_stp_task = None
|
||||
_use_celery = False
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
file_handler = FileHandler()
|
||||
|
||||
# ─── 批量元数据 Redis key 约定 ──────────────────────────────────────
|
||||
_BATCH_KEY_PREFIX = "batch:"
|
||||
_BATCH_TTL = 86400 # 24h
|
||||
|
||||
|
||||
def _batch_redis_key(batch_id: str) -> str:
|
||||
return f"{_BATCH_KEY_PREFIX}{batch_id}"
|
||||
|
||||
|
||||
@router.post("/batch-upload")
|
||||
async def batch_upload(
|
||||
files: List[UploadFile] = File(...),
|
||||
material: str = Form("ABS"),
|
||||
draft_angle: float = Form(2.0),
|
||||
shrinkage_rate: float = Form(0.5),
|
||||
parting_precision: float = Form(0.1),
|
||||
cavity_match: int = Form(95),
|
||||
db_session: AsyncSession = Depends(get_db_session),
|
||||
current_user: User = Depends(get_current_active_user),
|
||||
):
|
||||
"""批量上传多个 STP 文件,每个文件创建独立分析任务,用 batch_id 聚合。"""
|
||||
if not files:
|
||||
raise HTTPException(400, "请至少上传一个文件")
|
||||
if len(files) > 20:
|
||||
raise HTTPException(400, "单次批量上传最多 20 个文件")
|
||||
|
||||
process_params = {
|
||||
"material": material,
|
||||
"draft_angle": float(draft_angle),
|
||||
"shrinkage_rate": float(shrinkage_rate),
|
||||
"parting_precision": float(parting_precision),
|
||||
"cavity_match": int(cavity_match),
|
||||
}
|
||||
|
||||
batch_id = str(uuid.uuid4())
|
||||
tasks: List[Dict[str, Any]] = []
|
||||
storage_service = StorageIntegrationService()
|
||||
|
||||
for file in files:
|
||||
# 文件类型检查
|
||||
if not file.filename.lower().endswith(('.stp', '.step')):
|
||||
tasks.append({
|
||||
"filename": file.filename,
|
||||
"task_id": None,
|
||||
"status": "rejected",
|
||||
"error": "不支持的文件类型",
|
||||
})
|
||||
continue
|
||||
|
||||
task_id = str(uuid.uuid4())
|
||||
try:
|
||||
file_path, file_size, file_meta = await file_handler.save_uploaded_file(file)
|
||||
|
||||
stp_file = await storage_service.save_stp_file(
|
||||
session=db_session,
|
||||
file_path=file_path,
|
||||
original_filename=file_meta["safe_original_name"],
|
||||
user_id=current_user.id,
|
||||
)
|
||||
|
||||
await storage_service.create_processing_task(
|
||||
db_session, task_id, stp_file.id, parameters=process_params,
|
||||
)
|
||||
|
||||
task_info = create_task_info(
|
||||
task_id=task_id,
|
||||
status=ProcessingStatus.PROCESSING,
|
||||
filename=file.filename,
|
||||
file_path=str(file_path),
|
||||
file_size=file_size,
|
||||
upload_time=str(datetime.now()),
|
||||
)
|
||||
task_info["material"] = material
|
||||
task_info["parameters"] = process_params
|
||||
task_info["batch_id"] = batch_id
|
||||
await redis_task_manager.set_task(task_id, task_info)
|
||||
|
||||
# 调度处理
|
||||
if _use_celery:
|
||||
process_stp_task.delay(task_id, str(file_path), stp_file.id, process_params)
|
||||
else:
|
||||
import asyncio
|
||||
from moldinsight.services.processing_service import processing_service
|
||||
asyncio.create_task(processing_service.process_file_with_storage(
|
||||
task_id, str(file_path), stp_file.id, process_params
|
||||
))
|
||||
|
||||
tasks.append({
|
||||
"filename": file.filename,
|
||||
"task_id": task_id,
|
||||
"status": "processing",
|
||||
"stp_file_id": stp_file.id,
|
||||
})
|
||||
logger.info(
|
||||
f"[BATCH] batch_id={batch_id} task_id={task_id} "
|
||||
f"file={file.filename} user={current_user.username}"
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.warning(f"[BATCH] 文件 {file.filename} 上传失败: {exc}")
|
||||
tasks.append({
|
||||
"filename": file.filename,
|
||||
"task_id": task_id,
|
||||
"status": "error",
|
||||
"error": str(exc),
|
||||
})
|
||||
|
||||
# 将 batch 元数据写入 Redis
|
||||
batch_meta = {
|
||||
"batch_id": batch_id,
|
||||
"user_id": current_user.id,
|
||||
"created_at": str(datetime.now()),
|
||||
"task_ids": [t["task_id"] for t in tasks if t.get("task_id")],
|
||||
"total": len(tasks),
|
||||
"params": process_params,
|
||||
}
|
||||
await redis_task_manager.redis_client.set(
|
||||
_batch_redis_key(batch_id),
|
||||
__import__("json").dumps(batch_meta),
|
||||
ex=_BATCH_TTL,
|
||||
)
|
||||
|
||||
return {
|
||||
"batch_id": batch_id,
|
||||
"total": len(tasks),
|
||||
"accepted": sum(1 for t in tasks if t.get("status") != "rejected"),
|
||||
"tasks": tasks,
|
||||
}
|
||||
|
||||
|
||||
@router.get("/batch/{batch_id}")
|
||||
async def get_batch_status(
|
||||
batch_id: str,
|
||||
current_user: User = Depends(get_current_active_user),
|
||||
):
|
||||
"""聚合查询批量任务进度"""
|
||||
import json
|
||||
|
||||
raw = await redis_task_manager.redis_client.get(_batch_redis_key(batch_id))
|
||||
if not raw:
|
||||
raise HTTPException(404, "批量任务不存在或已过期")
|
||||
|
||||
batch_meta = json.loads(raw)
|
||||
|
||||
# 权限检查
|
||||
if batch_meta.get("user_id") and batch_meta["user_id"] != current_user.id:
|
||||
raise HTTPException(403, "无权访问该批量任务")
|
||||
|
||||
task_ids = batch_meta.get("task_ids", [])
|
||||
task_statuses = []
|
||||
completed = 0
|
||||
failed = 0
|
||||
processing = 0
|
||||
|
||||
for tid in task_ids:
|
||||
task_data = await redis_task_manager.get_task(tid)
|
||||
if not task_data:
|
||||
task_statuses.append({"task_id": tid, "status": "unknown"})
|
||||
continue
|
||||
status = task_data.get("status", "unknown")
|
||||
progress = task_data.get("progress", 0)
|
||||
filename = task_data.get("filename", "")
|
||||
error = task_data.get("error", "")
|
||||
html_file = task_data.get("html_file", "")
|
||||
|
||||
if status == ProcessingStatus.COMPLETED:
|
||||
completed += 1
|
||||
elif status == ProcessingStatus.FAILED:
|
||||
failed += 1
|
||||
else:
|
||||
processing += 1
|
||||
|
||||
task_statuses.append({
|
||||
"task_id": tid,
|
||||
"status": status,
|
||||
"progress": progress,
|
||||
"filename": filename,
|
||||
"error": error,
|
||||
"html_file": html_file,
|
||||
})
|
||||
|
||||
total = len(task_ids)
|
||||
return {
|
||||
"batch_id": batch_id,
|
||||
"created_at": batch_meta.get("created_at"),
|
||||
"total": total,
|
||||
"completed": completed,
|
||||
"failed": failed,
|
||||
"processing": processing,
|
||||
"progress_percent": round((completed + failed) / max(total, 1) * 100, 1),
|
||||
"tasks": task_statuses,
|
||||
}
|
||||
@@ -34,46 +34,3 @@ async def get_status(task_id: str, db_session: AsyncSession = Depends(get_db_ses
|
||||
except Exception as e:
|
||||
logger.error(f"获取任务状态失败: {e}")
|
||||
raise HTTPException(500, f"获取任务状态失败: {str(e)}")
|
||||
|
||||
|
||||
@router.get("/result/{task_id}")
|
||||
@router.post("/result/{task_id}")
|
||||
async def result_page(request: Request, task_id: str, db_session: AsyncSession = Depends(get_db_session)):
|
||||
"""结果详情页面"""
|
||||
# 从数据库查询任务详情
|
||||
result = await db_session.execute(
|
||||
select(ProcessingTask, STPFile)
|
||||
.join(STPFile, ProcessingTask.stp_file_id == STPFile.id)
|
||||
.where(ProcessingTask.task_id == task_id)
|
||||
)
|
||||
|
||||
task_record = result.first()
|
||||
|
||||
if not task_record:
|
||||
raise HTTPException(404, "任务不存在")
|
||||
|
||||
task, stp_file = task_record
|
||||
|
||||
# 构建任务详情数据
|
||||
task_data = {
|
||||
"task_id": task.task_id,
|
||||
"filename": stp_file.original_filename if stp_file else "",
|
||||
"file_size": stp_file.file_size if stp_file else 0,
|
||||
"status": task.status,
|
||||
"progress": task.progress,
|
||||
"current_step": task.current_step,
|
||||
"created_at": task.created_time.isoformat() if task.created_time else "",
|
||||
"completed_at": task.completed_time.isoformat() if task.completed_time else "",
|
||||
"error": task.error_message if task.error_message else ""
|
||||
}
|
||||
|
||||
from fastapi.templating import Jinja2Templates
|
||||
import os
|
||||
templates_dir = os.path.join(os.getcwd(), "templates")
|
||||
templates = Jinja2Templates(directory=templates_dir)
|
||||
return templates.TemplateResponse("result.html", {
|
||||
"request": request,
|
||||
"task": task_data,
|
||||
"pythonocc_available": True,
|
||||
"version": "3.0.0"
|
||||
})
|
||||
|
||||
@@ -0,0 +1,237 @@
|
||||
"""
|
||||
moldinsight/services/cost_estimate_service.py — 规则式成本估算
|
||||
|
||||
当 LLM 未启用时作为兜底,基于几何参数 + 材料库 + 模具尺寸
|
||||
计算材料费 / 加工费,完全不依赖 LLM。
|
||||
"""
|
||||
|
||||
from typing import Dict, Any, Optional
|
||||
|
||||
|
||||
# ─── 模具钢材料单价参考(元/kg,含税) ──────────────────────────────
|
||||
_MOLD_STEEL_PRICE = {
|
||||
"铝合金7075": 45,
|
||||
"P20": 25,
|
||||
"718H": 35,
|
||||
"NAK80": 55,
|
||||
"S136": 70,
|
||||
"H13": 40,
|
||||
"default": 30,
|
||||
}
|
||||
|
||||
# ─── 加工复杂度系数 ────────────────────────────────────────────────────
|
||||
_COMPLEXITY_FACTOR = {
|
||||
"low": 1.0,
|
||||
"medium": 1.3,
|
||||
"high": 1.7,
|
||||
"very_high": 2.2,
|
||||
}
|
||||
|
||||
# 侧向机构附加费用(元/个)
|
||||
_SIDE_ACTION_COST = {
|
||||
"slider": 8000, # 滑块
|
||||
"lifter": 6000, # 斜顶
|
||||
"mixed": 7000, # 混合
|
||||
}
|
||||
|
||||
# 型腔加工基础费用(元/型腔)
|
||||
_CAVITY_MACHINING_BASE = 25000
|
||||
|
||||
# 模架基础费用(元)
|
||||
_BASE_MOLD_FRAME = {
|
||||
"small": 15000, # 长宽 < 250mm
|
||||
"medium": 25000, # 长宽 250-400mm
|
||||
"large": 45000, # 长宽 > 400mm
|
||||
}
|
||||
|
||||
|
||||
def estimate_cost_by_rules(
|
||||
analysis_result: Dict[str, Any],
|
||||
detailed_cavity_json: Optional[Dict[str, Any]] = None,
|
||||
) -> Dict[str, Any]:
|
||||
"""规则式模具成本估算。
|
||||
|
||||
返回结构与 LLM estimate_cost 一致,前端可无缝展示。
|
||||
"""
|
||||
ctx = _extract_context(analysis_result, detailed_cavity_json)
|
||||
|
||||
# ── 模具材料费 ──────────────────────────────────────────────
|
||||
mold_weight_kg = _estimate_mold_weight(ctx)
|
||||
steel_price = _MOLD_STEEL_PRICE.get(ctx["mold_material"], _MOLD_STEEL_PRICE["default"])
|
||||
material_cost = int(mold_weight_kg * steel_price)
|
||||
|
||||
# ── 加工费 ─────────────────────────────────────────────────
|
||||
cavity_count = ctx["cavity_count"]
|
||||
complexity_key = ctx["complexity"]
|
||||
complexity_factor = _COMPLEXITY_FACTOR.get(complexity_key, 1.3)
|
||||
machining_cost = int(
|
||||
_CAVITY_MACHINING_BASE * cavity_count * complexity_factor
|
||||
+ _BASE_MOLD_FRAME[ctx["mold_frame_size"]]
|
||||
)
|
||||
|
||||
# ── 侧向机构附加费 ─────────────────────────────────────────
|
||||
side_action_extra = 0
|
||||
side_action_parts = []
|
||||
slider_count = ctx["slider_count"]
|
||||
lifter_count = ctx["lifter_count"]
|
||||
if slider_count > 0:
|
||||
side_action_extra += slider_count * _SIDE_ACTION_COST["slider"]
|
||||
side_action_parts.append(f"{slider_count} 个滑块")
|
||||
if lifter_count > 0:
|
||||
side_action_extra += lifter_count * _SIDE_ACTION_COST["lifter"]
|
||||
side_action_parts.append(f"{lifter_count} 个斜顶")
|
||||
|
||||
complexity_label = (
|
||||
f"{complexity_factor:.1f}"
|
||||
+ (f"(含 {', '.join(side_action_parts)})" if side_action_parts else "")
|
||||
)
|
||||
|
||||
# ── 合计 ───────────────────────────────────────────────────
|
||||
mold_subtotal = material_cost + machining_cost + side_action_extra
|
||||
|
||||
# ── 单件成本 ───────────────────────────────────────────────
|
||||
part_weight_g = ctx["part_weight_g"]
|
||||
cycle_time_s = ctx["cycle_time_s"]
|
||||
# 材料费:塑料粒 ~30 元/kg 均值
|
||||
material_price_per_kg = 30
|
||||
part_material_cost = (part_weight_g / 1000) * material_price_per_kg
|
||||
# 机时分摊:假设机时费 60 元/h
|
||||
machine_hourly_rate = 60
|
||||
part_cycle_cost = (cycle_time_s / 3600) * machine_hourly_rate * (1 / max(cavity_count, 1))
|
||||
# 人工 + 能耗分摊 ~15%
|
||||
part_overhead = (part_material_cost + part_cycle_cost) * 0.15
|
||||
cost_per_part = round(part_material_cost + part_cycle_cost + part_overhead, 2)
|
||||
|
||||
return {
|
||||
"mold_cost": {
|
||||
"material": f"¥{material_cost:,}({ctx['mold_material']},约 {mold_weight_kg:.0f} kg)",
|
||||
"machining": f"¥{machining_cost:,}(含 CNC/EDM/线切割,{cavity_count} 腔)",
|
||||
"complexity_factor": complexity_label,
|
||||
"subtotal": f"¥{mold_subtotal:,}",
|
||||
},
|
||||
"part_cost": {
|
||||
"material": f"¥{part_material_cost:.2f}({ctx['material_name']},约 {part_weight_g:.1f} g)",
|
||||
"cycle_time": f"{cycle_time_s} s",
|
||||
"cost_per_part": f"¥{cost_per_part:.2f}",
|
||||
},
|
||||
"total_mold_cost": f"¥{mold_subtotal:,}",
|
||||
"cost_per_part": f"¥{cost_per_part:.2f}",
|
||||
"confidence": 0.55,
|
||||
"assumptions": [
|
||||
"假设模具寿命 50 万模次",
|
||||
f"模具钢:{ctx['mold_material']}({steel_price} 元/kg)",
|
||||
f"型腔数:{cavity_count}",
|
||||
f"机时费:{machine_hourly_rate} 元/h",
|
||||
"塑料粒均价 30 元/kg",
|
||||
"人工+能耗分摊 15%",
|
||||
"规则估算,仅供参考",
|
||||
],
|
||||
"source": "rules",
|
||||
}
|
||||
|
||||
|
||||
def _extract_context(
|
||||
analysis_result: Dict[str, Any],
|
||||
detailed_cavity_json: Optional[Dict[str, Any]],
|
||||
) -> Dict[str, Any]:
|
||||
"""从分析结果提取成本估算所需上下文"""
|
||||
geometry = analysis_result.get("geometry_data", {})
|
||||
bbox = geometry.get("bounding_box", {})
|
||||
dims = bbox.get("dimensions", [0, 0, 0])
|
||||
volume_mm3 = geometry.get("volume", 0) or 0
|
||||
|
||||
schemes = (detailed_cavity_json or {}).get("candidate_schemes", [])
|
||||
best = schemes[0] if schemes else {}
|
||||
cavity_data = best.get("cavity_data", {}) if isinstance(best, dict) else {}
|
||||
mfg_info = cavity_data.get("manufacturing_info", {})
|
||||
metadata = cavity_data.get("metadata", {})
|
||||
|
||||
# 型腔数
|
||||
cavity_count = (cavity_data.get("mold_cavities", {}) or {}).get("cavity_count", 1)
|
||||
|
||||
# 模具材料
|
||||
mold_material = mfg_info.get("mold_material", "P20")
|
||||
|
||||
# 模具尺寸
|
||||
mold_size = mfg_info.get("estimated_mold_size", {})
|
||||
length = mold_size.get("length", 300)
|
||||
width = mold_size.get("width", 300)
|
||||
max_dim = max(length, width)
|
||||
if max_dim < 250:
|
||||
mold_frame_size = "small"
|
||||
elif max_dim < 400:
|
||||
mold_frame_size = "medium"
|
||||
else:
|
||||
mold_frame_size = "large"
|
||||
|
||||
# 复杂度
|
||||
side_actions = cavity_data.get("side_actions", {}) or {}
|
||||
summary = side_actions.get("summary", {})
|
||||
slider_count = summary.get("total_slider_count", 0) or 0
|
||||
lifter_count = summary.get("total_lifter_count", 0) or 0
|
||||
total_mechanism = slider_count + lifter_count
|
||||
if total_mechanism == 0:
|
||||
complexity = "low"
|
||||
elif total_mechanism <= 2:
|
||||
complexity = "medium"
|
||||
elif total_mechanism <= 4:
|
||||
complexity = "high"
|
||||
else:
|
||||
complexity = "very_high"
|
||||
|
||||
# 产品重量
|
||||
material_name = metadata.get("selected_material", "ABS")
|
||||
density = 1.05 # ABS 默认密度
|
||||
part_weight_g = (volume_mm3 / 1000) * density
|
||||
|
||||
# 成型周期
|
||||
cycle_time_raw = mfg_info.get("estimated_cycle_time", "30")
|
||||
try:
|
||||
cycle_time_s = int(str(cycle_time_raw).replace("秒", "").strip())
|
||||
except (ValueError, TypeError):
|
||||
cycle_time_s = 30
|
||||
|
||||
return {
|
||||
"dims": dims,
|
||||
"volume_mm3": volume_mm3,
|
||||
"cavity_count": cavity_count,
|
||||
"mold_material": mold_material,
|
||||
"mold_frame_size": mold_frame_size,
|
||||
"complexity": complexity,
|
||||
"slider_count": slider_count,
|
||||
"lifter_count": lifter_count,
|
||||
"material_name": material_name,
|
||||
"part_weight_g": part_weight_g,
|
||||
"cycle_time_s": cycle_time_s,
|
||||
}
|
||||
|
||||
|
||||
def _estimate_mold_weight(ctx: Dict[str, Any]) -> float:
|
||||
"""基于模具尺寸估算重量(kg),假设钢材密度 7.85 g/cm³"""
|
||||
cavity_count = ctx["cavity_count"]
|
||||
dims = ctx["dims"]
|
||||
dim_x = max(dims[0] if len(dims) > 0 else 120, 120)
|
||||
dim_y = max(dims[1] if len(dims) > 1 else 100, 100)
|
||||
dim_z = max(dims[2] if len(dims) > 2 else 60, 60)
|
||||
|
||||
edge_margin = 50
|
||||
if cavity_count == 1:
|
||||
length = dim_x + 2 * edge_margin
|
||||
width = dim_y + 2 * edge_margin
|
||||
elif cavity_count == 2:
|
||||
length = 2 * dim_x + 30 + 2 * edge_margin
|
||||
width = dim_y + 2 * edge_margin
|
||||
elif cavity_count == 4:
|
||||
length = 2 * dim_x + 30 + 2 * edge_margin
|
||||
width = 2 * dim_y + 30 + 2 * edge_margin
|
||||
else:
|
||||
length = 4 * dim_x + 90 + 2 * edge_margin
|
||||
width = 2 * dim_y + 30 + 2 * edge_margin
|
||||
|
||||
height = dim_z + 80 # 含冷却系统
|
||||
|
||||
# 体积 mm³ → cm³,再乘钢材密度 7.85 g/cm³,再转 kg
|
||||
# 模架不是实心钢块,取 40% 填充率
|
||||
volume_cm3 = (length * width * height) / 1000
|
||||
weight_kg = volume_cm3 * 7.85 * 0.40 / 1000
|
||||
return max(weight_kg, 50) # 最小 50 kg
|
||||
@@ -13,9 +13,6 @@ from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from moldinsight.core.stp_parser import STPParser
|
||||
from moldinsight.core.geometry_analyzer import GeometryAnalyzer
|
||||
from moldinsight.core.mold_generator import MoldCavityGenerator
|
||||
from moldinsight.core.aluminum_foam_mold import AluminumFoamMoldGenerator
|
||||
from moldinsight.core.mold_quality_inspector import AluminumFoamMoldQualityInspector
|
||||
from moldinsight.core.mesh_generator import MeshGenerator
|
||||
from moldinsight.core.multi_scheme_planner import MultiSchemeMoldPlanner
|
||||
from moldinsight.core.cad_exporter import CADExporter
|
||||
@@ -38,9 +35,6 @@ class ProcessingService:
|
||||
def __init__(self):
|
||||
self.stp_parser = STPParser()
|
||||
self.geometry_analyzer = GeometryAnalyzer()
|
||||
self.mold_generator = MoldCavityGenerator(shrinkage_rate=0.005)
|
||||
self.aluminum_foam_generator = AluminumFoamMoldGenerator(shrinkage_rate=0.015, draft_angle=3.0)
|
||||
self.mold_quality_inspector = AluminumFoamMoldQualityInspector()
|
||||
self.mesh_generator = MeshGenerator(quality="medium")
|
||||
self.html_generator = HTMLGenerator()
|
||||
self.storage_service = StorageIntegrationService()
|
||||
|
||||
@@ -0,0 +1,212 @@
|
||||
"""
|
||||
shared/app_factory.py — FastAPI 应用工厂
|
||||
|
||||
将 moldinsight.py / inventory.py 两个入口的重复引导代码
|
||||
(CORS、日志中间件、startup/shutdown、/health、SPA fallback)
|
||||
收敛到一个工厂函数,消除漂移风险。
|
||||
"""
|
||||
import os
|
||||
import time
|
||||
from pathlib import Path
|
||||
from typing import List, Optional, Callable, Awaitable
|
||||
|
||||
from fastapi import FastAPI, Request
|
||||
from fastapi.staticfiles import StaticFiles
|
||||
from fastapi.middleware.cors import CORSMiddleware
|
||||
from fastapi.responses import JSONResponse, FileResponse
|
||||
|
||||
from shared.config.settings import settings
|
||||
from shared.utils.logger import setup_logging, get_logger, generate_request_id, set_request_id
|
||||
|
||||
setup_logging()
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
def create_app(
|
||||
*,
|
||||
title: str,
|
||||
service_name: str,
|
||||
version: str = "4.0.0",
|
||||
mount_html: bool = False,
|
||||
startup_hooks: Optional[List[Callable[[], Awaitable[None]]]] = None,
|
||||
register_routers: Optional[Callable[[FastAPI], None]] = None,
|
||||
) -> FastAPI:
|
||||
"""创建标准化的 FastAPI 应用实例。
|
||||
|
||||
Args:
|
||||
title: 应用标题
|
||||
service_name: 服务名(用于 /health 响应)
|
||||
version: 版本号
|
||||
mount_html: 是否挂载 /html 静态目录(moldinsight 需要)
|
||||
startup_hooks: 额外的 startup 钩子列表(在数据库/RustFS/Redis 初始化后执行)
|
||||
register_routers: 回调函数,用于注册业务路由
|
||||
"""
|
||||
app = FastAPI(title=title, version=version)
|
||||
|
||||
# ── CORS 白名单 ──────────────────────────────────────────────
|
||||
cors_origins = settings.CORS_ORIGINS or ["*"]
|
||||
if cors_origins == ["*"]:
|
||||
logger.warning(
|
||||
"CORS 使用通配符 ['*'],生产环境请设置 CORS_ORIGINS 环境变量"
|
||||
)
|
||||
app.add_middleware(
|
||||
CORSMiddleware,
|
||||
allow_origins=cors_origins,
|
||||
allow_credentials=True,
|
||||
allow_methods=["*"],
|
||||
allow_headers=["*"],
|
||||
)
|
||||
|
||||
# ── 请求日志中间件(结构化 + request_id 追踪)─────────────────
|
||||
@app.middleware("http")
|
||||
async def log_requests(request: Request, call_next):
|
||||
# 生成/提取 request_id
|
||||
rid = request.headers.get("X-Request-ID") or generate_request_id()
|
||||
set_request_id(rid)
|
||||
|
||||
start_time = time.time()
|
||||
response = await call_next(request)
|
||||
duration_ms = round((time.time() - start_time) * 1000, 1)
|
||||
|
||||
# 跳过静态资源和健康检查的详细日志
|
||||
path = request.url.path
|
||||
is_static = path.startswith("/static") or path == "/health"
|
||||
|
||||
if not is_static:
|
||||
log_level = "warning" if response.status_code >= 400 else "info"
|
||||
extra = {
|
||||
"method": request.method,
|
||||
"path": path,
|
||||
"status": response.status_code,
|
||||
"duration_ms": duration_ms,
|
||||
"client_ip": request.client.host if request.client else "-",
|
||||
}
|
||||
getattr(logger, log_level)(
|
||||
f"{request.method} {path} -> {response.status_code} ({duration_ms}ms)",
|
||||
extra=extra,
|
||||
)
|
||||
|
||||
# 注入 X-Request-ID 响应头,方便前端/运维追踪
|
||||
response.headers["X-Request-ID"] = rid
|
||||
return response
|
||||
|
||||
# ── 目录准备 ─────────────────────────────────────────────────
|
||||
Path("uploads").mkdir(exist_ok=True)
|
||||
Path("static").mkdir(exist_ok=True)
|
||||
if mount_html:
|
||||
Path("html_output").mkdir(exist_ok=True)
|
||||
|
||||
# ── 静态文件挂载 ─────────────────────────────────────────────
|
||||
app.mount(
|
||||
"/static",
|
||||
StaticFiles(directory=os.path.join(os.getcwd(), "static")),
|
||||
name="static",
|
||||
)
|
||||
if mount_html:
|
||||
app.mount(
|
||||
"/html",
|
||||
StaticFiles(directory=os.path.join(os.getcwd(), "html_output")),
|
||||
name="html",
|
||||
)
|
||||
|
||||
# ── Startup ──────────────────────────────────────────────────
|
||||
@app.on_event("startup")
|
||||
async def startup_event():
|
||||
from shared.database.init_db import init_database
|
||||
|
||||
success = await init_database(keep_connected=True)
|
||||
print(f"[{'OK' if success else 'FAIL'}] 数据库初始化")
|
||||
|
||||
# RustFS(仅 moldinsight 需要)
|
||||
if mount_html:
|
||||
try:
|
||||
from moldinsight.storage.rustfs_storage import rustfs_manager
|
||||
await rustfs_manager.connect(
|
||||
endpoint=settings.RUSTFS_ENDPOINT,
|
||||
access_key=settings.RUSTFS_ACCESS_KEY,
|
||||
secret_key=settings.RUSTFS_SECRET_KEY,
|
||||
timeout=settings.RUSTFS_TIMEOUT,
|
||||
)
|
||||
print("[OK] RustFS连接成功")
|
||||
except Exception as e:
|
||||
print(f"[WARN] RustFS连接失败: {e}")
|
||||
|
||||
# Redis
|
||||
try:
|
||||
from shared.services.redis_task_manager import redis_task_manager
|
||||
await redis_task_manager.connect()
|
||||
print(f"[{'OK' if redis_task_manager.is_connected else 'WARN'}] Redis")
|
||||
except Exception as e:
|
||||
print(f"[WARN] Redis异常: {e}")
|
||||
|
||||
# 额外钩子
|
||||
for hook in (startup_hooks or []):
|
||||
try:
|
||||
await hook()
|
||||
except Exception as e:
|
||||
print(f"[WARN] startup hook 异常: {e}")
|
||||
|
||||
# ── Shutdown ─────────────────────────────────────────────────
|
||||
@app.on_event("shutdown")
|
||||
async def shutdown_event():
|
||||
try:
|
||||
from shared.services.redis_task_manager import redis_task_manager
|
||||
await redis_task_manager.disconnect()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# ── 认证路由 ─────────────────────────────────────────────────
|
||||
from shared.services.auth_routes import router as auth_router
|
||||
app.include_router(auth_router)
|
||||
|
||||
# ── 业务路由注册 ─────────────────────────────────────────────
|
||||
if register_routers:
|
||||
register_routers(app)
|
||||
|
||||
# ── /health 统一端点 ─────────────────────────────────────────
|
||||
@app.get("/health")
|
||||
@app.post("/health")
|
||||
async def health():
|
||||
from shared.database.database import db_manager
|
||||
from sqlalchemy import text
|
||||
|
||||
db_ok = False
|
||||
db_error = None
|
||||
try:
|
||||
if not db_manager.is_connected:
|
||||
await db_manager.connect()
|
||||
async with db_manager.engine.begin() as conn:
|
||||
await conn.execute(text("SELECT 1"))
|
||||
db_ok = True
|
||||
except Exception as e:
|
||||
db_error = str(e)
|
||||
|
||||
return {
|
||||
"status": "healthy" if db_ok else "degraded",
|
||||
"service": service_name,
|
||||
"version": version,
|
||||
"database_connected": db_ok,
|
||||
"database_error": db_error,
|
||||
}
|
||||
|
||||
# ── SPA fallback(排除 /api 前缀,避免吞掉 API 404)────────
|
||||
@app.get("/{full_path:path}")
|
||||
async def spa_fallback(full_path: str):
|
||||
# API 路径不走 SPA fallback,让 FastAPI 正常返回 404 JSON
|
||||
if full_path.startswith("api/") or full_path.startswith("api"):
|
||||
raise _api_not_found(full_path)
|
||||
# 健康检查 / 文档路径也排除
|
||||
if full_path.startswith("docs") or full_path.startswith("openapi"):
|
||||
raise _api_not_found(full_path)
|
||||
static_index = os.path.join(os.getcwd(), "static", "index.html")
|
||||
if os.path.exists(static_index):
|
||||
return FileResponse(static_index)
|
||||
return JSONResponse({"detail": "SPA index not found"}, status_code=404)
|
||||
|
||||
return app
|
||||
|
||||
|
||||
def _api_not_found(path: str):
|
||||
"""为 API 路径生成标准 404 异常"""
|
||||
from fastapi import HTTPException
|
||||
raise HTTPException(status_code=404, detail=f"Not Found: /{path}")
|
||||
@@ -1,6 +1,6 @@
|
||||
import os
|
||||
import urllib.parse
|
||||
from typing import Dict, Any
|
||||
from typing import Dict, Any, List
|
||||
from dotenv import load_dotenv
|
||||
|
||||
load_dotenv()
|
||||
@@ -75,6 +75,11 @@ class Settings:
|
||||
self.REDIS_PASSWORD = os.getenv("REDIS_PASSWORD", "")
|
||||
self.REDIS_DB = int(os.getenv("REDIS_DB", "0"))
|
||||
|
||||
# CORS 白名单(逗号分隔,默认允许本机开发地址)
|
||||
self.CORS_ORIGINS = self._parse_cors_origins(
|
||||
os.getenv("CORS_ORIGINS", "")
|
||||
)
|
||||
|
||||
# LLM 增强分析配置(可选)
|
||||
self.LLM_ENABLED = os.getenv("LLM_ENABLED", "false").lower() == "true"
|
||||
self.LLM_API_URL = os.getenv("LLM_API_URL", "https://api.openai.com/v1")
|
||||
@@ -95,5 +100,14 @@ class Settings:
|
||||
def allowed_extensions_set(self) -> set:
|
||||
return set(ext.strip() for ext in self.ALLOWED_EXTENSIONS.split(","))
|
||||
|
||||
@staticmethod
|
||||
def _parse_cors_origins(raw: str) -> List[str]:
|
||||
"""解析 CORS_ORIGINS 环境变量,逗号分隔。
|
||||
为空时返回空列表(由 app_factory 决定是否降级为 ['*'])。
|
||||
"""
|
||||
if not raw or not raw.strip():
|
||||
return []
|
||||
return [o.strip().rstrip("/") for o in raw.split(",") if o.strip()]
|
||||
|
||||
|
||||
settings = Settings()
|
||||
|
||||
@@ -13,6 +13,24 @@ from shared.utils.logger import get_logger
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
def _get_pool_config(role: str = "web") -> dict:
|
||||
"""按角色返回连接池参数。
|
||||
|
||||
web 入口:适中并发;celery worker:少量长连接。
|
||||
通过 DB_POOL_SIZE / DB_MAX_OVERFLOW 环境变量可覆盖默认值。
|
||||
"""
|
||||
defaults = {
|
||||
"web": {"pool_size": 10, "max_overflow": 20},
|
||||
"celery": {"pool_size": 5, "max_overflow": 10},
|
||||
}
|
||||
role_cfg = defaults.get(role, defaults["web"])
|
||||
# 允许环境变量覆盖
|
||||
pool_size = int(os.getenv("DB_POOL_SIZE", str(role_cfg["pool_size"])))
|
||||
max_overflow = int(os.getenv("DB_MAX_OVERFLOW", str(role_cfg["max_overflow"])))
|
||||
return {"pool_size": pool_size, "max_overflow": max_overflow}
|
||||
|
||||
|
||||
class DatabaseManager:
|
||||
"""数据库管理器"""
|
||||
|
||||
@@ -21,21 +39,27 @@ class DatabaseManager:
|
||||
self.async_session = None
|
||||
self.is_connected = False
|
||||
|
||||
async def connect(self):
|
||||
"""连接数据库"""
|
||||
async def connect(self, role: str = "web"):
|
||||
"""连接数据库
|
||||
|
||||
Args:
|
||||
role: 连接角色,"web" 或 "celery",决定连接池大小
|
||||
"""
|
||||
if not settings.DATABASE_URL:
|
||||
logger.warning("未配置数据库连接,跳过数据库初始化")
|
||||
self.is_connected = False
|
||||
return
|
||||
|
||||
try:
|
||||
pool_cfg = _get_pool_config(role)
|
||||
# 创建异步引擎
|
||||
self.engine = create_async_engine(
|
||||
settings.DATABASE_URL,
|
||||
echo=settings.DEBUG,
|
||||
pool_size=20,
|
||||
max_overflow=30,
|
||||
pool_recycle=3600
|
||||
pool_size=pool_cfg["pool_size"],
|
||||
max_overflow=pool_cfg["max_overflow"],
|
||||
pool_recycle=3600,
|
||||
pool_pre_ping=True, # 自动检测失效连接,避免 PG 断连报错
|
||||
)
|
||||
|
||||
# 创建异步会话工厂
|
||||
@@ -50,7 +74,10 @@ class DatabaseManager:
|
||||
await conn.execute(text("SELECT 1"))
|
||||
|
||||
self.is_connected = True
|
||||
logger.info("数据库连接成功")
|
||||
logger.info(
|
||||
"数据库连接成功 (pool_size=%d, max_overflow=%d)",
|
||||
pool_cfg["pool_size"], pool_cfg["max_overflow"],
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"数据库连接失败: {e}")
|
||||
@@ -66,18 +93,22 @@ class DatabaseManager:
|
||||
|
||||
@asynccontextmanager
|
||||
async def session(self):
|
||||
"""获取数据库会话的异步上下文管理器"""
|
||||
"""获取数据库会话的异步上下文管理器(用于后台任务/Celery)"""
|
||||
if not self.is_connected:
|
||||
await self.connect()
|
||||
|
||||
session = self.async_session()
|
||||
try:
|
||||
yield session
|
||||
await session.commit()
|
||||
except Exception:
|
||||
await session.rollback()
|
||||
raise
|
||||
finally:
|
||||
await session.close()
|
||||
|
||||
async def get_session(self) -> AsyncSession:
|
||||
"""获取数据库会话"""
|
||||
"""获取数据库会话(非上下文管理器,配合 get_db_session 依赖使用)"""
|
||||
if not self.is_connected:
|
||||
await self.connect()
|
||||
|
||||
@@ -100,9 +131,20 @@ db_manager = DatabaseManager()
|
||||
|
||||
# 数据库依赖注入
|
||||
async def get_db_session():
|
||||
"""获取数据库会话的依赖函数"""
|
||||
"""获取数据库会话的依赖函数
|
||||
|
||||
统一事务边界:
|
||||
- 路由正常返回 → 自动 commit
|
||||
- 路由抛出异常 → 自动 rollback
|
||||
|
||||
路由中应使用 flush() 代替 commit(),以便在提交前仍能 refresh()。
|
||||
"""
|
||||
session = await db_manager.get_session()
|
||||
try:
|
||||
yield session
|
||||
await session.commit()
|
||||
except Exception:
|
||||
await session.rollback()
|
||||
raise
|
||||
finally:
|
||||
await session.close()
|
||||
await session.close()
|
||||
|
||||
+121
-13
@@ -1,17 +1,125 @@
|
||||
# utils/logger.py
|
||||
"""
|
||||
shared/utils/logger.py — 结构化日志 + 请求追踪
|
||||
|
||||
功能:
|
||||
- JSON 结构化日志输出(生产友好,方便 ELK/Loki 采集)
|
||||
- request_id 自动注入(通过 contextvars,跨 async 传播)
|
||||
- 向后兼容:get_logger(name) / setup_logging() API 不变
|
||||
- 支持 LOG_FORMAT 环境变量切换(json / text,默认 json)
|
||||
"""
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import sys
|
||||
import uuid
|
||||
from contextvars import ContextVar
|
||||
from datetime import datetime, timezone
|
||||
from typing import Optional
|
||||
|
||||
def setup_logging():
|
||||
"""设置日志配置"""
|
||||
logging.basicConfig(
|
||||
level=logging.INFO,
|
||||
format='%(asctime)s - %(name)s - %(levelname)s - %(message)s',
|
||||
handlers=[
|
||||
logging.StreamHandler(sys.stdout)
|
||||
]
|
||||
)
|
||||
# ── request_id 上下文变量(跨 async task 自动传播)────────────
|
||||
request_id_var: ContextVar[Optional[str]] = ContextVar("request_id", default=None)
|
||||
|
||||
def get_logger(name: str):
|
||||
"""获取日志器"""
|
||||
return logging.getLogger(name)
|
||||
|
||||
def generate_request_id() -> str:
|
||||
"""生成短 request_id(8 位 hex,便于日志阅读)"""
|
||||
return uuid.uuid4().hex[:8]
|
||||
|
||||
|
||||
def set_request_id(rid: Optional[str]) -> None:
|
||||
"""设置当前请求的 request_id"""
|
||||
request_id_var.set(rid)
|
||||
|
||||
|
||||
def get_request_id() -> Optional[str]:
|
||||
"""获取当前请求的 request_id"""
|
||||
return request_id_var.get()
|
||||
|
||||
|
||||
# ── JSON 结构化 Formatter ─────────────────────────────────────
|
||||
class JSONFormatter(logging.Formatter):
|
||||
"""将日志记录格式化为单行 JSON 字符串。
|
||||
|
||||
输出字段:
|
||||
- timestamp: ISO-8601 UTC 时间戳
|
||||
- level: 日志级别
|
||||
- logger: logger 名称
|
||||
- message: 日志消息
|
||||
- request_id: 当前请求 ID(如果有)
|
||||
- module/function/line: 代码位置
|
||||
- exc_info: 异常信息(如果有)
|
||||
"""
|
||||
|
||||
def format(self, record: logging.LogRecord) -> str:
|
||||
log_entry = {
|
||||
"timestamp": datetime.now(timezone.utc).isoformat(),
|
||||
"level": record.levelname,
|
||||
"logger": record.name,
|
||||
"message": record.getMessage(),
|
||||
"request_id": request_id_var.get(),
|
||||
"module": record.module,
|
||||
"func": record.funcName,
|
||||
"line": record.lineno,
|
||||
}
|
||||
|
||||
if record.exc_info and record.exc_info[0] is not None:
|
||||
log_entry["exc_info"] = self.formatException(record.exc_info)
|
||||
|
||||
# 支持 extra 字段(通过 logger.info("msg", extra={"key": "val"}))
|
||||
for key in ("method", "path", "status", "duration_ms", "client_ip",
|
||||
"user_agent", "user_id"):
|
||||
val = getattr(record, key, None)
|
||||
if val is not None:
|
||||
log_entry[key] = val
|
||||
|
||||
return json.dumps(log_entry, ensure_ascii=False)
|
||||
|
||||
|
||||
# ── 文本 Formatter(开发环境友好)─────────────────────────────
|
||||
class TextFormatter(logging.Formatter):
|
||||
"""带 request_id 的文本格式,适合本地开发阅读。"""
|
||||
|
||||
def __init__(self):
|
||||
super().__init__(
|
||||
fmt="%(asctime)s [%(levelname)s] %(name)s [rid=%(request_id)s] %(message)s",
|
||||
datefmt="%Y-%m-%d %H:%M:%S",
|
||||
)
|
||||
|
||||
def format(self, record: logging.LogRecord) -> str:
|
||||
if not hasattr(record, "request_id"):
|
||||
record.request_id = request_id_var.get() or "-"
|
||||
return super().format(record)
|
||||
|
||||
|
||||
# ── 公共 API ──────────────────────────────────────────────────
|
||||
def setup_logging(level: Optional[str] = None):
|
||||
"""初始化日志系统。
|
||||
|
||||
Args:
|
||||
level: 日志级别,默认从 LOG_LEVEL 环境变量读取(INFO)
|
||||
|
||||
环境变量:
|
||||
LOG_FORMAT: json(默认)或 text
|
||||
LOG_LEVEL: 日志级别(DEBUG/INFO/WARNING/ERROR)
|
||||
"""
|
||||
log_level = getattr(logging, (level or os.getenv("LOG_LEVEL", "INFO")).upper(), logging.INFO)
|
||||
log_format = os.getenv("LOG_FORMAT", "json").lower()
|
||||
|
||||
formatter = JSONFormatter() if log_format == "json" else TextFormatter()
|
||||
|
||||
handler = logging.StreamHandler(sys.stdout)
|
||||
handler.setFormatter(formatter)
|
||||
|
||||
root = logging.getLogger()
|
||||
root.setLevel(log_level)
|
||||
# 清除已有 handler 避免重复输出
|
||||
root.handlers.clear()
|
||||
root.addHandler(handler)
|
||||
|
||||
# 降低第三方库的日志级别
|
||||
for noisy in ("uvicorn.access", "uvicorn.error", "httpx", "httpcore"):
|
||||
logging.getLogger(noisy).setLevel(logging.WARNING)
|
||||
|
||||
|
||||
def get_logger(name: str) -> logging.Logger:
|
||||
"""获取带模块名的 logger(API 不变,向后兼容)。"""
|
||||
return logging.getLogger(name)
|
||||
|
||||
Reference in New Issue
Block a user