feat: 凭证生成、成本分析、AI问答、前端页面、集成测试与E2E测试

- 后端: 凭证生成引擎、金蝶导出器、凭证模板服务
- 后端: 成本分析服务、AI问答服务
- 后端: 科目映射CRUD API、分析API、QA API
- 后端: 集成测试(认证/任务/凭证) 49个测试全部通过
- 前端: 凭证管理、成本分析、导出中心、知识库、系统设置页面
- 前端: AuthGuard认证守卫、Dashboard AI聊天功能
- 前端: Playwright E2E测试 16 passed, 1 skipped
- 基础设施: Docker Compose、Nginx反向代理、.env.example
- 文档: 用户手册、管理员手册、发布检查清单
This commit is contained in:
selfrelease
2026-07-07 21:21:29 +08:00
parent feebbb10ac
commit 5278190750
47 changed files with 6567 additions and 156 deletions
+241
View File
@@ -0,0 +1,241 @@
"""
成本分析 API
提供人工成本分析、部门拆分、费用科目拆分、环比变化等接口
"""
import csv
import io
from typing import Any, Dict, List, Optional
from fastapi import APIRouter, Depends, HTTPException, Query
from fastapi.responses import StreamingResponse
from pydantic import BaseModel
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.database import get_db
from app.core.tenant import get_current_company_id
from app.services.analysis.cost_calculator import CostCalculatorService
from app.services.ai.cost_analyzer import CostAnalyzerService
router = APIRouter(prefix="/api/analysis", tags=["成本分析"])
class CostSummaryResponse(BaseModel):
"""成本汇总响应"""
total_cost: float
salary_cost: float
social_security_cost: float
fund_cost: float
employee_count: int
class DepartmentCostResponse(BaseModel):
"""部门成本响应"""
department: str
employee_count: int
salary_cost: float
social_security_cost: float
fund_cost: float
total_cost: float
class ExpenseBreakdownResponse(BaseModel):
"""费用科目拆分响应"""
expense_type: str
amount: float
class MonthOverMonthResponse(BaseModel):
"""环比变化响应"""
total_cost: Dict[str, Any]
salary_cost: Dict[str, Any]
social_security_cost: Dict[str, Any]
fund_cost: Dict[str, Any]
employee_count: Dict[str, Any]
class CostChangeAnalysisResponse(BaseModel):
"""成本变化分析响应"""
new_employees: List[Dict[str, Any]]
left_employees: List[Dict[str, Any]]
salary_adjustments: List[Dict[str, Any]]
new_employee_cost: float
left_employee_saving: float
adjustment_cost: float
net_change: float
class FullAnalysisResponse(BaseModel):
"""完整成本分析响应"""
summary: CostSummaryResponse
departments: List[DepartmentCostResponse]
expenses: List[ExpenseBreakdownResponse]
changes: Optional[Dict[str, Any]] = None
ai_summary: Optional[str] = None
@router.get("/labor-cost/{task_id}", response_model=FullAnalysisResponse)
async def get_labor_cost_analysis(
task_id: int,
prev_task_id: Optional[int] = Query(None, description="上月任务ID,用于环比"),
db: AsyncSession = Depends(get_db),
company_id: int = Depends(get_current_company_id),
) -> FullAnalysisResponse:
"""
获取人工成本分析
包含总额、部门拆分、费用科目拆分,可选环比
"""
service = CostCalculatorService(db)
summary = await service.calculate_total_cost(task_id)
departments = await service.calculate_by_department(task_id)
expenses = await service.calculate_by_expense_type(task_id)
changes = None
ai_summary = None
if prev_task_id:
changes = await service.calculate_month_over_month(task_id, prev_task_id)
change_detail = await service.analyze_cost_changes(task_id, prev_task_id)
# 尝试 AI 分析
try:
analyzer = CostAnalyzerService()
prev_summary = await service.calculate_total_cost(prev_task_id)
ai_summary = await analyzer.analyze_cost_changes(
summary.to_dict(),
prev_summary.to_dict(),
change_detail,
)
except Exception:
pass
return FullAnalysisResponse(
summary=CostSummaryResponse(**summary.to_dict()),
departments=[
DepartmentCostResponse(**d.to_dict()) for d in departments
],
expenses=[
ExpenseBreakdownResponse(expense_type=k, amount=v)
for k, v in expenses.items()
],
changes=changes,
ai_summary=ai_summary,
)
@router.get("/labor-cost/{task_id}/summary", response_model=CostSummaryResponse)
async def get_cost_summary(
task_id: int,
db: AsyncSession = Depends(get_db),
company_id: int = Depends(get_current_company_id),
) -> CostSummaryResponse:
"""获取成本汇总"""
service = CostCalculatorService(db)
summary = await service.calculate_total_cost(task_id)
return CostSummaryResponse(**summary.to_dict())
@router.get("/labor-cost/{task_id}/departments", response_model=List[DepartmentCostResponse])
async def get_department_costs(
task_id: int,
db: AsyncSession = Depends(get_db),
company_id: int = Depends(get_current_company_id),
) -> List[DepartmentCostResponse]:
"""获取部门成本拆分"""
service = CostCalculatorService(db)
departments = await service.calculate_by_department(task_id)
return [DepartmentCostResponse(**d.to_dict()) for d in departments]
@router.get("/labor-cost/{task_id}/expenses", response_model=List[ExpenseBreakdownResponse])
async def get_expense_breakdown(
task_id: int,
db: AsyncSession = Depends(get_db),
company_id: int = Depends(get_current_company_id),
) -> List[ExpenseBreakdownResponse]:
"""获取费用科目拆分"""
service = CostCalculatorService(db)
expenses = await service.calculate_by_expense_type(task_id)
return [
ExpenseBreakdownResponse(expense_type=k, amount=v)
for k, v in expenses.items()
]
@router.get("/labor-cost/{task_id}/changes", response_model=CostChangeAnalysisResponse)
async def get_cost_changes(
task_id: int,
prev_task_id: int = Query(..., description="上月任务ID"),
db: AsyncSession = Depends(get_db),
company_id: int = Depends(get_current_company_id),
) -> CostChangeAnalysisResponse:
"""获取成本变化分析"""
service = CostCalculatorService(db)
changes = await service.analyze_cost_changes(task_id, prev_task_id)
return CostChangeAnalysisResponse(**changes)
@router.get("/labor-cost/{task_id}/export")
async def export_cost_analysis(
task_id: int,
db: AsyncSession = Depends(get_db),
company_id: int = Depends(get_current_company_id),
) -> StreamingResponse:
"""
导出成本分析为 Excel
返回包含成本汇总、部门拆分、费用拆分的 Excel 文件
"""
service = CostCalculatorService(db)
summary = await service.calculate_total_cost(task_id)
departments = await service.calculate_by_department(task_id)
expenses = await service.calculate_by_expense_type(task_id)
# 使用 openpyxl 生成 Excel
from openpyxl import Workbook
wb = Workbook()
# Sheet1: 成本汇总
ws1 = wb.active
ws1.title = "成本汇总"
ws1.append(["项目", "金额(元)"])
ws1.append(["人工成本总额", summary.total_cost])
ws1.append(["工资成本", summary.salary_cost])
ws1.append(["社保成本(公司部分)", summary.social_security_cost])
ws1.append(["公积金成本(公司部分)", summary.fund_cost])
ws1.append(["员工人数", summary.employee_count])
# Sheet2: 部门拆分
ws2 = wb.create_sheet("部门拆分")
ws2.append(["部门", "人数", "工资成本", "社保成本", "公积金成本", "合计"])
for dept in departments:
ws2.append([
dept.department,
dept.employee_count,
dept.salary_cost,
dept.social_security_cost,
dept.fund_cost,
dept.total_cost,
])
# Sheet3: 费用科目拆分
ws3 = wb.create_sheet("费用科目拆分")
ws3.append(["费用科目", "金额(元)"])
for expense_type, amount in expenses.items():
ws3.append([expense_type, amount])
# 输出到内存
output = io.BytesIO()
wb.save(output)
output.seek(0)
filename = f"cost_analysis_{task_id}.xlsx"
return StreamingResponse(
output,
media_type="application/vnd.openxmlformats-officedocument.spreadsheetml.sheet",
headers={"Content-Disposition": f'attachment; filename="{filename}"'},
)
+68
View File
@@ -0,0 +1,68 @@
"""
问答 API
提供预置问题列表和问题回答接口
"""
from typing import List
from fastapi import APIRouter, Depends, HTTPException, Body
from pydantic import BaseModel
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.database import get_db
from app.core.tenant import get_current_company_id
from app.services.ai.qa_service import QAService
router = APIRouter(prefix="/api/qa", tags=["问答"])
class QuestionRequest(BaseModel):
"""问题请求"""
task_id: int
question: str
context: str = ""
class QuestionResponse(BaseModel):
"""问题响应"""
answer: str
data_points: List[str] = []
class SuggestedQuestionsResponse(BaseModel):
"""建议问题响应"""
questions: List[str]
@router.get("/suggested-questions", response_model=SuggestedQuestionsResponse)
async def get_suggested_questions(
task_id: int,
context: str = "",
db: AsyncSession = Depends(get_db),
company_id: int = Depends(get_current_company_id),
) -> SuggestedQuestionsResponse:
"""
获取预置问题列表
根据任务状态返回合适的预置问题
"""
service = QAService(db)
questions = await service.get_suggested_questions(task_id, context)
return SuggestedQuestionsResponse(questions=questions)
@router.post("/ask", response_model=QuestionResponse)
async def ask_question(
request: QuestionRequest,
db: AsyncSession = Depends(get_db),
company_id: int = Depends(get_current_company_id),
) -> QuestionResponse:
"""
回答用户问题
基于真实数据回答预置问题
"""
service = QAService(db)
answer = await service.answer_preset_question(request.task_id, request.question)
return QuestionResponse(answer=answer, data_points=[])
+266
View File
@@ -0,0 +1,266 @@
"""
凭证 API
提供凭证生成、查询、确认和金蝶导出接口
"""
import io
from typing import Any, Dict, List, Optional
from fastapi import APIRouter, Depends, HTTPException, Query
from fastapi.responses import StreamingResponse
from pydantic import BaseModel
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.database import get_db
from app.core.tenant import get_current_company_id
from app.models.voucher import Voucher, VoucherStatus
from app.models.account_mapping import AccountMapping
from app.services.voucher.generator import VoucherGeneratorService
from app.services.voucher.kingdee_exporter import KingdeeExporterService
from app.services.voucher.template import VoucherTemplate
router = APIRouter(prefix="/api/vouchers", tags=["凭证管理"])
class VoucherEntryResponse(BaseModel):
"""凭证分录响应"""
account_code: str
account_name: str
debit_amount: float
credit_amount: float
summary: str
department: str = ""
class VoucherResponse(BaseModel):
"""凭证响应"""
id: int
voucher_number: str
voucher_date: str
period: str
summary: str
entries: List[Dict[str, Any]]
total_debit: float
total_credit: float
status: str
confirmed_by: Optional[int] = None
confirmed_at: Optional[str] = None
created_at: str
updated_at: str
class GenerateVoucherRequest(BaseModel):
"""生成凭证请求"""
task_id: int
period: str = ""
class ConfirmVoucherRequest(BaseModel):
"""确认凭证请求"""
user_id: int
class AccountMappingResponse(BaseModel):
"""科目映射响应"""
id: int
standard_field: str
debit_account: str
debit_account_name: str
credit_account: str
credit_account_name: str
cost_center: Optional[str] = None
is_active: bool
class AccountMappingRequest(BaseModel):
"""科目映射请求"""
standard_field: str
debit_account: str
debit_account_name: str
credit_account: str
credit_account_name: str
cost_center: Optional[str] = None
@router.post("/generate", response_model=VoucherResponse)
async def generate_voucher(
request: GenerateVoucherRequest,
db: AsyncSession = Depends(get_db),
company_id: int = Depends(get_current_company_id),
) -> VoucherResponse:
"""生成会计凭证"""
service = VoucherGeneratorService(db)
# 检查是否已有凭证
existing = await service.get_voucher(request.task_id)
if existing:
return _to_response(existing)
voucher = await service.generate_voucher(
task_id=request.task_id,
company_id=company_id,
period=request.period,
)
return _to_response(voucher)
@router.get("/task/{task_id}", response_model=Optional[VoucherResponse])
async def get_task_voucher(
task_id: int,
db: AsyncSession = Depends(get_db),
company_id: int = Depends(get_current_company_id),
) -> Optional[VoucherResponse]:
"""获取任务的凭证"""
service = VoucherGeneratorService(db)
voucher = await service.get_voucher(task_id)
if not voucher:
return None
return _to_response(voucher)
@router.get("/{voucher_id}", response_model=VoucherResponse)
async def get_voucher(
voucher_id: int,
db: AsyncSession = Depends(get_db),
company_id: int = Depends(get_current_company_id),
) -> VoucherResponse:
"""获取凭证详情"""
voucher = await db.get(Voucher, voucher_id)
if not voucher:
raise HTTPException(status_code=404, detail="凭证不存在")
return _to_response(voucher)
@router.post("/{voucher_id}/confirm", response_model=VoucherResponse)
async def confirm_voucher(
voucher_id: int,
request: ConfirmVoucherRequest,
db: AsyncSession = Depends(get_db),
company_id: int = Depends(get_current_company_id),
) -> VoucherResponse:
"""确认凭证"""
service = VoucherGeneratorService(db)
voucher = await service.confirm_voucher(voucher_id, request.user_id)
if not voucher:
raise HTTPException(status_code=404, detail="凭证不存在")
return _to_response(voucher)
@router.get("/task/{task_id}/export")
async def export_voucher(
task_id: int,
format: str = Query("csv", description="导出格式: csv 或 excel"),
db: AsyncSession = Depends(get_db),
company_id: int = Depends(get_current_company_id),
) -> StreamingResponse:
"""导出凭证为金蝶格式"""
service = KingdeeExporterService(db)
try:
return await service.export_task_vouchers(task_id, format)
except ValueError as e:
raise HTTPException(status_code=404, detail=str(e))
# ===== 科目映射管理 =====
@router.get("/account-mappings/list", response_model=List[AccountMappingResponse])
async def list_account_mappings(
db: AsyncSession = Depends(get_db),
company_id: int = Depends(get_current_company_id),
) -> List[AccountMappingResponse]:
"""获取科目映射列表"""
result = await db.execute(
select(AccountMapping).where(AccountMapping.company_id == company_id)
)
mappings = result.scalars().all()
return [AccountMappingResponse(**m.__dict__) for m in mappings]
@router.post("/account-mappings", response_model=AccountMappingResponse)
async def create_account_mapping(
request: AccountMappingRequest,
db: AsyncSession = Depends(get_db),
company_id: int = Depends(get_current_company_id),
) -> AccountMappingResponse:
"""创建科目映射"""
mapping = AccountMapping(
company_id=company_id,
standard_field=request.standard_field,
debit_account=request.debit_account,
debit_account_name=request.debit_account_name,
credit_account=request.credit_account,
credit_account_name=request.credit_account_name,
cost_center=request.cost_center,
is_active=True,
)
db.add(mapping)
await db.commit()
await db.refresh(mapping)
return AccountMappingResponse(**mapping.__dict__)
@router.put("/account-mappings/{mapping_id}", response_model=AccountMappingResponse)
async def update_account_mapping(
mapping_id: int,
request: AccountMappingRequest,
db: AsyncSession = Depends(get_db),
company_id: int = Depends(get_current_company_id),
) -> AccountMappingResponse:
"""更新科目映射"""
mapping = await db.get(AccountMapping, mapping_id)
if not mapping:
raise HTTPException(status_code=404, detail="科目映射不存在")
mapping.standard_field = request.standard_field
mapping.debit_account = request.debit_account
mapping.debit_account_name = request.debit_account_name
mapping.credit_account = request.credit_account
mapping.credit_account_name = request.credit_account_name
mapping.cost_center = request.cost_center
await db.commit()
await db.refresh(mapping)
return AccountMappingResponse(**mapping.__dict__)
@router.delete("/account-mappings/{mapping_id}")
async def delete_account_mapping(
mapping_id: int,
db: AsyncSession = Depends(get_db),
company_id: int = Depends(get_current_company_id),
) -> Dict[str, str]:
"""删除科目映射"""
mapping = await db.get(AccountMapping, mapping_id)
if not mapping:
raise HTTPException(status_code=404, detail="科目映射不存在")
await db.delete(mapping)
await db.commit()
return {"message": "已删除"}
@router.get("/templates/list")
async def get_voucher_templates() -> Dict[str, Any]:
"""获取默认凭证模板"""
return VoucherTemplate.get_all_templates()
def _to_response(voucher: Voucher) -> VoucherResponse:
"""转换凭证模型为响应"""
return VoucherResponse(
id=voucher.id,
voucher_number=voucher.voucher_number,
voucher_date=voucher.voucher_date,
period=voucher.period,
summary=voucher.summary,
entries=voucher.entries or [],
total_debit=voucher.total_debit,
total_credit=voucher.total_credit,
status=voucher.status,
confirmed_by=voucher.confirmed_by,
confirmed_at=voucher.confirmed_at.isoformat() if voucher.confirmed_at else None,
created_at=voucher.created_at.isoformat() if voucher.created_at else "",
updated_at=voucher.updated_at.isoformat() if voucher.updated_at else "",
)
+4 -1
View File
@@ -3,7 +3,7 @@ from fastapi.exceptions import RequestValidationError
from fastapi.middleware.cors import CORSMiddleware
from sqlalchemy.exc import IntegrityError
from app.api import auth, mappings, exceptions, exports, reconciliation, tasks
from app.api import auth, mappings, exceptions, exports, reconciliation, tasks, analysis, qa, vouchers
from app.core.config import get_settings
from app.core.error_handlers import (
generic_exception_handler,
@@ -45,6 +45,9 @@ app.include_router(exceptions.router)
app.include_router(exports.router)
app.include_router(reconciliation.router)
app.include_router(tasks.router)
app.include_router(analysis.router)
app.include_router(qa.router)
app.include_router(vouchers.router)
@app.get("/api/health", tags=["health"])
+14
View File
@@ -0,0 +1,14 @@
from app.models.base import BaseModel
from app.models.company import Company
from app.models.user import User
from app.models.uploaded_file import UploadedFile
from app.models.field_mapping import FieldMapping
from app.models.company_rule import CompanyRule
from app.models.reconciliation_task import ReconciliationTask
from app.models.exception_item import ExceptionItem
from app.models.standard_field import StandardField
from app.models.parsed_file_record import ParsedFileRecord
from app.models.audit_log import AuditLog
from app.models.cost_analysis import CostAnalysis
from app.models.account_mapping import AccountMapping
from app.models.voucher import Voucher
+47
View File
@@ -0,0 +1,47 @@
"""
科目映射模型
存储标准字段到会计科目的映射关系
"""
from typing import Optional
from sqlalchemy import ForeignKey, Integer, String, Float, Boolean, JSON
from sqlalchemy.orm import Mapped, mapped_column, relationship
from app.models.base import BaseModel
class AccountMapping(BaseModel):
"""
科目映射模型
将标准字段(如基本工资、社保等)映射到会计科目
Attributes:
company_id: 企业ID
standard_field: 标准字段名
debit_account: 借方科目代码
debit_account_name: 借方科目名称
credit_account: 贷方科目代码
credit_account_name: 贷方科目名称
cost_center: 成本中心(可选)
is_active: 是否启用
"""
__tablename__ = "account_mappings"
company_id: Mapped[int] = mapped_column(Integer, ForeignKey("companies.id"), nullable=False, index=True)
standard_field: Mapped[str] = mapped_column(String(100), nullable=False, index=True, comment="标准字段名")
debit_account: Mapped[str] = mapped_column(String(50), nullable=False, comment="借方科目代码")
debit_account_name: Mapped[str] = mapped_column(String(200), nullable=False, comment="借方科目名称")
credit_account: Mapped[str] = mapped_column(String(50), nullable=False, comment="贷方科目代码")
credit_account_name: Mapped[str] = mapped_column(String(200), nullable=False, comment="贷方科目名称")
cost_center: Mapped[Optional[str]] = mapped_column(String(100), nullable=True, comment="成本中心")
is_active: Mapped[bool] = mapped_column(Boolean, nullable=False, default=True, comment="是否启用")
# 关系
company = relationship("Company", backref="account_mappings")
def __repr__(self) -> str:
return f"<AccountMapping(field='{self.standard_field}', debit='{self.debit_account}', credit='{self.credit_account}')>"
+55
View File
@@ -0,0 +1,55 @@
"""
成本分析模型
存储人工成本分析结果,包括总额、部门拆分、费用科目拆分、环比变化等
"""
from typing import Optional
from sqlalchemy import ForeignKey, Integer, String, Float, JSON, DateTime
from sqlalchemy.orm import Mapped, mapped_column, relationship
from app.models.base import BaseModel
class CostAnalysis(BaseModel):
"""
成本分析模型
记录一次成本分析的计算结果
Attributes:
task_id: 关联的对账任务ID
company_id: 企业ID
total_cost: 人工成本总额
salary_cost: 工资成本
social_security_cost: 社保成本(公司部分)
fund_cost: 公积金成本(公司部分)
department_breakdown: 部门成本拆分(JSON
expense_breakdown: 费用科目拆分(JSON
month_over_month: 环比变化数据(JSON
ai_summary: AI 生成的分析摘要
"""
__tablename__ = "cost_analyses"
task_id: Mapped[int] = mapped_column(Integer, ForeignKey("reconciliation_tasks.id"), nullable=False, index=True)
company_id: Mapped[int] = mapped_column(Integer, ForeignKey("companies.id"), nullable=False, index=True)
total_cost: Mapped[float] = mapped_column(Float, nullable=False, default=0.0, comment="人工成本总额")
salary_cost: Mapped[float] = mapped_column(Float, nullable=False, default=0.0, comment="工资成本")
social_security_cost: Mapped[float] = mapped_column(Float, nullable=False, default=0.0, comment="社保成本(公司部分)")
fund_cost: Mapped[float] = mapped_column(Float, nullable=False, default=0.0, comment="公积金成本(公司部分)")
department_breakdown: Mapped[Optional[dict]] = mapped_column(JSON, nullable=True, comment="部门成本拆分")
expense_breakdown: Mapped[Optional[dict]] = mapped_column(JSON, nullable=True, comment="费用科目拆分")
month_over_month: Mapped[Optional[dict]] = mapped_column(JSON, nullable=True, comment="环比变化数据")
ai_summary: Mapped[Optional[str]] = mapped_column(String(2000), nullable=True, comment="AI 生成的分析摘要")
# 关系
task = relationship("ReconciliationTask", backref="cost_analyses")
company = relationship("Company", backref="cost_analyses")
def __repr__(self) -> str:
return f"<CostAnalysis(task_id={self.task_id}, total_cost={self.total_cost})>"
+67
View File
@@ -0,0 +1,67 @@
"""
凭证模型
存储生成的会计凭证
"""
from datetime import datetime
from typing import Optional
from sqlalchemy import ForeignKey, Integer, String, Float, DateTime, JSON, Text
from sqlalchemy.orm import Mapped, mapped_column, relationship
from app.models.base import BaseModel
class VoucherStatus:
"""凭证状态"""
DRAFT = "DRAFT"
CONFIRMED = "CONFIRMED"
EXPORTED = "EXPORTED"
class Voucher(BaseModel):
"""
会计凭证模型
Attributes:
company_id: 企业ID
task_id: 对账任务ID
voucher_number: 凭证编号
voucher_date: 凭证日期
period: 会计期间
summary: 摘要
entries: 凭证分录(JSON
total_debit: 借方合计
total_credit: 贷方合计
status: 状态
confirmed_by: 确认人
confirmed_at: 确认时间
"""
__tablename__ = "vouchers"
company_id: Mapped[int] = mapped_column(Integer, ForeignKey("companies.id"), nullable=False, index=True)
task_id: Mapped[int] = mapped_column(Integer, ForeignKey("reconciliation_tasks.id"), nullable=False, index=True)
voucher_number: Mapped[str] = mapped_column(String(50), nullable=False, index=True, comment="凭证编号")
voucher_date: Mapped[str] = mapped_column(String(20), nullable=False, comment="凭证日期 YYYY-MM-DD")
period: Mapped[str] = mapped_column(String(20), nullable=False, index=True, comment="会计期间 YYYY-MM")
summary: Mapped[str] = mapped_column(String(500), nullable=False, comment="摘要")
entries: Mapped[dict] = mapped_column(JSON, nullable=False, comment="凭证分录列表")
total_debit: Mapped[float] = mapped_column(Float, nullable=False, default=0.0, comment="借方合计")
total_credit: Mapped[float] = mapped_column(Float, nullable=False, default=0.0, comment="贷方合计")
status: Mapped[str] = mapped_column(String(20), nullable=False, default=VoucherStatus.DRAFT, index=True, comment="状态")
confirmed_by: Mapped[Optional[int]] = mapped_column(Integer, ForeignKey("users.id"), nullable=True)
confirmed_at: Mapped[Optional[datetime]] = mapped_column(DateTime, nullable=True)
# 关系
company = relationship("Company", backref="vouchers")
task = relationship("ReconciliationTask", backref="vouchers")
confirmer = relationship("User", foreign_keys=[confirmed_by])
def __repr__(self) -> str:
return f"<Voucher(number='{self.voucher_number}', period='{self.period}', status='{self.status}')>"
View File
+231
View File
@@ -0,0 +1,231 @@
"""
AI 成本变化分析服务
使用 LLM 生成成本变化原因摘要和建议追问问题
"""
import json
import logging
from typing import Any, Dict, List, Optional
from app.core.config import get_settings
logger = logging.getLogger(__name__)
settings = get_settings()
# Prompt 模板
COST_ANALYSIS_PROMPT = """你是一个专业的财务分析师。请根据以下成本对比数据,生成简洁的成本变化分析摘要。
## 当前月数据
- 工资成本: {curr_salary}
- 社保成本: {curr_social}
- 公积金成本: {curr_fund}
- 总成本: {curr_total}
- 员工人数: {curr_count}
## 上月数据
- 工资成本: {prev_salary}
- 社保成本: {prev_social}
- 公积金成本: {prev_fund}
- 总成本: {prev_total}
- 员工人数: {prev_count}
## 变化分析
- 新增员工: {new_count} 人,新增成本 {new_cost}
- 离职员工: {left_count} 人,减少成本 {left_cost}
- 薪资调整: {adjustment_count} 人,净变化 {adjustment_cost}
## 要求
1. 用2-3句话概括成本变化的主要原因
2. 包含具体金额和比例
3. 指出top变化项
4. 输出格式:纯文本,不要Markdown
请直接输出分析摘要:
"""
SUGGESTED_QUESTIONS_PROMPT = """基于以下成本分析数据,生成3-5个用户可能追问的问题。
## 成本变化数据
{change_data}
## 要求
1. 问题要具体、有针对性
2. 关注关键变化项
3. 输出JSON数组格式:["问题1", "问题2", ...]
请直接输出JSON
"""
class CostAnalyzerService:
"""AI 成本变化分析服务"""
async def analyze_cost_changes(
self,
current_data: Dict[str, Any],
previous_data: Dict[str, Any],
changes: Dict[str, Any],
) -> str:
"""
使用 LLM 生成成本变化原因摘要
Args:
current_data: 当前月成本数据
previous_data: 上月成本数据
changes: 变化分析数据
Returns:
AI 生成的分析摘要文本
"""
prompt = COST_ANALYSIS_PROMPT.format(
curr_salary=current_data.get("salary_cost", 0),
curr_social=current_data.get("social_security_cost", 0),
curr_fund=current_data.get("fund_cost", 0),
curr_total=current_data.get("total_cost", 0),
curr_count=current_data.get("employee_count", 0),
prev_salary=previous_data.get("salary_cost", 0),
prev_social=previous_data.get("social_security_cost", 0),
prev_fund=previous_data.get("fund_cost", 0),
prev_total=previous_data.get("total_cost", 0),
prev_count=previous_data.get("employee_count", 0),
new_count=len(changes.get("new_employees", [])),
new_cost=changes.get("new_employee_cost", 0),
left_count=len(changes.get("left_employees", [])),
left_cost=changes.get("left_employee_saving", 0),
adjustment_count=len(changes.get("salary_adjustments", [])),
adjustment_cost=changes.get("adjustment_cost", 0),
)
try:
result = await self._call_llm(prompt)
return result.strip()
except Exception as e:
logger.error(f"AI 成本分析失败: {e}")
# 兜底:生成规则化摘要
return self._generate_fallback_summary(current_data, previous_data, changes)
async def generate_suggested_questions(
self, analysis: Dict[str, Any]
) -> List[str]:
"""
基于变化分析生成可追问的问题
Args:
analysis: 变化分析数据
Returns:
建议问题列表
"""
change_summary = json.dumps(analysis, ensure_ascii=False, default=str)
prompt = SUGGESTED_QUESTIONS_PROMPT.format(change_data=change_summary)
try:
result = await self._call_llm(prompt)
questions = json.loads(result.strip())
if isinstance(questions, list):
return questions[:5]
except Exception as e:
logger.error(f"生成建议问题失败: {e}")
# 兜底问题
return self._generate_fallback_questions(analysis)
def _generate_fallback_summary(
self,
current_data: Dict[str, Any],
previous_data: Dict[str, Any],
changes: Dict[str, Any],
) -> str:
"""规则兜底:生成摘要"""
curr_total = current_data.get("total_cost", 0)
prev_total = previous_data.get("total_cost", 0)
change = curr_total - prev_total
ratio = (change / prev_total * 100) if prev_total > 0 else 0
parts = []
if change > 0:
parts.append(f"本月人工成本较上月增加 {change:.2f} 元({ratio:.1f}%")
elif change < 0:
parts.append(f"本月人工成本较上月减少 {abs(change):.2f} 元({abs(ratio):.1f}%")
else:
parts.append("本月人工成本与上月持平")
new_count = len(changes.get("new_employees", []))
left_count = len(changes.get("left_employees", []))
if new_count > 0:
parts.append(f"新增 {new_count} 名员工增加成本 {changes.get('new_employee_cost', 0):.2f}")
if left_count > 0:
parts.append(f"离职 {left_count} 名员工减少成本 {changes.get('left_employee_saving', 0):.2f}")
adj_count = len(changes.get("salary_adjustments", []))
if adj_count > 0:
parts.append(f"{adj_count} 名员工薪资调整净变化 {changes.get('adjustment_cost', 0):.2f}")
return "".join(parts) + ""
def _generate_fallback_questions(self, analysis: Dict[str, Any]) -> List[str]:
"""规则兜底:生成问题"""
questions = []
new_emps = analysis.get("new_employees", [])
left_emps = analysis.get("left_employees", [])
adjustments = analysis.get("salary_adjustments", [])
if new_emps:
questions.append(f"新增的 {len(new_emps)} 名员工分布在哪些部门?")
if left_emps:
questions.append(f"离职的 {len(left_emps)} 名员工减少了多少成本?")
if adjustments:
top = max(adjustments, key=lambda x: abs(x.get("change", 0)))
questions.append(f"薪资调整幅度最大的是谁?变化了多少?")
questions.append("哪个部门成本变化最大?")
questions.append("社保和公积金成本占比如何?")
return questions[:5]
async def _call_llm(self, prompt: str) -> str:
"""调用 LLM API"""
provider = settings.ai_provider
if provider == "zhipu" and settings.zhipu_api_key:
return await self._call_zhipu(prompt)
elif settings.openai_api_key:
return await self._call_openai(prompt)
else:
raise ValueError("未配置 AI API Key")
async def _call_openai(self, prompt: str) -> str:
"""调用 OpenAI API"""
from openai import AsyncOpenAI
client = AsyncOpenAI(api_key=settings.openai_api_key)
response = await client.chat.completions.create(
model=settings.openai_model,
messages=[{"role": "user", "content": prompt}],
temperature=0.3,
max_tokens=500,
)
return response.choices[0].message.content or ""
async def _call_zhipu(self, prompt: str) -> str:
"""调用智谱 AI API"""
import httpx
url = "https://open.bigmodel.cn/api/paas/v4/chat/completions"
headers = {
"Authorization": f"Bearer {settings.zhipu_api_key}",
"Content-Type": "application/json",
}
payload = {
"model": "glm-4-flash",
"messages": [{"role": "user", "content": prompt}],
"temperature": 0.3,
"max_tokens": 500,
}
async with httpx.AsyncClient(timeout=30) as client:
response = await client.post(url, json=payload, headers=headers)
response.raise_for_status()
data = response.json()
return data["choices"][0]["message"]["content"]
+276
View File
@@ -0,0 +1,276 @@
"""
预置问题问答服务
根据任务状态生成预置问题,并基于真实数据回答用户问题
"""
import json
import logging
from typing import List
from app.core.config import get_settings
from app.services.analysis.cost_calculator import CostCalculatorService
logger = logging.getLogger(__name__)
settings = get_settings()
# 预置问题模板
PRESET_QUESTIONS = {
"CREATED": [
"本月需要对账哪些文件?",
"如何开始一个新的对账任务?",
],
"FILES_UPLOADED": [
"本月还缺哪些文件?",
"文件上传后下一步做什么?",
],
"PARSING": [
"文件解析需要多长时间?",
"解析过程中发现了什么问题?",
],
"WAITING_MAPPING_CONFIRM": [
"哪些字段需要人工确认?",
"AI 识别的准确率如何?",
"低置信度字段有哪些?",
],
"MAPPING_CONFIRMED": [
"字段映射确认后下一步做什么?",
"可以开始对账了吗?",
],
"RECONCILING": [
"对账进度如何?",
"对账过程中发现了什么异常?",
],
"COMPLETED": [
"本次对账发现了多少异常?",
"哪些异常最严重?",
"本月人工成本是多少?",
"为什么本月人工成本上涨?",
"哪些部门变化最大?",
"现在可以生成金蝶凭证吗?",
],
"FAILED": [
"对账失败的原因是什么?",
"如何修复错误?",
],
}
# 问题到数据查询的映射
QUESTION_KEYWORDS = {
"异常": "exceptions",
"成本": "cost",
"上涨": "cost",
"下降": "cost",
"部门": "department",
"凭证": "voucher",
"文件": "files",
"字段": "fields",
"置信度": "fields",
}
class QAService:
"""预置问题问答服务"""
def __init__(self, db):
self.db = db
self.cost_calculator = CostCalculatorService(db)
async def get_suggested_questions(self, task_id: int, context: str = "") -> List[str]:
"""
根据当前任务状态生成合适的预置问题
Args:
task_id: 任务ID
context: 上下文信息(可选)
Returns:
预置问题列表(5-8个)
"""
from app.models.reconciliation_task import ReconciliationTask
task = await self.db.get(ReconciliationTask, task_id)
if not task:
return PRESET_QUESTIONS.get("CREATED", [])
status = task.status
questions = PRESET_QUESTIONS.get(status, []).copy()
# 如果 context 指定了特定场景,追加相关问题
if context == "cost_analysis":
cost_questions = [
"本月人工成本是多少?",
"为什么本月人工成本上涨?",
"哪些部门变化最大?",
"社保和公积金成本占比如何?",
"新增员工对成本的影响有多大?",
]
questions = cost_questions
# 确保至少5个问题
if len(questions) < 5:
questions.extend(PRESET_QUESTIONS.get("COMPLETED", [])[: 5 - len(questions)])
return questions[:8]
async def answer_preset_question(self, task_id: int, question: str) -> str:
"""
回答预置问题
Args:
task_id: 任务ID
question: 用户问题
Returns:
包含数据点的简洁答案
"""
# 根据问题关键词决定数据来源
data_type = self._classify_question(question)
if data_type == "cost":
return await self._answer_cost_question(task_id, question)
elif data_type == "exceptions":
return await self._answer_exception_question(task_id, question)
elif data_type == "department":
return await self._answer_department_question(task_id, question)
elif data_type == "files":
return await self._answer_files_question(task_id, question)
elif data_type == "fields":
return await self._answer_fields_question(task_id, question)
elif data_type == "voucher":
return "凭证生成功能即将上线,请先完成对账和异常处理。"
else:
return await self._answer_with_ai(task_id, question)
def _classify_question(self, question: str) -> str:
"""根据问题关键词分类"""
for keyword, data_type in QUESTION_KEYWORDS.items():
if keyword in question:
return data_type
return "general"
async def _answer_cost_question(self, task_id: int, question: str) -> str:
"""回答成本相关问题"""
summary = await self.cost_calculator.calculate_total_cost(task_id)
if "上涨" in question or "增加" in question or "变化" in question:
changes = await self.cost_calculator.analyze_cost_changes(task_id, task_id)
new_cost = changes.get("new_employee_cost", 0)
left_cost = changes.get("left_employee_saving", 0)
adj_cost = changes.get("adjustment_cost", 0)
return (
f"本月人工成本总额 {summary.total_cost:.2f} 元,"
f"其中工资 {summary.salary_cost:.2f} 元、社保 {summary.social_security_cost:.2f} 元、公积金 {summary.fund_cost:.2f} 元。"
f"新增员工增加成本 {new_cost:.2f} 元,离职员工减少 {left_cost:.2f} 元,薪资调整净变化 {adj_cost:.2f} 元。"
)
return (
f"本月人工成本总额 {summary.total_cost:.2f} 元,"
f"{summary.employee_count} 人。"
f"其中工资成本 {summary.salary_cost:.2f} 元,"
f"社保成本 {summary.social_security_cost:.2f} 元,"
f"公积金成本 {summary.fund_cost:.2f} 元。"
)
async def _answer_department_question(self, task_id: int, question: str) -> str:
"""回答部门相关问题"""
dept_costs = await self.cost_calculator.calculate_by_department(task_id)
if not dept_costs:
return "暂无部门成本数据。"
sorted_depts = sorted(dept_costs, key=lambda d: d.total_cost, reverse=True)
top_dept = sorted_depts[0]
return (
f"{len(sorted_depts)} 个部门,"
f"成本最高的是 {top_dept.department}{top_dept.total_cost:.2f} 元,{top_dept.employee_count} 人),"
f"其次是 {sorted_depts[1].department if len(sorted_depts) > 1 else ''}"
)
async def _answer_exception_question(self, task_id: int, question: str) -> str:
"""回答异常相关问题"""
from app.models.exception_item import ExceptionItem
from sqlalchemy import select, func
count_result = await self.db.execute(
select(func.count(ExceptionItem.id)).where(ExceptionItem.task_id == task_id)
)
total = count_result.scalar() or 0
high_result = await self.db.execute(
select(func.count(ExceptionItem.id)).where(
ExceptionItem.task_id == task_id,
ExceptionItem.severity == "HIGH",
)
)
high_count = high_result.scalar() or 0
if "严重" in question:
return f"本次对账共发现 {total} 个异常,其中高严重程度 {high_count} 个,建议优先处理。"
return f"本次对账共发现 {total} 个异常,其中高严重程度 {high_count} 个。"
async def _answer_files_question(self, task_id: int, question: str) -> str:
"""回答文件相关问题"""
from app.models.uploaded_file import UploadedFile
from sqlalchemy import select
result = await self.db.execute(
select(UploadedFile).where(UploadedFile.task_id == task_id)
)
files = result.scalars().all()
if not files:
return "本月尚未上传任何文件,请上传工资表、社保表和个税表。"
file_types = [f.file_type for f in files]
missing = []
if "工资表" not in file_types:
missing.append("工资表")
if "社保表" not in file_types:
missing.append("社保表")
if "个税表" not in file_types:
missing.append("个税表")
if missing:
return f"已上传 {len(files)} 个文件,还缺少:{''.join(missing)}"
return f"已上传 {len(files)} 个文件(工资表、社保表、个税表),可以进入下一步。"
async def _answer_fields_question(self, task_id: int, question: str) -> str:
"""回答字段相关问题"""
from app.models.field_mapping import FieldMapping
from sqlalchemy import select, func
low_result = await self.db.execute(
select(func.count(FieldMapping.id)).where(
FieldMapping.confidence < 0.7,
FieldMapping.is_skipped == False,
)
)
low_count = low_result.scalar() or 0
total_result = await self.db.execute(
select(func.count(FieldMapping.id))
)
total = total_result.scalar() or 0
if "低置信度" in question or "确认" in question:
return f"共有 {total} 个字段需要映射,其中 {low_count} 个低置信度字段需要人工确认。"
return f"AI 共识别 {total} 个字段,其中 {low_count} 个低置信度字段需要人工确认。"
async def _answer_with_ai(self, task_id: int, question: str) -> str:
"""使用 AI 回答通用问题"""
try:
summary = await self.cost_calculator.calculate_total_cost(task_id)
context = f"任务ID: {task_id}, 总成本: {summary.total_cost}, 员工数: {summary.employee_count}"
prompt = f"基于以下数据回答问题:\n数据:{context}\n问题:{question}\n要求:简洁2-3句话,包含数字。"
from app.services.ai.cost_analyzer import CostAnalyzerService
analyzer = CostAnalyzerService()
result = await analyzer._call_llm(prompt)
return result.strip()
except Exception as e:
logger.error(f"AI 回答失败: {e}")
return "暂无法回答该问题,请稍后重试。"
@@ -0,0 +1,356 @@
"""
成本分析计算服务
计算人工成本总额、部门拆分、费用科目拆分、环比变化等
"""
from datetime import datetime
from decimal import Decimal
from typing import Any, Dict, List, Optional
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.models.cost_analysis import CostAnalysis
from app.models.reconciliation_task import ReconciliationTask
from app.services.data_cleaner import DataCleaner
class CostSummary:
"""成本汇总结果"""
def __init__(
self,
total_cost: float = 0.0,
salary_cost: float = 0.0,
social_security_cost: float = 0.0,
fund_cost: float = 0.0,
employee_count: int = 0,
):
self.total_cost = total_cost
self.salary_cost = salary_cost
self.social_security_cost = social_security_cost
self.fund_cost = fund_cost
self.employee_count = employee_count
def to_dict(self) -> Dict[str, Any]:
return {
"total_cost": self.total_cost,
"salary_cost": self.salary_cost,
"social_security_cost": self.social_security_cost,
"fund_cost": self.fund_cost,
"employee_count": self.employee_count,
}
class DepartmentCost:
"""部门成本"""
def __init__(
self,
department: str,
employee_count: int = 0,
salary_cost: float = 0.0,
social_security_cost: float = 0.0,
fund_cost: float = 0.0,
):
self.department = department
self.employee_count = employee_count
self.salary_cost = salary_cost
self.social_security_cost = social_security_cost
self.fund_cost = fund_cost
@property
def total_cost(self) -> float:
return self.salary_cost + self.social_security_cost + self.fund_cost
def to_dict(self) -> Dict[str, Any]:
return {
"department": self.department,
"employee_count": self.employee_count,
"salary_cost": self.salary_cost,
"social_security_cost": self.social_security_cost,
"fund_cost": self.fund_cost,
"total_cost": self.total_cost,
}
class CostCalculatorService:
"""成本分析计算服务"""
def __init__(self, db: AsyncSession):
self.db = db
async def calculate_total_cost(self, task_id: int) -> CostSummary:
"""
计算人工成本总额
Args:
task_id: 对账任务ID
Returns:
成本汇总结果
"""
cleaned_data = await self._load_cleaned_data(task_id)
salary_cost = 0.0
social_security_cost = 0.0
fund_cost = 0.0
employee_count = 0
for record in cleaned_data:
salary_cost += float(record.get("应发工资", 0) or 0)
social_security_cost += float(record.get("养老保险(公司)", 0) or 0)
social_security_cost += float(record.get("医疗保险(公司)", 0) or 0)
social_security_cost += float(record.get("失业保险(公司)", 0) or 0)
fund_cost += float(record.get("公积金(公司)", 0) or 0)
employee_count += 1
total_cost = salary_cost + social_security_cost + fund_cost
return CostSummary(
total_cost=total_cost,
salary_cost=salary_cost,
social_security_cost=social_security_cost,
fund_cost=fund_cost,
employee_count=employee_count,
)
async def calculate_by_department(self, task_id: int) -> List[DepartmentCost]:
"""
按部门汇总人工成本
Args:
task_id: 对账任务ID
Returns:
部门成本列表
"""
cleaned_data = await self._load_cleaned_data(task_id)
dept_map: Dict[str, DepartmentCost] = {}
for record in cleaned_data:
department = record.get("部门", "未分配") or "未分配"
if department not in dept_map:
dept_map[department] = DepartmentCost(department=department)
dept = dept_map[department]
dept.employee_count += 1
dept.salary_cost += float(record.get("应发工资", 0) or 0)
dept.social_security_cost += float(record.get("养老保险(公司)", 0) or 0)
dept.social_security_cost += float(record.get("医疗保险(公司)", 0) or 0)
dept.social_security_cost += float(record.get("失业保险(公司)", 0) or 0)
dept.fund_cost += float(record.get("公积金(公司)", 0) or 0)
return list(dept_map.values())
async def calculate_by_expense_type(self, task_id: int) -> Dict[str, float]:
"""
按费用科目拆分人工成本
Args:
task_id: 对账任务ID
Returns:
费用科目 -> 金额 的映射
"""
dept_costs = await self.calculate_by_department(task_id)
expense_map: Dict[str, float] = {
"管理费用": 0.0,
"销售费用": 0.0,
"研发费用": 0.0,
"其他": 0.0,
}
dept_to_expense = {
"研发部": "研发费用",
"研发中心": "研发费用",
"技术部": "研发费用",
"销售部": "销售费用",
"市场部": "销售费用",
"商务部": "销售费用",
}
for dept in dept_costs:
expense_type = dept_to_expense.get(dept.department, "管理费用")
expense_map[expense_type] = expense_map.get(expense_type, 0.0) + dept.total_cost
return expense_map
async def calculate_month_over_month(
self, task_id: int, prev_task_id: int
) -> Dict[str, Any]:
"""
计算环比变化
Args:
task_id: 当前任务ID
prev_task_id: 上月任务ID
Returns:
环比变化数据
"""
current = await self.calculate_total_cost(task_id)
previous = await self.calculate_total_cost(prev_task_id)
def _calc_change(curr: float, prev: float) -> Dict[str, Any]:
amount_change = curr - prev
ratio_change = (amount_change / prev * 100) if prev > 0 else 0.0
return {
"current": curr,
"previous": prev,
"amount_change": amount_change,
"ratio_change": round(ratio_change, 2),
}
return {
"total_cost": _calc_change(current.total_cost, previous.total_cost),
"salary_cost": _calc_change(current.salary_cost, previous.salary_cost),
"social_security_cost": _calc_change(
current.social_security_cost, previous.social_security_cost
),
"fund_cost": _calc_change(current.fund_cost, previous.fund_cost),
"employee_count": {
"current": current.employee_count,
"previous": previous.employee_count,
"amount_change": current.employee_count - previous.employee_count,
},
}
async def analyze_cost_changes(
self, curr_task_id: int, prev_task_id: int
) -> Dict[str, Any]:
"""
分析成本变化原因
Args:
curr_task_id: 当前任务ID
prev_task_id: 上月任务ID
Returns:
变化分析结果,包含新增/离职员工影响、薪资调整影响等
"""
curr_data = await self._load_cleaned_data(curr_task_id)
prev_data = await self._load_cleaned_data(prev_task_id)
curr_map = {r.get("员工姓名", ""): r for r in curr_data if r.get("员工姓名")}
prev_map = {r.get("员工姓名", ""): r for r in prev_data if r.get("员工姓名")}
new_employees = []
left_employees = []
salary_adjustments = []
for name, record in curr_map.items():
if name not in prev_map:
new_employees.append({
"name": name,
"salary": float(record.get("应发工资", 0) or 0),
})
else:
prev_salary = float(prev_map[name].get("应发工资", 0) or 0)
curr_salary = float(record.get("应发工资", 0) or 0)
if abs(curr_salary - prev_salary) > 0.01:
salary_adjustments.append({
"name": name,
"previous": prev_salary,
"current": curr_salary,
"change": curr_salary - prev_salary,
})
for name, record in prev_map.items():
if name not in curr_map:
left_employees.append({
"name": name,
"salary": float(record.get("应发工资", 0) or 0),
})
new_cost = sum(e["salary"] for e in new_employees)
left_cost = sum(e["salary"] for e in left_employees)
adjustment_cost = sum(a["change"] for a in salary_adjustments)
return {
"new_employees": new_employees,
"left_employees": left_employees,
"salary_adjustments": salary_adjustments,
"new_employee_cost": new_cost,
"left_employee_saving": left_cost,
"adjustment_cost": adjustment_cost,
"net_change": new_cost - left_cost + adjustment_cost,
}
async def save_analysis(
self,
task_id: int,
company_id: int,
summary: CostSummary,
department_breakdown: Optional[List[DepartmentCost]] = None,
expense_breakdown: Optional[Dict[str, float]] = None,
month_over_month: Optional[Dict[str, Any]] = None,
ai_summary: Optional[str] = None,
) -> CostAnalysis:
"""
保存成本分析结果到数据库
Args:
task_id: 任务ID
company_id: 企业ID
summary: 成本汇总
department_breakdown: 部门拆分
expense_breakdown: 费用科目拆分
month_over_month: 环比变化
ai_summary: AI 摘要
Returns:
创建的成本分析记录
"""
analysis = CostAnalysis(
task_id=task_id,
company_id=company_id,
total_cost=summary.total_cost,
salary_cost=summary.salary_cost,
social_security_cost=summary.social_security_cost,
fund_cost=summary.fund_cost,
department_breakdown=(
[d.to_dict() for d in department_breakdown] if department_breakdown else None
),
expense_breakdown=expense_breakdown,
month_over_month=month_over_month,
ai_summary=ai_summary,
)
self.db.add(analysis)
await self.db.commit()
await self.db.refresh(analysis)
return analysis
async def get_analysis(self, task_id: int) -> Optional[CostAnalysis]:
"""获取任务的成本分析结果"""
result = await self.db.execute(
select(CostAnalysis)
.where(CostAnalysis.task_id == task_id)
.order_by(CostAnalysis.created_at.desc())
.limit(1)
)
return result.scalar_one_or_none()
async def _load_cleaned_data(self, task_id: int) -> List[Dict[str, Any]]:
"""
加载清洗后的数据
Args:
task_id: 任务ID
Returns:
清洗后的数据列表
"""
task = await self.db.get(ReconciliationTask, task_id)
if not task:
return []
if task.reconciliation_result and isinstance(task.reconciliation_result, dict):
records = task.reconciliation_result.get("records", [])
if records:
return records
return []
+286
View File
@@ -0,0 +1,286 @@
"""
凭证生成引擎
根据对账数据和科目映射生成会计凭证
"""
import logging
from datetime import datetime
from typing import Any, Dict, List, Optional, Tuple
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.models.account_mapping import AccountMapping
from app.models.voucher import Voucher, VoucherStatus
from app.models.reconciliation_task import ReconciliationTask
from app.services.voucher.template import VoucherTemplate, DEFAULT_ACCOUNT_TEMPLATES
logger = logging.getLogger(__name__)
class VoucherEntry:
"""凭证分录"""
def __init__(
self,
account_code: str,
account_name: str,
debit_amount: float = 0.0,
credit_amount: float = 0.0,
summary: str = "",
department: str = "",
):
self.account_code = account_code
self.account_name = account_name
self.debit_amount = debit_amount
self.credit_amount = credit_amount
self.summary = summary
self.department = department
def to_dict(self) -> Dict[str, Any]:
return {
"account_code": self.account_code,
"account_name": self.account_name,
"debit_amount": self.debit_amount,
"credit_amount": self.credit_amount,
"summary": self.summary,
"department": self.department,
}
class VoucherGeneratorService:
"""凭证生成引擎"""
def __init__(self, db: AsyncSession):
self.db = db
async def generate_voucher(
self,
task_id: int,
company_id: int,
period: str = "",
) -> Voucher:
"""
根据对账任务数据生成会计凭证
Args:
task_id: 对账任务ID
company_id: 企业ID
period: 会计期间(如 2024-01),为空则从任务中获取
Returns:
生成的凭证对象
"""
# 1. 加载对账数据
task = await self.db.get(ReconciliationTask, task_id)
if not task:
raise ValueError(f"任务不存在: {task_id}")
if not period:
period = task.period
# 2. 获取科目映射
mappings = await self._get_account_mappings(company_id)
# 3. 加载清洗后的数据
records = await self._load_cleaned_data(task_id)
# 4. 按字段汇总金额
field_totals = self._aggregate_by_field(records)
# 5. 生成凭证分录
entries = self._generate_entries(field_totals, mappings)
# 6. 计算合计
total_debit = sum(e.debit_amount for e in entries)
total_credit = sum(e.credit_amount for e in entries)
# 7. 生成凭证编号
voucher_number = f"PAY-{period.replace('-', '')}-{task_id:04d}"
# 8. 创建凭证
voucher = Voucher(
company_id=company_id,
task_id=task_id,
voucher_number=voucher_number,
voucher_date=datetime.now().strftime("%Y-%m-%d"),
period=period,
summary=f"{period} 工资薪酬凭证",
entries=[e.to_dict() for e in entries],
total_debit=total_debit,
total_credit=total_credit,
status=VoucherStatus.DRAFT,
)
self.db.add(voucher)
await self.db.commit()
await self.db.refresh(voucher)
logger.info(f"凭证已生成: {voucher.voucher_number}, 借方合计={total_debit}, 贷方合计={total_credit}")
return voucher
async def get_voucher(self, task_id: int) -> Optional[Voucher]:
"""获取任务的凭证"""
result = await self.db.execute(
select(Voucher)
.where(Voucher.task_id == task_id)
.order_by(Voucher.created_at.desc())
.limit(1)
)
return result.scalar_one_or_none()
async def confirm_voucher(self, voucher_id: int, user_id: int) -> Optional[Voucher]:
"""确认凭证"""
voucher = await self.db.get(Voucher, voucher_id)
if not voucher:
return None
voucher.status = VoucherStatus.CONFIRMED
voucher.confirmed_by = user_id
voucher.confirmed_at = datetime.utcnow()
await self.db.commit()
await self.db.refresh(voucher)
return voucher
async def _get_account_mappings(self, company_id: int) -> Dict[str, Dict]:
"""
获取企业的科目映射,不存在则使用默认模板
Args:
company_id: 企业ID
Returns:
标准字段 -> 科目映射配置
"""
result = await self.db.execute(
select(AccountMapping).where(
AccountMapping.company_id == company_id,
AccountMapping.is_active == True,
)
)
db_mappings = result.scalars().all()
# 先用默认模板
mappings = dict(DEFAULT_ACCOUNT_TEMPLATES)
# 用数据库中的映射覆盖
for m in db_mappings:
mappings[m.standard_field] = {
"debit_account": m.debit_account,
"debit_account_name": m.debit_account_name,
"credit_account": m.credit_account,
"credit_account_name": m.credit_account_name,
"cost_center": m.cost_center or "",
}
return mappings
async def _load_cleaned_data(self, task_id: int) -> List[Dict[str, Any]]:
"""加载清洗后的数据"""
task = await self.db.get(ReconciliationTask, task_id)
if not task:
return []
if task.reconciliation_result and isinstance(task.reconciliation_result, dict):
records = task.reconciliation_result.get("records", [])
if records:
return records
return []
def _aggregate_by_field(self, records: List[Dict[str, Any]]) -> Dict[str, float]:
"""
按标准字段汇总金额
Args:
records: 清洗后的数据列表
Returns:
标准字段 -> 总金额
"""
totals: Dict[str, float] = {}
cost_fields = [
"基本工资", "奖金", "补贴", "加班费", "应发工资",
"养老保险(公司)", "医疗保险(公司)", "失业保险(公司)", "公积金(公司)",
"养老保险", "医疗保险", "失业保险", "公积金",
"应缴个税", "实发工资",
]
for record in records:
for field in cost_fields:
value = float(record.get(field, 0) or 0)
if value != 0:
totals[field] = totals.get(field, 0.0) + value
return totals
def _generate_entries(
self,
field_totals: Dict[str, float],
mappings: Dict[str, Dict],
) -> List[VoucherEntry]:
"""
根据字段汇总和科目映射生成凭证分录
Args:
field_totals: 字段金额汇总
mappings: 科目映射
Returns:
凭证分录列表
"""
entries: List[VoucherEntry] = []
for field, amount in field_totals.items():
if abs(amount) < 0.01:
continue
template = mappings.get(field)
if not template:
logger.warning(f"字段 '{field}' 无科目映射,跳过")
continue
# 借方分录
entries.append(VoucherEntry(
account_code=template["debit_account"],
account_name=template["debit_account_name"],
debit_amount=amount,
summary=f"{field}",
))
# 贷方分录
entries.append(VoucherEntry(
account_code=template["credit_account"],
account_name=template["credit_account_name"],
credit_amount=amount,
summary=f"{field}",
))
# 合并相同科目的分录
entries = self._merge_entries(entries)
return entries
def _merge_entries(self, entries: List[VoucherEntry]) -> List[VoucherEntry]:
"""合并相同科目的分录"""
merged: Dict[Tuple[str, str], VoucherEntry] = {}
for entry in entries:
key = (entry.account_code, "debit" if entry.debit_amount > 0 else "credit")
if key in merged:
if entry.debit_amount > 0:
merged[key].debit_amount += entry.debit_amount
else:
merged[key].credit_amount += entry.credit_amount
merged[key].summary += f", {entry.summary}"
else:
merged[key] = VoucherEntry(
account_code=entry.account_code,
account_name=entry.account_name,
debit_amount=entry.debit_amount,
credit_amount=entry.credit_amount,
summary=entry.summary,
)
return list(merged.values())
@@ -0,0 +1,246 @@
"""
金蝶格式导出服务
将凭证导出为金蝶K3/星空可导入的格式
"""
import csv
import io
import logging
from datetime import datetime
from typing import List, Optional
from fastapi.responses import StreamingResponse
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.models.voucher import Voucher
logger = logging.getLogger(__name__)
class KingdeeExporterService:
"""金蝶格式导出服务"""
def __init__(self, db: AsyncSession):
self.db = db
async def export_voucher(
self,
voucher_id: int,
format_type: str = "csv",
) -> StreamingResponse:
"""
导出凭证为金蝶格式
Args:
voucher_id: 凭证ID
format_type: 导出格式 (csv/excel)
Returns:
StreamingResponse
"""
voucher = await self.db.get(Voucher, voucher_id)
if not voucher:
raise ValueError(f"凭证不存在: {voucher_id}")
if format_type == "excel":
return self._export_excel(voucher)
else:
return self._export_csv(voucher)
async def export_task_vouchers(
self,
task_id: int,
format_type: str = "csv",
) -> StreamingResponse:
"""
导出任务的所有凭证
Args:
task_id: 任务ID
format_type: 导出格式
Returns:
StreamingResponse
"""
result = await self.db.execute(
select(Voucher).where(Voucher.task_id == task_id)
)
vouchers = list(result.scalars().all())
if not vouchers:
raise ValueError(f"任务 {task_id} 无凭证")
if len(vouchers) == 1:
return await self.export_voucher(vouchers[0].id, format_type)
# 多凭证导出
if format_type == "excel":
return self._export_multiple_excel(vouchers)
else:
return self._export_multiple_csv(vouchers)
def _export_csv(self, voucher: Voucher) -> StreamingResponse:
"""
金蝶K3 CSV导入格式
格式: 凭证日期, 凭证号, 摘要, 科目代码, 科目名称, 借方金额, 贷方金额, 制单人
"""
output = io.StringIO()
writer = csv.writer(output)
# 金蝶K3导入格式头
writer.writerow([
"凭证日期", "凭证字", "凭证号", "摘要",
"科目代码", "科目名称",
"借方金额", "贷方金额",
"制单人", "审核人",
])
entries = voucher.entries or []
for i, entry in enumerate(entries):
writer.writerow([
voucher.voucher_date,
"",
voucher.voucher_number,
entry.get("summary", voucher.summary),
entry.get("account_code", ""),
entry.get("account_name", ""),
f'{entry.get("debit_amount", 0):.2f}',
f'{entry.get("credit_amount", 0):.2f}',
"AI助手",
"",
])
output.seek(0)
# 转为 bytes
content = output.getvalue().encode("utf-8-sig") # BOM for Excel compatibility
return StreamingResponse(
io.BytesIO(content),
media_type="text/csv",
headers={
"Content-Disposition": f'attachment; filename="voucher_{voucher.voucher_number}.csv"',
},
)
def _export_excel(self, voucher: Voucher) -> StreamingResponse:
"""导出为 Excel 格式"""
from openpyxl import Workbook
wb = Workbook()
ws = wb.active
ws.title = "凭证"
# 标题行
ws.append(["凭证编号", voucher.voucher_number])
ws.append(["凭证日期", voucher.voucher_date])
ws.append(["会计期间", voucher.period])
ws.append(["摘要", voucher.summary])
ws.append([])
# 分录表头
ws.append(["序号", "科目代码", "科目名称", "摘要", "借方金额", "贷方金额"])
entries = voucher.entries or []
for i, entry in enumerate(entries, 1):
ws.append([
i,
entry.get("account_code", ""),
entry.get("account_name", ""),
entry.get("summary", ""),
entry.get("debit_amount", 0),
entry.get("credit_amount", 0),
])
# 合计行
ws.append([])
ws.append(["", "", "", "合计", voucher.total_debit, voucher.total_credit])
output = io.BytesIO()
wb.save(output)
output.seek(0)
return StreamingResponse(
output,
media_type="application/vnd.openxmlformats-officedocument.spreadsheetml.sheet",
headers={
"Content-Disposition": f'attachment; filename="voucher_{voucher.voucher_number}.xlsx"',
},
)
def _export_multiple_csv(self, vouchers: List[Voucher]) -> StreamingResponse:
"""多凭证 CSV 导出"""
output = io.StringIO()
writer = csv.writer(writer)
writer.writerow([
"凭证日期", "凭证字", "凭证号", "摘要",
"科目代码", "科目名称",
"借方金额", "贷方金额",
])
for voucher in vouchers:
entries = voucher.entries or []
for entry in entries:
writer.writerow([
voucher.voucher_date,
"",
voucher.voucher_number,
entry.get("summary", voucher.summary),
entry.get("account_code", ""),
entry.get("account_name", ""),
f'{entry.get("debit_amount", 0):.2f}',
f'{entry.get("credit_amount", 0):.2f}',
])
output.seek(0)
content = output.getvalue().encode("utf-8-sig")
return StreamingResponse(
io.BytesIO(content),
media_type="text/csv",
headers={
"Content-Disposition": f'attachment; filename="vouchers_batch.csv"',
},
)
def _export_multiple_excel(self, vouchers: List[Voucher]) -> StreamingResponse:
"""多凭证 Excel 导出"""
from openpyxl import Workbook
wb = Workbook()
ws = wb.active
ws.title = "凭证汇总"
ws.append([
"凭证日期", "凭证号", "会计期间", "摘要",
"科目代码", "科目名称",
"借方金额", "贷方金额",
])
for voucher in vouchers:
entries = voucher.entries or []
for entry in entries:
ws.append([
voucher.voucher_date,
voucher.voucher_number,
voucher.period,
entry.get("summary", voucher.summary),
entry.get("account_code", ""),
entry.get("account_name", ""),
entry.get("debit_amount", 0),
entry.get("credit_amount", 0),
])
output = io.BytesIO()
wb.save(output)
output.seek(0)
return StreamingResponse(
output,
media_type="application/vnd.openxmlformats-officedocument.spreadsheetml.sheet",
headers={
"Content-Disposition": f'attachment; filename="vouchers_batch.xlsx"',
},
)
+129
View File
@@ -0,0 +1,129 @@
"""
凭证模板服务
定义标准字段到会计科目的默认映射模板
"""
from typing import Dict, List
# 默认科目映射模板(标准字段 -> 借方科目/贷方科目)
DEFAULT_ACCOUNT_TEMPLATES: Dict[str, Dict] = {
"基本工资": {
"debit_account": "6601.01",
"debit_account_name": "管理费用-工资",
"credit_account": "2211.01",
"credit_account_name": "应付职工薪酬-工资",
},
"奖金": {
"debit_account": "6601.01",
"debit_account_name": "管理费用-工资",
"credit_account": "2211.01",
"credit_account_name": "应付职工薪酬-工资",
},
"补贴": {
"debit_account": "6601.01",
"debit_account_name": "管理费用-工资",
"credit_account": "2211.01",
"credit_account_name": "应付职工薪酬-工资",
},
"加班费": {
"debit_account": "6601.01",
"debit_account_name": "管理费用-工资",
"credit_account": "2211.01",
"credit_account_name": "应付职工薪酬-工资",
},
"应发工资": {
"debit_account": "6601.01",
"debit_account_name": "管理费用-工资",
"credit_account": "2211.01",
"credit_account_name": "应付职工薪酬-工资",
},
"养老保险(公司)": {
"debit_account": "6601.02",
"debit_account_name": "管理费用-社保",
"credit_account": "2211.02",
"credit_account_name": "应付职工薪酬-社保",
},
"医疗保险(公司)": {
"debit_account": "6601.02",
"debit_account_name": "管理费用-社保",
"credit_account": "2211.02",
"credit_account_name": "应付职工薪酬-社保",
},
"失业保险(公司)": {
"debit_account": "6601.02",
"debit_account_name": "管理费用-社保",
"credit_account": "2211.02",
"credit_account_name": "应付职工薪酬-社保",
},
"公积金(公司)": {
"debit_account": "6601.03",
"debit_account_name": "管理费用-公积金",
"credit_account": "2211.03",
"credit_account_name": "应付职工薪酬-公积金",
},
"养老保险": {
"debit_account": "2211.01",
"debit_account_name": "应付职工薪酬-工资",
"credit_account": "2221.01",
"credit_account_name": "其他应付款-养老",
},
"医疗保险": {
"debit_account": "2211.01",
"debit_account_name": "应付职工薪酬-工资",
"credit_account": "2221.02",
"credit_account_name": "其他应付款-医疗",
},
"失业保险": {
"debit_account": "2211.01",
"debit_account_name": "应付职工薪酬-工资",
"credit_account": "2221.03",
"credit_account_name": "其他应付款-失业",
},
"公积金": {
"debit_account": "2211.01",
"debit_account_name": "应付职工薪酬-工资",
"credit_account": "2221.04",
"credit_account_name": "其他应付款-公积金",
},
"应缴个税": {
"debit_account": "2211.01",
"debit_account_name": "应付职工薪酬-工资",
"credit_account": "2221.05",
"credit_account_name": "应交税费-个人所得税",
},
"实发工资": {
"debit_account": "2211.01",
"debit_account_name": "应付职工薪酬-工资",
"credit_account": "1001.01",
"credit_account_name": "银行存款",
},
}
class VoucherTemplate:
"""凭证模板服务"""
@staticmethod
def get_default_template(standard_field: str) -> Dict:
"""
获取标准字段的默认科目映射
Args:
standard_field: 标准字段名
Returns:
科目映射配置
"""
return DEFAULT_ACCOUNT_TEMPLATES.get(standard_field, {})
@staticmethod
def get_all_templates() -> Dict[str, Dict]:
"""获取所有默认模板"""
return DEFAULT_ACCOUNT_TEMPLATES
@staticmethod
def get_template_fields() -> List[str]:
"""获取所有有模板的字段列表"""
return list(DEFAULT_ACCOUNT_TEMPLATES.keys())
+2
View File
@@ -0,0 +1,2 @@
"""
"""
+97
View File
@@ -0,0 +1,97 @@
"""
集成测试公共夹具
使用同步 TestClient + Mock AsyncSession 进行 API 集成测试
避免 AsyncClient + ASGITransport 在 pytest-asyncio 下的死锁问题
"""
from typing import AsyncGenerator
from unittest.mock import AsyncMock, MagicMock
import pytest
from fastapi.testclient import TestClient
from app.core.database import get_db
from app.main import app
@pytest.fixture
def mock_db() -> AsyncMock:
"""
创建 Mock AsyncSession
返回一个 AsyncMock 对象,模拟异步数据库会话
"""
session = AsyncMock()
session.commit = AsyncMock()
session.refresh = AsyncMock()
session.rollback = AsyncMock()
session.close = AsyncMock()
session.add = MagicMock()
session.delete = AsyncMock()
session.get = AsyncMock(return_value=None)
result_mock = MagicMock()
result_mock.scalars.return_value.all.return_value = []
result_mock.scalars.return_value.first.return_value = None
result_mock.scalar_one_or_none.return_value = None
result_mock.scalar.return_value = 0
result_mock.one.return_value = MagicMock()
session.execute = AsyncMock(return_value=result_mock)
return session
@pytest.fixture
def client(mock_db: AsyncMock) -> TestClient:
"""
创建同步测试客户端,覆盖数据库依赖
"""
async def override_get_db() -> AsyncGenerator[AsyncMock, None]:
yield mock_db
app.dependency_overrides[get_db] = override_get_db
with TestClient(app) as tc:
yield tc
app.dependency_overrides.clear()
def make_db_result(items: list = None, scalar=None, one=None):
"""
构造 db.execute() 返回值的辅助函数
Args:
items: scalars().all() 返回的列表
scalar: scalar() 返回的标量值
one: one() 返回的行对象
"""
result = MagicMock()
result.scalars.return_value.all.return_value = items or []
result.scalars.return_value.first.return_value = items[0] if items else None
result.scalar_one_or_none.return_value = scalar
result.scalar.return_value = scalar if scalar is not None else 0
if one:
result.one.return_value = one
return result
@pytest.fixture
def auth_token() -> str:
"""
生成测试用 JWT Token(不依赖数据库)
"""
from app.core.security import create_access_token
return create_access_token(data={"sub": "1", "email": "test@example.com"})
@pytest.fixture
def auth_headers(auth_token: str) -> dict:
"""
返回认证请求头
Returns:
包含 Authorization 和 X-Company-ID 的请求头
"""
return {
"Authorization": f"Bearer {auth_token}",
"X-Company-ID": "1",
}
+174
View File
@@ -0,0 +1,174 @@
"""
认证 API 集成测试
测试注册、登录、获取当前用户、刷新 Token 等端点
使用同步 TestClient + Mock 数据库会话
"""
from datetime import datetime
from unittest.mock import AsyncMock, MagicMock, patch
from fastapi.testclient import TestClient
from app.core.security import hash_password
from app.models.user import User, UserStatus
from tests.conftest import make_db_result
class TestAuthAPI:
"""认证 API 集成测试"""
def test_register_user_success(self, client: TestClient, mock_db: AsyncMock):
"""测试用户注册成功"""
# mock: 邮箱不存在
mock_db.execute = AsyncMock(
return_value=make_db_result(scalar=None)
)
# mock: commit + refresh 后返回带 id 的 user
def refresh_side_effect(obj, *args, **kwargs):
obj.id = 1
obj.created_at = datetime.utcnow()
obj.updated_at = datetime.utcnow()
obj.last_login_at = None
obj.permissions = None
obj.status = UserStatus.ACTIVE.value
mock_db.refresh.side_effect = refresh_side_effect
response = client.post(
"/api/auth/register",
json={
"email": "newuser@example.com",
"password": "password123",
"full_name": "新用户",
"company_id": 1,
},
)
assert response.status_code == 201
data = response.json()
assert data["email"] == "newuser@example.com"
assert data["full_name"] == "新用户"
def test_register_duplicate_email(self, client: TestClient, mock_db: AsyncMock):
"""测试重复邮箱注册失败"""
existing_user = MagicMock()
mock_db.execute = AsyncMock(
return_value=make_db_result(scalar=existing_user)
)
response = client.post(
"/api/auth/register",
json={
"email": "dup@example.com",
"password": "password123",
"full_name": "用户1",
"company_id": 1,
},
)
assert response.status_code == 400
def test_login_success(self, client: TestClient, mock_db: AsyncMock):
"""测试登录成功"""
password = "pass123456"
user = User(
id=1,
company_id=1,
email="login@example.com",
hashed_password=hash_password(password),
full_name="登录用户",
role="会计",
status=UserStatus.ACTIVE.value,
created_at=datetime.utcnow(),
updated_at=datetime.utcnow(),
)
mock_db.execute = AsyncMock(
return_value=make_db_result(scalar=user)
)
response = client.post(
"/api/auth/login",
json={"email": "login@example.com", "password": password},
)
assert response.status_code == 200
data = response.json()
assert "access_token" in data
assert data["token_type"] == "bearer"
assert data["user"]["email"] == "login@example.com"
def test_login_wrong_password(self, client: TestClient, mock_db: AsyncMock):
"""测试密码错误登录失败"""
user = User(
id=1,
company_id=1,
email="wrong@example.com",
hashed_password=hash_password("correctpass"),
full_name="用户",
status=UserStatus.ACTIVE.value,
created_at=datetime.utcnow(),
updated_at=datetime.utcnow(),
)
mock_db.execute = AsyncMock(
return_value=make_db_result(scalar=user)
)
response = client.post(
"/api/auth/login",
json={"email": "wrong@example.com", "password": "wrongpass"},
)
assert response.status_code == 401
def test_login_nonexistent_user(self, client: TestClient, mock_db: AsyncMock):
"""测试不存在的用户登录失败"""
mock_db.execute = AsyncMock(
return_value=make_db_result(scalar=None)
)
response = client.post(
"/api/auth/login",
json={"email": "nobody@example.com", "password": "anypass"},
)
assert response.status_code == 401
def test_get_current_user_no_token(self, client: TestClient):
"""测试无 Token 访问被拒"""
response = client.get("/api/auth/me")
assert response.status_code == 401
def test_refresh_token(self, client: TestClient, auth_headers: dict, mock_db: AsyncMock):
"""测试刷新 Token"""
user = User(
id=1,
company_id=1,
email="test@example.com",
hashed_password="x",
full_name="测试用户",
status=UserStatus.ACTIVE.value,
created_at=datetime.utcnow(),
updated_at=datetime.utcnow(),
)
mock_db.execute = AsyncMock(
return_value=make_db_result(scalar=user)
)
response = client.post("/api/auth/refresh", headers=auth_headers)
assert response.status_code == 200
data = response.json()
assert "access_token" in data
def test_logout(self, client: TestClient, auth_headers: dict, mock_db: AsyncMock):
"""测试登出"""
user = User(
id=1,
company_id=1,
email="test@example.com",
hashed_password="x",
full_name="测试用户",
status=UserStatus.ACTIVE.value,
created_at=datetime.utcnow(),
updated_at=datetime.utcnow(),
)
mock_db.execute = AsyncMock(
return_value=make_db_result(scalar=user)
)
response = client.post("/api/auth/logout", headers=auth_headers)
assert response.status_code == 200
+149
View File
@@ -0,0 +1,149 @@
"""
成本分析计算服务单元测试
测试 CostCalculatorService 和 CostSummary / DepartmentCost 的核心逻辑
"""
import pytest
from unittest.mock import AsyncMock, patch
from app.services.analysis.cost_calculator import (
CostCalculatorService,
CostSummary,
DepartmentCost,
)
@pytest.fixture
def mock_db():
"""模拟数据库会话"""
return AsyncMock()
@pytest.fixture
def cost_calculator(mock_db):
"""创建成本计算服务实例"""
return CostCalculatorService(mock_db)
class TestCostSummary:
"""成本汇总结果测试"""
def test_init_defaults(self):
"""测试默认值"""
summary = CostSummary()
assert summary.total_cost == 0.0
assert summary.salary_cost == 0.0
assert summary.social_security_cost == 0.0
assert summary.fund_cost == 0.0
assert summary.employee_count == 0
def test_init_with_values(self):
"""测试带值初始化"""
summary = CostSummary(
total_cost=50000,
salary_cost=30000,
social_security_cost=10000,
fund_cost=10000,
employee_count=10,
)
assert summary.total_cost == 50000
assert summary.salary_cost == 30000
assert summary.employee_count == 10
def test_to_dict(self):
"""测试转字典"""
summary = CostSummary(total_cost=10000, salary_cost=8000, employee_count=5)
d = summary.to_dict()
assert d["total_cost"] == 10000
assert d["salary_cost"] == 8000
assert d["employee_count"] == 5
class TestDepartmentCost:
"""部门成本测试"""
def test_init(self):
"""测试初始化"""
dept = DepartmentCost(department="技术部", employee_count=10, salary_cost=100000)
assert dept.department == "技术部"
assert dept.employee_count == 10
assert dept.salary_cost == 100000
def test_total_cost_property(self):
"""测试 total_cost 属性计算"""
dept = DepartmentCost(
department="财务部",
salary_cost=10000,
social_security_cost=3000,
fund_cost=1200,
)
assert dept.total_cost == 14200
def test_to_dict(self):
"""测试转字典"""
dept = DepartmentCost(department="技术部", salary_cost=10000, employee_count=5)
d = dept.to_dict()
assert d["department"] == "技术部"
assert d["total_cost"] == 10000
assert d["employee_count"] == 5
class TestCostCalculatorService:
"""成本计算服务测试"""
def test_init(self, mock_db):
"""测试初始化"""
service = CostCalculatorService(mock_db)
assert service.db == mock_db
@pytest.mark.asyncio
async def test_calculate_total_cost(self, cost_calculator):
"""测试计算总成本"""
mock_data = [
{"应发工资": 10000, "养老保险(公司)": 2000, "医疗保险(公司)": 1000, "失业保险(公司)": 500, "公积金(公司)": 1200},
{"应发工资": 8000, "养老保险(公司)": 1600, "医疗保险(公司)": 800, "失业保险(公司)": 400, "公积金(公司)": 960},
]
with patch.object(
cost_calculator, "_load_cleaned_data", return_value=mock_data
):
summary = await cost_calculator.calculate_total_cost(task_id=1)
assert summary.employee_count == 2
assert summary.salary_cost == 18000
assert summary.social_security_cost == 6300 # (2000+1000+500) + (1600+800+400)
assert summary.fund_cost == 2160 # 1200 + 960
assert summary.total_cost == 18000 + 6300 + 2160
@pytest.mark.asyncio
async def test_calculate_total_cost_empty(self, cost_calculator):
"""测试空数据计算"""
with patch.object(
cost_calculator, "_load_cleaned_data", return_value=[]
):
summary = await cost_calculator.calculate_total_cost(task_id=1)
assert summary.employee_count == 0
assert summary.total_cost == 0.0
@pytest.mark.asyncio
async def test_calculate_by_department(self, cost_calculator):
"""测试按部门汇总"""
mock_data = [
{"部门": "技术部", "应发工资": 10000, "养老保险(公司)": 2000, "公积金(公司)": 1200},
{"部门": "技术部", "应发工资": 8000, "养老保险(公司)": 1600, "公积金(公司)": 960},
{"部门": "财务部", "应发工资": 12000, "养老保险(公司)": 2400, "公积金(公司)": 1440},
]
with patch.object(
cost_calculator, "_load_cleaned_data", return_value=mock_data
):
departments = await cost_calculator.calculate_by_department(task_id=1)
assert len(departments) == 2
tech = [d for d in departments if d.department == "技术部"][0]
assert tech.employee_count == 2
assert tech.salary_cost == 18000
finance = [d for d in departments if d.department == "财务部"][0]
assert finance.employee_count == 1
assert finance.salary_cost == 12000
+110
View File
@@ -0,0 +1,110 @@
"""
对账任务 API 集成测试
测试任务列表、统计、详情等端点
使用同步 TestClient + Mock 数据库会话
"""
from datetime import datetime
from unittest.mock import AsyncMock, MagicMock
from fastapi.testclient import TestClient
from app.models.reconciliation_task import ReconciliationTask
from tests.conftest import make_db_result
def _make_task(task_id=1, period="2026-07", status="COMPLETED",
total=100, matched=95, exceptions=5):
"""构造测试任务对象"""
return ReconciliationTask(
id=task_id,
company_id=1,
period=period,
status=status,
total_employees=total,
matched_count=matched,
exception_count=exceptions,
file_ids=[],
reconciliation_result={},
created_at=datetime.utcnow(),
updated_at=datetime.utcnow(),
)
class TestTasksAPI:
"""对账任务 API 集成测试"""
def test_list_tasks_empty(self, client: TestClient, auth_headers: dict, mock_db: AsyncMock):
"""测试空任务列表"""
mock_db.execute = AsyncMock(return_value=make_db_result(items=[], scalar=0))
response = client.get("/api/tasks/", headers=auth_headers)
assert response.status_code == 200
data = response.json()
assert data["items"] == []
assert data["total"] == 0
def test_list_tasks_with_data(self, client: TestClient, auth_headers: dict, mock_db: AsyncMock):
"""测试有数据的任务列表"""
task = _make_task()
# tasks.py 先执行 count 查询,再执行 list 查询
call_count = [0]
async def side_effect(*args, **kwargs):
call_count[0] += 1
if call_count[0] == 1:
return make_db_result(scalar=1) # count
return make_db_result(items=[task]) # list
mock_db.execute = AsyncMock(side_effect=side_effect)
response = client.get("/api/tasks/", headers=auth_headers)
assert response.status_code == 200
data = response.json()
assert data["total"] == 1
assert data["items"][0]["period"] == "2026-07"
assert data["items"][0]["status"] == "COMPLETED"
def test_list_tasks_pagination(self, client: TestClient, auth_headers: dict, mock_db: AsyncMock):
"""测试分页"""
tasks = [_make_task(task_id=i, period=f"2025-{i:02d}") for i in range(1, 6)]
call_count = [0]
async def side_effect(*args, **kwargs):
call_count[0] += 1
if call_count[0] == 1:
return make_db_result(scalar=15)
return make_db_result(items=tasks)
mock_db.execute = AsyncMock(side_effect=side_effect)
response = client.get("/api/tasks/?page=1&page_size=5", headers=auth_headers)
assert response.status_code == 200
data = response.json()
assert len(data["items"]) == 5
assert data["total"] == 15
assert data["page"] == 1
def test_get_task_detail(self, client: TestClient, auth_headers: dict, mock_db: AsyncMock):
"""测试获取任务详情"""
task = _make_task(task_id=42)
mock_db.get = AsyncMock(return_value=task)
response = client.get("/api/tasks/42", headers=auth_headers)
assert response.status_code == 200
data = response.json()
assert data["id"] == 42
assert data["period"] == "2026-07"
def test_get_task_not_found(self, client: TestClient, auth_headers: dict, mock_db: AsyncMock):
"""测试获取不存在的任务"""
mock_db.get = AsyncMock(return_value=None)
response = client.get("/api/tasks/99999", headers=auth_headers)
assert response.status_code == 404
def test_list_tasks_no_company_header(self, client: TestClient, auth_headers: dict, mock_db: AsyncMock):
"""测试无企业 ID 头的任务列表(company_id=None 时查不到数据)"""
mock_db.execute = AsyncMock(return_value=make_db_result(items=[], scalar=0))
response = client.get("/api/tasks/")
assert response.status_code == 200
data = response.json()
assert data["total"] == 0
+113
View File
@@ -0,0 +1,113 @@
"""
凭证生成引擎单元测试
测试 VoucherGeneratorService 的核心逻辑
"""
import pytest
from unittest.mock import AsyncMock, MagicMock, patch
from app.services.voucher.generator import VoucherGeneratorService, VoucherEntry
@pytest.fixture
def mock_db():
return AsyncMock()
@pytest.fixture
def generator(mock_db):
return VoucherGeneratorService(mock_db)
class TestVoucherEntry:
"""凭证分录测试"""
def test_entry_creation(self):
"""测试创建借方分录"""
entry = VoucherEntry(
account_code="6601.01",
account_name="管理费用-工资",
debit_amount=10000,
summary="基本工资",
)
assert entry.account_code == "6601.01"
assert entry.debit_amount == 10000
assert entry.credit_amount == 0
def test_entry_to_dict(self):
"""测试分录转字典"""
entry = VoucherEntry(
account_code="2211.01",
account_name="应付职工薪酬",
credit_amount=10000,
summary="基本工资",
department="技术部",
)
d = entry.to_dict()
assert d["account_code"] == "2211.01"
assert d["credit_amount"] == 10000
assert d["department"] == "技术部"
class TestVoucherGeneratorService:
"""凭证生成引擎测试"""
def test_aggregate_by_field(self, generator):
"""测试按字段汇总"""
records = [
{"基本工资": 10000, "奖金": 2000, "养老保险(公司)": 2000},
{"基本工资": 8000, "奖金": 1000, "养老保险(公司)": 1600},
]
totals = generator._aggregate_by_field(records)
assert totals["基本工资"] == 18000
assert totals["奖金"] == 3000
assert totals["养老保险(公司)"] == 3600
def test_aggregate_empty(self, generator):
"""测试空数据汇总"""
totals = generator._aggregate_by_field([])
assert totals == {}
def test_generate_entries_basic(self, generator):
"""测试基本分录生成"""
field_totals = {"基本工资": 18000}
mappings = {
"基本工资": {
"debit_account": "6601.01",
"debit_account_name": "管理费用-工资",
"credit_account": "2211.01",
"credit_account_name": "应付职工薪酬-工资",
}
}
entries = generator._generate_entries(field_totals, mappings)
assert len(entries) >= 2
total_debit = sum(e.debit_amount for e in entries)
total_credit = sum(e.credit_amount for e in entries)
assert total_debit == 18000
assert total_credit == 18000
def test_generate_entries_skip_missing_mapping(self, generator):
"""测试跳过无映射的字段"""
field_totals = {"基本工资": 18000, "未知字段": 5000}
mappings = {
"基本工资": {
"debit_account": "6601.01",
"debit_account_name": "管理费用-工资",
"credit_account": "2211.01",
"credit_account_name": "应付职工薪酬-工资",
}
}
entries = generator._generate_entries(field_totals, mappings)
total_debit = sum(e.debit_amount for e in entries)
assert total_debit == 18000
def test_merge_entries_same_account(self, generator):
"""测试合并相同科目"""
entries = [
VoucherEntry("6601.01", "管理费用-工资", debit_amount=10000, summary="基本工资"),
VoucherEntry("6601.01", "管理费用-工资", debit_amount=8000, summary="奖金"),
]
merged = generator._merge_entries(entries)
assert len(merged) == 1
assert merged[0].debit_amount == 18000
+65
View File
@@ -0,0 +1,65 @@
"""
凭证模板服务单元测试
测试 VoucherTemplate 的默认科目映射
"""
import pytest
from app.services.voucher.template import VoucherTemplate, DEFAULT_ACCOUNT_TEMPLATES
class TestVoucherTemplate:
"""凭证模板服务测试"""
def test_get_default_template_existing(self):
"""测试获取已存在的字段模板"""
template = VoucherTemplate.get_default_template("基本工资")
assert template is not None
assert template["debit_account"] == "6601.01"
assert template["credit_account"] == "2211.01"
assert "管理费用" in template["debit_account_name"]
assert "应付职工薪酬" in template["credit_account_name"]
def test_get_default_template_nonexistent(self):
"""测试获取不存在的字段模板"""
template = VoucherTemplate.get_default_template("不存在的字段")
assert template == {}
def test_get_all_templates(self):
"""测试获取所有模板"""
templates = VoucherTemplate.get_all_templates()
assert len(templates) > 0
assert "基本工资" in templates
assert "养老保险(公司)" in templates
assert "公积金(公司)" in templates
assert "实发工资" in templates
def test_get_template_fields(self):
"""测试获取模板字段列表"""
fields = VoucherTemplate.get_template_fields()
assert len(fields) > 0
assert "基本工资" in fields
assert isinstance(fields, list)
def test_template_structure(self):
"""测试模板结构完整性"""
for field, template in DEFAULT_ACCOUNT_TEMPLATES.items():
assert "debit_account" in template, f"字段 {field} 缺少 debit_account"
assert "debit_account_name" in template, f"字段 {field} 缺少 debit_account_name"
assert "credit_account" in template, f"字段 {field} 缺少 credit_account"
assert "credit_account_name" in template, f"字段 {field} 缺少 credit_account_name"
def test_social_security_templates(self):
"""测试社保相关模板"""
for field in ["养老保险(公司)", "医疗保险(公司)", "失业保险(公司)"]:
template = VoucherTemplate.get_default_template(field)
assert template != {}
assert "社保" in template["debit_account_name"]
def test_fund_template(self):
"""测试公积金模板"""
template = VoucherTemplate.get_default_template("公积金(公司)")
assert template != {}
assert "公积金" in template["debit_account_name"]
assert "公积金" in template["credit_account_name"]
+178
View File
@@ -0,0 +1,178 @@
"""
凭证 API 集成测试
测试科目映射 CRUD、凭证模板等端点
使用同步 TestClient + Mock 数据库会话
"""
from unittest.mock import AsyncMock, MagicMock
from fastapi.testclient import TestClient
from app.models.account_mapping import AccountMapping
from tests.conftest import make_db_result
def _make_mapping(mapping_id=1, field="基本工资", debit="6601.01", credit="2211.01"):
"""构造测试科目映射对象"""
return AccountMapping(
id=mapping_id,
company_id=1,
standard_field=field,
debit_account=debit,
debit_account_name=f"管理费用-{field}",
credit_account=credit,
credit_account_name="应付职工薪酬",
cost_center=None,
is_active=True,
)
class TestVoucherTemplatesAPI:
"""凭证模板 API 测试"""
def test_get_templates(self, client: TestClient, auth_headers: dict):
"""测试获取默认凭证模板"""
response = client.get("/api/vouchers/templates/list", headers=auth_headers)
assert response.status_code == 200
data = response.json()
assert "基本工资" in data
assert "养老保险(公司)" in data
def test_get_templates_no_auth(self, client: TestClient):
"""测试无认证访问模板(该端点不需要认证)"""
response = client.get("/api/vouchers/templates/list")
assert response.status_code == 200
class TestAccountMappingsAPI:
"""科目映射 CRUD API 测试"""
def test_list_mappings_empty(self, client: TestClient, auth_headers: dict, mock_db: AsyncMock):
"""测试空映射列表"""
mock_db.execute = AsyncMock(return_value=make_db_result(items=[]))
response = client.get(
"/api/vouchers/account-mappings/list", headers=auth_headers
)
assert response.status_code == 200
assert response.json() == []
def test_create_mapping(self, client: TestClient, auth_headers: dict, mock_db: AsyncMock):
"""测试创建科目映射"""
def refresh_side_effect(obj, *args, **kwargs):
obj.id = 1
obj.is_active = True
mock_db.refresh.side_effect = refresh_side_effect
response = client.post(
"/api/vouchers/account-mappings",
headers=auth_headers,
json={
"standard_field": "基本工资",
"debit_account": "6601.01",
"debit_account_name": "管理费用-工资",
"credit_account": "2211.01",
"credit_account_name": "应付职工薪酬-工资",
},
)
assert response.status_code == 200
data = response.json()
assert data["standard_field"] == "基本工资"
assert data["debit_account"] == "6601.01"
assert data["is_active"] is True
def test_create_and_list_mappings(self, client: TestClient, auth_headers: dict, mock_db: AsyncMock):
"""测试创建后查询映射列表"""
mappings = [_make_mapping(1, "基本工资"), _make_mapping(2, "奖金", "6601.02")]
mock_db.execute = AsyncMock(return_value=make_db_result(items=mappings))
response = client.get(
"/api/vouchers/account-mappings/list", headers=auth_headers
)
assert response.status_code == 200
data = response.json()
assert len(data) == 2
def test_update_mapping(self, client: TestClient, auth_headers: dict, mock_db: AsyncMock):
"""测试更新科目映射"""
existing = _make_mapping(1)
mock_db.get = AsyncMock(return_value=existing)
def refresh_side_effect(obj, *args, **kwargs):
pass
mock_db.refresh.side_effect = refresh_side_effect
response = client.put(
"/api/vouchers/account-mappings/1",
headers=auth_headers,
json={
"standard_field": "基本工资",
"debit_account": "6601.03",
"debit_account_name": "销售费用-工资",
"credit_account": "2211.01",
"credit_account_name": "应付职工薪酬-工资",
},
)
assert response.status_code == 200
data = response.json()
assert data["debit_account"] == "6601.03"
def test_update_mapping_not_found(self, client: TestClient, auth_headers: dict, mock_db: AsyncMock):
"""测试更新不存在的映射"""
mock_db.get = AsyncMock(return_value=None)
response = client.put(
"/api/vouchers/account-mappings/99999",
headers=auth_headers,
json={
"standard_field": "基本工资",
"debit_account": "6601.01",
"debit_account_name": "管理费用-工资",
"credit_account": "2211.01",
"credit_account_name": "应付职工薪酬-工资",
},
)
assert response.status_code == 404
def test_delete_mapping(self, client: TestClient, auth_headers: dict, mock_db: AsyncMock):
"""测试删除科目映射"""
existing = _make_mapping(1, "奖金")
mock_db.get = AsyncMock(return_value=existing)
response = client.delete(
"/api/vouchers/account-mappings/1", headers=auth_headers
)
assert response.status_code == 200
def test_delete_mapping_not_found(self, client: TestClient, auth_headers: dict, mock_db: AsyncMock):
"""测试删除不存在的映射"""
mock_db.get = AsyncMock(return_value=None)
response = client.delete(
"/api/vouchers/account-mappings/99999", headers=auth_headers
)
assert response.status_code == 404
def test_create_mapping_with_cost_center(self, client: TestClient, auth_headers: dict, mock_db: AsyncMock):
"""测试带成本中心创建映射"""
def refresh_side_effect(obj, *args, **kwargs):
obj.id = 1
obj.is_active = True
mock_db.refresh.side_effect = refresh_side_effect
response = client.post(
"/api/vouchers/account-mappings",
headers=auth_headers,
json={
"standard_field": "基本工资",
"debit_account": "6601.01",
"debit_account_name": "管理费用-工资",
"credit_account": "2211.01",
"credit_account_name": "应付职工薪酬-工资",
"cost_center": "技术部",
},
)
assert response.status_code == 200
data = response.json()
assert data["cost_center"] == "技术部"