This commit is contained in:
2026-07-30 10:30:50 +08:00
parent cf6d708566
commit 853c478657
85 changed files with 4711 additions and 1052 deletions
+4
View File
@@ -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
+15 -64
View File
@@ -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,
)
+15 -74
View File
@@ -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,
)
+2
View File
@@ -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)
+3 -3
View File
@@ -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": "客户已删除"}
+4 -4
View File
@@ -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": "物料供应商关联已删除"}
+5 -5
View File
@@ -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)
+3 -3
View File
@@ -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": "供应商已删除"}
+1 -1
View File
@@ -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)
+6
View File
@@ -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
View File
@@ -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
+1
View File
@@ -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", "调试")
+10 -4
View File
@@ -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")
+226
View File
@@ -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,
}
-43
View File
@@ -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()
+212
View File
@@ -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}")
+15 -1
View File
@@ -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()
+52 -10
View File
@@ -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
View File
@@ -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)