feat: 完善前后端核心功能模块

后端:
- 新增认证(auth)、任务(tasks)、映射(mappings)、对账(reconciliation)、异常(exceptions)、导出(exports) API
- 新增核心模块: database, security, permissions, tenant, exceptions, error_handlers
- 新增数据模型: user, company, reconciliation_task, field_mapping, uploaded_file 等
- 新增服务层: ai_recognizer, file_parser, file_storage, mapping, reconciliation 等
- 添加数据库迁移脚本

前端:
- 新增登录页面和仪表盘页面
- 新增任务列表、任务详情、字段映射页面
- 新增异常处理页面和规则设置页面
- 新增 API 代理路由 /api/[...path]
- 新增 UI 组件库 (button, card, dialog, input, table 等)
- 新增 auth 组件 (ProtectedRoute, PermissionGate)
- 新增 layout 组件 (Header, Sidebar)
- 新增 mapping 组件 (FieldMappingTable, AISuggestionPanel)
- 新增 API 客户端和 hooks (useAsync, useToast, usePermission 等)
- 新增状态管理 (auth-store, company-store, ui-store)
- 集成 Tailwind CSS 和 shadcn/ui 组件库

其他:
- 添加 Alembic 数据库迁移配置
- 添加初始化示例数据脚本
- 更新项目文档
This commit is contained in:
freedakgmail
2026-07-07 09:04:47 +08:00
parent 8487f6eadf
commit 33b4c734aa
117 changed files with 22070 additions and 295 deletions
+117
View File
@@ -0,0 +1,117 @@
# A generic, single database configuration.
[alembic]
# path to migration scripts
script_location = migrations
# template used to generate migration file names; The default value is %%(rev)s_%%(slug)s
# Uncomment the line below if you want the files to be prepended with date and time
# see https://alembic.sqlalchemy.org/en/latest/tutorial.html#editing-the-ini-file
# for all available tokens
# file_template = %%(year)d_%%(month).2d_%%(day).2d_%%(hour).2d%%(minute).2d-%%(rev)s_%%(slug)s
# sys.path path, will be prepended to sys.path if present.
# defaults to the current working directory.
prepend_sys_path = .
# timezone to use when rendering the date within the migration file
# as well as the filename.
# If specified, requires the python>=3.9 or backports.zoneinfo library.
# Any required deps can installed by adding `alembic[tz]` to the pip requirements
# string value is passed to ZoneInfo()
# leave blank for localtime
# timezone =
# max length of characters to apply to the
# "slug" field
# truncate_slug_length = 40
# set to 'true' to run the environment during
# the 'revision' command, regardless of autogenerate
# revision_environment = false
# set to 'true' to allow .pyc and .pyo files without
# a source .py file to be detected as revisions in the
# versions/ directory
# sourceless = false
# version location specification; This defaults
# to migrations/versions. When using multiple version
# directories, initial revisions must be specified with --version-path.
# The path separator used here should be the separator specified by "version_path_separator" below.
# version_locations = %(here)s/bar:%(here)s/bat:migrations/versions
# version path separator; As mentioned above, this is the character used to split
# version_locations. The default within new alembic.ini files is "os", which uses os.pathsep.
# If this key is omitted entirely, it falls back to the legacy behavior of splitting on spaces and/or commas.
# Valid values for version_path_separator are:
#
# version_path_separator = :
# version_path_separator = ;
# version_path_separator = space
version_path_separator = os # Use os.pathsep. Default configuration used for new projects.
# set to 'true' to search source files recursively
# in each "version_locations" directory
# new in Alembic version 1.10
# recursive_version_locations = false
# the output encoding used when revision files
# are written from script.py.mako
# output_encoding = utf-8
# sqlalchemy.url = driver://user:pass@localhost/dbname
# 实际数据库 URL 从环境变量读取
[post_write_hooks]
# post_write_hooks defines scripts or Python functions that are run
# on newly generated revision scripts. See the documentation for further
# detail and examples
# format using "black" - use the console_scripts runner, against the "black" entrypoint
# hooks = black
# black.type = console_scripts
# black.entrypoint = black
# black.options = -l 79 REVISION_SCRIPT_FILENAME
# lint with attempts to fix using "ruff" - use the exec runner, execute a binary
# hooks = ruff
# ruff.type = exec
# ruff.executable = %(here)s/.venv/bin/ruff
# ruff.options = --fix REVISION_SCRIPT_FILENAME
# Logging configuration
[loggers]
keys = root,sqlalchemy,alembic
[handlers]
keys = console
[formatters]
keys = generic
[logger_root]
level = WARN
handlers = console
qualname =
[logger_sqlalchemy]
level = WARN
handlers =
qualname = sqlalchemy.engine
[logger_alembic]
level = INFO
handlers =
qualname = alembic
[handler_console]
class = StreamHandler
args = (sys.stderr,)
level = NOTSET
formatter = generic
[formatter_generic]
format = %(levelname)-5.5s [%(name)s] %(message)s
datefmt = %H:%M:%S
+71
View File
@@ -0,0 +1,71 @@
from fastapi import APIRouter, Depends, HTTPException, status
from fastapi.security import OAuth2PasswordRequestForm
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.database import get_db
from app.models.user import User
from app.schemas.user import Token, UserCreate, UserLogin, UserResponse
from app.services.auth import AuthService, auth_service
router = APIRouter(prefix="/api/auth", tags=["auth"])
@router.post("/register", response_model=UserResponse, status_code=status.HTTP_201_CREATED)
async def register(
user_data: UserCreate,
db: AsyncSession = Depends(get_db),
):
"""
注册新用户
"""
user = await auth_service.register_user(db, user_data)
return user
@router.post("/login", response_model=dict)
async def login(
login_data: UserLogin,
db: AsyncSession = Depends(get_db),
):
"""
用户登录
"""
user, access_token = await auth_service.authenticate_user(db, login_data)
return {
"access_token": access_token,
"token_type": "bearer",
"user": UserResponse.model_validate(user)
}
@router.post("/logout")
async def logout(
current_user: User = Depends(auth_service.get_current_user),
):
"""
用户登出(客户端需清除 Token)
"""
return {"message": "登出成功"}
@router.get("/me", response_model=UserResponse)
async def get_current_user_info(
current_user: User = Depends(auth_service.get_current_user),
):
"""
获取当前用户信息
"""
return current_user
@router.post("/refresh", response_model=Token)
async def refresh_token(
current_user: User = Depends(auth_service.get_current_user),
):
"""
刷新 Token
"""
from app.core.security import create_access_token
access_token = create_access_token(data={"sub": str(current_user.id), "email": current_user.email})
return Token(access_token=access_token)
+219
View File
@@ -0,0 +1,219 @@
"""
异常处理 API
提供异常的查询、更新、批量操作等接口
"""
from typing import Any, Dict, List, Optional
from fastapi import APIRouter, Depends, HTTPException, Query, Body
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.exception_item import ExceptionItem
from app.services.exception_service import (
get_exceptions,
update_exception,
batch_resolve_exceptions,
ExceptionService,
)
router = APIRouter(prefix="/api/exceptions", tags=["异常处理"])
@router.get("/")
async def list_exceptions(
task_id: Optional[int] = Query(None, description="任务ID"),
status: Optional[str] = Query(None, description="状态"),
severity: Optional[str] = Query(None, description="严重程度"),
exception_type: Optional[str] = Query(None, description="异常类型"),
page: int = Query(1, ge=1, description="页码"),
page_size: int = Query(20, ge=1, le=100, description="每页数量"),
db: AsyncSession = Depends(get_db),
company_id: int = Depends(get_current_company_id),
) -> Dict[str, Any]:
"""
查询异常列表
支持按任务、状态、严重程度、类型筛选
"""
exceptions, total = await get_exceptions(
db=db,
task_id=task_id,
company_id=company_id,
status=status,
severity=severity,
exception_type=exception_type,
page=page,
page_size=page_size,
)
return {
"items": [
{
"id": e.id,
"task_id": e.task_id,
"exception_type": e.exception_type,
"severity": e.severity.lower() if e.severity else None,
"status": e.status.lower() if e.status else None,
"employee_id": e.employee_id,
"employee_name": e.employee_name,
"description": e.description,
"detail": e.detail,
"salary_amount": e.salary_amount,
"social_security_amount": e.social_security_amount,
"tax_amount": e.tax_amount,
"bank_amount": e.bank_amount,
"difference_amount": e.difference_amount,
"suggested_action": e.suggested_action,
"ai_suggestion": e.ai_suggestion,
"handling_note": e.handling_note,
"handled_by": e.handled_by,
"handler_name": e.handler.full_name if e.handler else None,
"handled_at": e.handled_at.isoformat() if e.handled_at else None,
"created_at": e.created_at.isoformat() if e.created_at else None,
"updated_at": e.updated_at.isoformat() if e.updated_at else None,
}
for e in exceptions
],
"total": total,
"page": page,
"page_size": page_size,
}
@router.get("/summary")
async def get_exception_summary(
task_id: Optional[int] = Query(None, description="任务ID"),
db: AsyncSession = Depends(get_db),
company_id: int = Depends(get_current_company_id),
) -> Dict[str, Any]:
"""获取异常汇总统计"""
service = ExceptionService(db)
summary = await service.get_exception_summary(
task_id=task_id,
company_id=company_id,
)
return summary
@router.get("/{exception_id}")
async def get_exception_detail(
exception_id: int,
db: AsyncSession = Depends(get_db),
company_id: int = Depends(get_current_company_id),
) -> Dict[str, Any]:
"""获取异常详情"""
service = ExceptionService(db)
exception = await service.get_exception(exception_id)
if not exception:
raise HTTPException(status_code=404, detail="异常不存在")
if exception.company_id != company_id:
raise HTTPException(status_code=403, detail="无权访问该异常")
return {
"id": exception.id,
"task_id": exception.task_id,
"exception_type": exception.exception_type,
"severity": exception.severity,
"status": exception.status,
"employee_id": exception.employee_id,
"employee_name": exception.employee_name,
"description": exception.description,
"detail": exception.detail,
"salary_amount": exception.salary_amount,
"social_security_amount": exception.social_security_amount,
"tax_amount": exception.tax_amount,
"bank_amount": exception.bank_amount,
"difference_amount": exception.difference_amount,
"suggested_action": exception.suggested_action,
"ai_suggestion": exception.ai_suggestion,
"handling_note": exception.handling_note,
"handled_by": exception.handled_by,
"handler_name": exception.handler.full_name if exception.handler else None,
"handled_at": exception.handled_at.isoformat() if exception.handled_at else None,
"created_at": exception.created_at.isoformat() if exception.created_at else None,
"updated_at": exception.updated_at.isoformat() if exception.updated_at else None,
}
@router.patch("/{exception_id}")
async def update_exception_status(
exception_id: int,
status: Optional[str] = Body(None),
handling_note: Optional[str] = Body(None),
db: AsyncSession = Depends(get_db),
company_id: int = Depends(get_current_company_id),
):
"""更新异常状态"""
service = ExceptionService(db)
# 检查权限
exception = await service.get_exception(exception_id)
if not exception:
raise HTTPException(status_code=404, detail="异常不存在")
if exception.company_id != company_id:
raise HTTPException(status_code=403, detail="无权访问该异常")
updated = await service.update_exception(
exception_id=exception_id,
status=status,
handling_note=handling_note,
)
return {"success": True, "data": updated}
@router.post("/batch-update")
async def batch_update_exceptions(
exception_ids: List[int] = Body(..., description="异常ID列表"),
status: str = Body(..., description="新状态"),
db: AsyncSession = Depends(get_db),
company_id: int = Depends(get_current_company_id),
) -> Dict[str, Any]:
"""批量更新异常状态"""
service = ExceptionService(db)
# 检查权限
query = select(ExceptionItem).where(ExceptionItem.id.in_(exception_ids))
result = await db.execute(query)
exceptions = result.scalars().all()
# 验证权限
for e in exceptions:
if e.company_id != company_id:
raise HTTPException(status_code=403, detail="存在无权访问的异常")
count = await service.batch_update_status(exception_ids, status)
return {"success": True, "updated_count": count}
@router.post("/batch-resolve")
async def batch_resolve(
exception_ids: List[int] = Body(..., description="异常ID列表"),
handling_note: str = Body(..., description="处理备注"),
db: AsyncSession = Depends(get_db),
company_id: int = Depends(get_current_company_id),
) -> Dict[str, Any]:
"""批量处理异常(标记为已解决)"""
service = ExceptionService(db)
# 检查权限
query = select(ExceptionItem).where(ExceptionItem.id.in_(exception_ids))
result = await db.execute(query)
exceptions = result.scalars().all()
# 验证权限
for e in exceptions:
if e.company_id != company_id:
raise HTTPException(status_code=403, detail="存在无权访问的异常")
count = await service.batch_resolve(exception_ids, handling_note)
return {"success": True, "resolved_count": count}
+95
View File
@@ -0,0 +1,95 @@
"""
导出 API
提供对账结果导出、金蝶凭证导出等接口
"""
from typing import Any, Dict
from fastapi import APIRouter, Depends, HTTPException, BackgroundTasks
from fastapi.responses import StreamingResponse
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.export_service import export_task_result, export_kingdee_voucher
router = APIRouter(prefix="/api/exports", tags=["导出"])
@router.get("/task/{task_id}")
async def export_task(
task_id: int,
db: AsyncSession = Depends(get_db),
company_id: int = Depends(get_current_company_id),
):
"""
导出任务对账结果
返回 Excel 文件,包含:
- 概览
- 异常列表
- 异常统计
"""
try:
data = await export_task_result(db, task_id)
return StreamingResponse(
iter([data]),
media_type="application/vnd.openxmlformats-officedocument.spreadsheetml.sheet",
headers={
"Content-Disposition": f"attachment; filename=对账结果_{task_id}.xlsx"
},
)
except ValueError as e:
raise HTTPException(status_code=404, detail=str(e))
except ImportError as e:
raise HTTPException(status_code=500, detail="导出功能未启用,请安装 openpyxl")
@router.get("/kingdee/{task_id}")
async def export_kingdee(
task_id: int,
db: AsyncSession = Depends(get_db),
company_id: int = Depends(get_current_company_id),
):
"""
导出金蝶凭证
返回 Excel 文件,可导入到金蝶系统
"""
try:
data = await export_kingdee_voucher(db, task_id)
return StreamingResponse(
iter([data]),
media_type="application/vnd.openxmlformats-officedocument.spreadsheetml.sheet",
headers={
"Content-Disposition": f"attachment; filename=金蝶凭证_{task_id}.xlsx"
},
)
except ValueError as e:
raise HTTPException(status_code=404, detail=str(e))
except ImportError as e:
raise HTTPException(status_code=500, detail="导出功能未启用,请安装 openpyxl")
@router.post("/task/{task_id}/async")
async def export_task_async(
task_id: int,
background_tasks: BackgroundTasks,
db: AsyncSession = Depends(get_db),
company_id: int = Depends(get_current_company_id),
):
"""
异步导出任务对账结果
适用于大数据量导出的场景,返回导出任务ID
"""
# TODO: 实现异步导出任务
return {
"success": True,
"message": "异步导出任务已创建",
"task_id": task_id,
}
+288
View File
@@ -0,0 +1,288 @@
"""
字段映射 API 端点
提供字段识别、映射确认和规则管理功能
"""
from typing import List, Optional
from fastapi import APIRouter, Depends, HTTPException, Query
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.database import get_db
from app.services.auth import AuthService
from app.core.tenant import get_current_company
from app.schemas.mapping import (
FieldMappingResponse,
FieldMappingUpdate,
ConfirmMappingsRequest,
ConfirmMappingsResponse,
SaveAsRuleRequest,
SaveAsRuleResponse,
CompanyRuleResponse,
UpdateRuleStatusRequest,
RecognizeFieldsRequest,
RecognizeFieldsResponse,
TaskMappingOverview,
FileMappingSummary,
)
from app.services.mapping import MappingService
from app.models.user import User
from app.models.company import Company
router = APIRouter(prefix="/mappings", tags=["字段映射"])
def get_mapping_service(db: AsyncSession = Depends(get_db)) -> MappingService:
return MappingService(db)
@router.post("/recognize", response_model=RecognizeFieldsResponse)
async def recognize_fields(
request: RecognizeFieldsRequest,
company_id: int = Depends(get_current_company),
current_user: User = Depends(AuthService.get_current_user),
service: MappingService = Depends(get_mapping_service),
):
"""
触发 AI 字段识别
- 读取已上传文件的表头和样例数据
- 调用 AI 识别服务
- 保存映射建议
- 返回映射结果
"""
try:
mappings = await service.recognize_and_save(
company_id=company_id,
file_id=request.file_id,
file_type=request.file_type
)
# 统计置信度分布
high = sum(1 for m in mappings if m.confidence > 0.9)
medium = sum(1 for m in mappings if 0.7 <= m.confidence <= 0.9)
low = sum(1 for m in mappings if m.confidence < 0.7)
return RecognizeFieldsResponse(
file_id=request.file_id,
mappings=[FieldMappingResponse.model_validate(m) for m in mappings],
high_confidence_count=high,
medium_confidence_count=medium,
low_confidence_count=low,
needs_review=low > 0 or medium > 0
)
except ValueError as e:
raise HTTPException(status_code=400, detail=str(e))
except Exception as e:
raise HTTPException(status_code=500, detail=f"字段识别失败: {str(e)}")
@router.get("/task/{task_id}", response_model=TaskMappingOverview)
async def get_task_mappings(
task_id: int,
company_id: int = Depends(get_current_company),
current_user: User = Depends(AuthService.get_current_user),
service: MappingService = Depends(get_mapping_service),
):
"""
获取任务的字段映射概览
- 按文件分组返回所有映射
- 包含置信度统计
"""
try:
grouped_mappings = await service.get_mappings_by_task(task_id, company_id)
files = []
total_fields = 0
confirmed_fields = 0
for file_id, mappings in grouped_mappings.items():
file = mappings[0].file if mappings else None
high = sum(1 for m in mappings if m.confidence > 0.9)
medium = sum(1 for m in mappings if 0.7 <= m.confidence <= 0.9)
low = sum(1 for m in mappings if m.confidence < 0.7)
confirmed = sum(1 for m in mappings if m.confirmed)
files.append(FileMappingSummary(
file_id=file_id,
file_type=file.file_type if file else "",
file_name=file.original_filename if file else "",
total_fields=len(mappings),
confirmed_fields=confirmed,
high_confidence=high,
medium_confidence=medium,
low_confidence=low,
mappings=[FieldMappingResponse.model_validate(m) for m in mappings]
))
total_fields += len(mappings)
confirmed_fields += confirmed
needs_review = any(
f.medium_confidence > 0 or f.low_confidence > 0
for f in files
)
return TaskMappingOverview(
task_id=task_id,
total_files=len(files),
total_fields=total_fields,
confirmed_fields=confirmed_fields,
needs_review=needs_review,
files=files
)
except Exception as e:
raise HTTPException(status_code=500, detail=f"获取映射失败: {str(e)}")
@router.get("/{mapping_id}", response_model=FieldMappingResponse)
async def get_mapping(
mapping_id: int,
current_user: User = Depends(AuthService.get_current_user),
service: MappingService = Depends(get_mapping_service),
):
"""获取单个字段映射"""
# 这里需要添加 company_id 检查
mapping = await service.db.get(mapping_id)
if not mapping:
raise HTTPException(status_code=404, detail="映射不存在")
return FieldMappingResponse.model_validate(mapping)
@router.put("/{mapping_id}", response_model=FieldMappingResponse)
async def update_mapping(
mapping_id: int,
update: FieldMappingUpdate,
current_user: User = Depends(AuthService.get_current_user),
service: MappingService = Depends(get_mapping_service),
):
"""
修改字段映射
- 修改标准字段
- 标记跳过
"""
mapping = await service.update_mapping(
mapping_id=mapping_id,
standard_field=update.standard_field,
is_skipped=update.is_skipped
)
if not mapping:
raise HTTPException(status_code=404, detail="映射不存在")
return FieldMappingResponse.model_validate(mapping)
@router.post("/confirm", response_model=ConfirmMappingsResponse)
async def confirm_mappings(
request: ConfirmMappingsRequest,
current_user: User = Depends(AuthService.get_current_user),
service: MappingService = Depends(get_mapping_service),
):
"""
批量确认字段映射
- 确认后映射将用于数据处理
- 不会自动保存为规则
"""
confirmed = await service.confirm_mappings(
mapping_ids=request.mapping_ids,
user_id=current_user.id
)
return ConfirmMappingsResponse(
confirmed_count=len(confirmed),
mappings=[FieldMappingResponse.model_validate(m) for m in confirmed]
)
@router.post("/{mapping_id}/save-as-rule", response_model=SaveAsRuleResponse)
async def save_mapping_as_rule(
mapping_id: int,
current_user: User = Depends(AuthService.get_current_user),
service: MappingService = Depends(get_mapping_service),
):
"""
将字段映射保存为企业规则
- 下次上传同类文件时自动应用规则
- 可在规则管理页面查看和管理
"""
from app.models.field_mapping import FieldMapping
# 获取映射
mapping = await service.db.get(FieldMapping, mapping_id)
if not mapping:
raise HTTPException(status_code=404, detail="映射不存在")
# 保存为规则
rule = await service.save_as_rule(mapping, current_user.id)
return SaveAsRuleResponse(
rule_id=rule.id,
source_field=mapping.source_field,
target_field=mapping.standard_field,
message=f"规则已保存:'{mapping.source_field}' -> '{mapping.standard_field}'"
)
# ===== 规则管理 API =====
@router.get("/rules/list", response_model=List[CompanyRuleResponse])
async def list_rules(
rule_type: Optional[str] = Query(None, description="规则类型"),
company_id: int = Depends(get_current_company),
current_user: User = Depends(AuthService.get_current_user),
service: MappingService = Depends(get_mapping_service),
):
"""获取企业的规则列表"""
rules = await service.get_company_rules(company_id, rule_type)
return [CompanyRuleResponse.model_validate(r) for r in rules]
@router.patch("/rules/{rule_id}/status", response_model=CompanyRuleResponse)
async def update_rule_status(
rule_id: int,
request: UpdateRuleStatusRequest,
company_id: int = Depends(get_current_company),
current_user: User = Depends(AuthService.get_current_user),
service: MappingService = Depends(get_mapping_service),
):
"""更新规则状态(启用/停用)"""
rule = await service.update_rule_status(rule_id, request.status)
if not rule:
raise HTTPException(status_code=404, detail="规则不存在")
if rule.company_id != company_id:
raise HTTPException(status_code=403, detail="无权操作此规则")
return CompanyRuleResponse.model_validate(rule)
@router.delete("/rules/{rule_id}")
async def delete_rule(
rule_id: int,
company_id: int = Depends(get_current_company),
current_user: User = Depends(AuthService.get_current_user),
service: MappingService = Depends(get_mapping_service),
):
"""删除规则"""
from app.models.company_rule import CompanyRule
rule = await service.db.get(CompanyRule, rule_id)
if not rule:
raise HTTPException(status_code=404, detail="规则不存在")
if rule.company_id != company_id:
raise HTTPException(status_code=403, detail="无权操作此规则")
success = await service.delete_rule(rule_id)
if not success:
raise HTTPException(status_code=500, detail="删除失败")
return {"message": "规则已删除"}
+141
View File
@@ -0,0 +1,141 @@
"""
对账执行 API
执行对账任务
"""
from typing import Any, Dict, List, Optional
from fastapi import APIRouter, Depends, HTTPException
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.reconciliation_task import ReconciliationTask
from app.services.reconciliation.engine import run_reconciliation
router = APIRouter(prefix="/api/reconciliation", tags=["对账执行"])
@router.post("/execute/{task_id}")
async def execute_reconciliation(
task_id: int,
db: AsyncSession = Depends(get_db),
company_id: int = Depends(get_current_company_id),
) -> Dict[str, Any]:
"""
执行对账
触发对账引擎,处理上传的数据并返回结果
"""
# 获取任务
task = await db.get(ReconciliationTask, task_id)
if not task:
raise HTTPException(status_code=404, detail="任务不存在")
if task.company_id != company_id:
raise HTTPException(status_code=403, detail="无权访问该任务")
if task.status not in ["MAPPING_COMPLETED", "DATA_UPLOADED"]:
raise HTTPException(status_code=400, detail="任务状态不允许执行对账")
try:
# 获取上传的数据
salary_data = task.salary_data or []
social_security_data = task.social_security_data or []
tax_data = task.tax_data or []
bank_data = task.bank_data if hasattr(task, 'bank_data') else None
# 执行对账
result = await run_reconciliation(
db=db,
company_id=company_id,
task_id=task_id,
period=task.period,
salary_records=salary_data,
social_security_records=social_security_data,
tax_records=tax_data,
bank_records=bank_data,
enable_bank_rules=bool(bank_data),
)
return {
"success": True,
"task_id": task_id,
"result": {
"total_employees": result.total_employees,
"matched_employees": result.matched_employees,
"exception_count": result.exception_count,
"exceptions_by_type": result.exceptions_by_type,
"exceptions_by_severity": result.exceptions_by_severity,
"execution_time_ms": result.execution_time_ms,
"completed_at": result.completed_at,
},
}
except Exception as e:
raise HTTPException(status_code=500, detail=f"对账执行失败: {str(e)}")
@router.get("/result/{task_id}")
async def get_reconciliation_result(
task_id: int,
db: AsyncSession = Depends(get_db),
company_id: int = Depends(get_current_company_id),
) -> Dict[str, Any]:
"""
获取对账结果
返回已执行的对账结果汇总
"""
task = await db.get(ReconciliationTask, task_id)
if not task:
raise HTTPException(status_code=404, detail="任务不存在")
if task.company_id != company_id:
raise HTTPException(status_code=403, detail="无权访问该任务")
# 获取异常统计
from app.services.exception_service import ExceptionService
service = ExceptionService(db)
summary = await service.get_exception_summary(task_id=task_id, company_id=company_id)
return {
"task_id": task_id,
"period": task.period,
"total_employees": task.total_employees,
"matched_employees": task.matched_count,
"exception_count": task.exception_count,
"exceptions_by_type": summary.get("by_type", {}),
"exceptions_by_severity": summary.get("by_severity", {}),
"exceptions_by_status": summary.get("by_status", {}),
"completed_at": task.updated_at.isoformat() if task.updated_at else None,
}
@router.post("/retry/{task_id}")
async def retry_reconciliation(
task_id: int,
db: AsyncSession = Depends(get_db),
company_id: int = Depends(get_current_company_id),
) -> Dict[str, Any]:
"""
重试对账
重新执行已完成的或失败的对账任务
"""
task = await db.get(ReconciliationTask, task_id)
if not task:
raise HTTPException(status_code=404, detail="任务不存在")
if task.company_id != company_id:
raise HTTPException(status_code=403, detail="无权访问该任务")
# 重置任务状态
task.status = "PROCESSING"
task.matched_count = 0
task.exception_count = 0
await db.commit()
# 执行对账
return await execute_reconciliation(task_id, db, company_id)
+181
View File
@@ -0,0 +1,181 @@
"""
对账任务 API
管理对账任务的生命周期
"""
from typing import Any, Dict, List, Optional
from fastapi import APIRouter, Depends, HTTPException, Query
from sqlalchemy import Float
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy import select, func
from app.core.database import get_db
from app.core.tenant import get_current_company_id
from app.models.reconciliation_task import ReconciliationTask
from app.models.user import User
from app.schemas.user import UserResponse
router = APIRouter(prefix="/api/tasks", tags=["对账任务"])
@router.get("/", response_model=List[Dict[str, Any]])
async def list_tasks(
period: Optional[str] = Query(None, description="筛选月份,格式: 2026-04"),
status: Optional[str] = Query(None, description="筛选状态"),
skip: int = Query(0, ge=0, description="跳过记录数"),
limit: int = Query(50, ge=1, le=100, description="返回记录数"),
db: AsyncSession = Depends(get_db),
company_id: Optional[int] = Depends(get_current_company_id),
) -> List[Dict[str, Any]]:
"""
获取对账任务列表
支持按月份和状态筛选
"""
# 构建查询
query = select(ReconciliationTask).where(
ReconciliationTask.company_id == company_id
)
if period:
query = query.where(ReconciliationTask.period == period)
if status:
query = query.where(ReconciliationTask.status == status)
# 按创建时间倒序
query = query.order_by(ReconciliationTask.created_at.desc())
# 分页
query = query.offset(skip).limit(limit)
result = await db.execute(query)
tasks = result.scalars().all()
return [
{
"id": task.id,
"period": task.period,
"status": task.status,
"total_employees": task.total_employees,
"matched_count": task.matched_count,
"exception_count": task.exception_count,
"created_at": task.created_at.isoformat() if task.created_at else None,
}
for task in tasks
]
@router.get("/stats", response_model=Dict[str, Any])
async def get_task_stats(
db: AsyncSession = Depends(get_db),
company_id: int = Depends(get_current_company_id),
) -> Dict[str, Any]:
"""
获取任务统计信息
返回本月任务数、待处理异常、已完成、匹配率等统计
"""
from datetime import datetime
from dateutil.relativedelta import relativedelta
# 计算本月时间范围
now = datetime.now()
first_day_of_month = now.replace(day=1, hour=0, minute=0, second=0, microsecond=0)
next_month = first_day_of_month + relativedelta(months=1)
current_period = now.strftime("%Y-%m")
# 统计本月任务
from sqlalchemy import case as sql_case
query = select(
func.count(ReconciliationTask.id).label("total"),
func.sum(
sql_case(
(ReconciliationTask.status == "COMPLETED", 1),
else_=0
)
).label("completed"),
func.sum(
sql_case(
(ReconciliationTask.exception_count > 0, 1),
else_=0
)
).label("with_exceptions"),
).where(
ReconciliationTask.company_id == company_id,
ReconciliationTask.period == current_period,
)
result = await db.execute(query)
row = result.one()
total_tasks = row.total or 0
completed_tasks = row.completed or 0
# 获取待处理异常数
from app.models.exception_item import ExceptionItem
exception_query = select(
func.count(ExceptionItem.id)
).where(
ExceptionItem.company_id == company_id,
ExceptionItem.status == "PENDING",
)
exception_result = await db.execute(exception_query)
pending_exceptions = exception_result.scalar() or 0
# 计算匹配率
match_rate = 0.0
if total_tasks > 0:
match_query = select(
func.avg(
func.cast(ReconciliationTask.matched_count * 100.0 / func.nullif(ReconciliationTask.total_employees, 0), Float)
)
).where(
ReconciliationTask.company_id == company_id,
ReconciliationTask.period == current_period,
ReconciliationTask.status == "COMPLETED",
)
match_result = await db.execute(match_query)
match_rate = match_result.scalar() or 0.0
return {
"monthly_tasks": total_tasks,
"pending_exceptions": pending_exceptions,
"completed_tasks": completed_tasks,
"match_rate": round(match_rate, 1),
"current_period": current_period,
}
@router.get("/{task_id}", response_model=Dict[str, Any])
async def get_task(
task_id: int,
db: AsyncSession = Depends(get_db),
company_id: int = Depends(get_current_company_id),
) -> Dict[str, Any]:
"""
获取任务详情
"""
task = await db.get(ReconciliationTask, task_id)
if not task:
raise HTTPException(status_code=404, detail="任务不存在")
if task.company_id != company_id:
raise HTTPException(status_code=403, detail="无权访问该任务")
return {
"id": task.id,
"period": task.period,
"status": task.status,
"total_employees": task.total_employees,
"matched_count": task.matched_count,
"exception_count": task.exception_count,
"file_ids": task.file_ids,
"reconciliation_result": task.reconciliation_result,
"created_at": task.created_at.isoformat() if task.created_at else None,
"error_message": task.error_message,
}
+65
View File
@@ -0,0 +1,65 @@
from typing import AsyncGenerator
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine
from sqlalchemy.orm import DeclarativeBase
from app.core.config import get_settings
settings = get_settings()
# 创建异步数据库引擎
engine = create_async_engine(
settings.database_url,
echo=settings.debug,
pool_size=5,
max_overflow=10,
pool_pre_ping=True,
)
# 创建异步会话工厂
async_session_maker = async_sessionmaker(
engine,
class_=AsyncSession,
expire_on_commit=False,
autocommit=False,
autoflush=False,
)
# 声明式基类
class Base(DeclarativeBase):
pass
async def get_db() -> AsyncGenerator[AsyncSession, None]:
"""
数据库会话依赖注入
Yields:
AsyncSession: 数据库会话
"""
async with async_session_maker() as session:
try:
yield session
await session.commit()
except Exception:
await session.rollback()
raise
finally:
await session.close()
async def init_db() -> None:
"""
初始化数据库(创建所有表)
仅用于开发环境,生产环境使用 Alembic 迁移
"""
async with engine.begin() as conn:
await conn.run_sync(Base.metadata.create_all)
async def close_db() -> None:
"""
关闭数据库连接
"""
await engine.dispose()
+191
View File
@@ -0,0 +1,191 @@
from typing import Any
from fastapi import HTTPException, Request, status
from fastapi.exceptions import RequestValidationError
from fastapi.responses import JSONResponse
from sqlalchemy.exc import IntegrityError
from app.core.exceptions import (
BusinessException,
ExternalServiceException,
ForbiddenException,
NotFoundException,
S2FException,
UnauthorizedException,
ValidationException,
)
from app.core.logging import logger
def create_error_response(
code: str,
message: str,
details: Any = None,
status_code: int = status.HTTP_500_INTERNAL_SERVER_ERROR,
) -> JSONResponse:
"""
创建统一错误响应
Args:
code: 错误代码
message: 错误消息
details: 错误详情
status_code: HTTP 状态码
Returns:
JSON 响应
"""
error_data = {
"error": {
"code": code,
"message": message,
}
}
if details:
error_data["error"]["details"] = details
return JSONResponse(
status_code=status_code,
content=error_data,
)
async def s2f_exception_handler(request: Request, exc: S2FException) -> JSONResponse:
"""
处理自定义业务异常
Args:
request: 请求对象
exc: 异常对象
Returns:
JSON 响应
"""
logger.warning(
"business_exception",
code=exc.code,
message=exc.message,
details=exc.details,
path=request.url.path,
)
# 根据异常类型映射 HTTP 状态码
status_code_map = {
NotFoundException: status.HTTP_404_NOT_FOUND,
UnauthorizedException: status.HTTP_401_UNAUTHORIZED,
ForbiddenException: status.HTTP_403_FORBIDDEN,
ValidationException: status.HTTP_422_UNPROCESSABLE_ENTITY,
BusinessException: status.HTTP_400_BAD_REQUEST,
ExternalServiceException: status.HTTP_503_SERVICE_UNAVAILABLE,
}
status_code = status_code_map.get(type(exc), status.HTTP_500_INTERNAL_SERVER_ERROR)
return create_error_response(
code=exc.code,
message=exc.message,
details=exc.details,
status_code=status_code,
)
async def validation_exception_handler(request: Request, exc: RequestValidationError) -> JSONResponse:
"""
处理请求验证异常
Args:
request: 请求对象
exc: 验证异常
Returns:
JSON 响应
"""
logger.warning(
"validation_error",
errors=exc.errors(),
body=exc.body,
path=request.url.path,
)
return create_error_response(
code="VALIDATION_ERROR",
message="请求参数验证失败",
details={"errors": exc.errors()},
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
)
async def http_exception_handler(request: Request, exc: HTTPException) -> JSONResponse:
"""
处理 HTTP 异常
Args:
request: 请求对象
exc: HTTP 异常
Returns:
JSON 响应
"""
logger.warning(
"http_exception",
status_code=exc.status_code,
detail=exc.detail,
path=request.url.path,
)
return create_error_response(
code="HTTP_ERROR",
message=str(exc.detail),
status_code=exc.status_code,
)
async def integrity_error_handler(request: Request, exc: IntegrityError) -> JSONResponse:
"""
处理数据库完整性错误
Args:
request: 请求对象
exc: 完整性错误
Returns:
JSON 响应
"""
logger.error(
"database_integrity_error",
error=str(exc),
path=request.url.path,
)
return create_error_response(
code="DATABASE_ERROR",
message="数据库操作失败,可能存在重复或约束冲突",
status_code=status.HTTP_409_CONFLICT,
)
async def generic_exception_handler(request: Request, exc: Exception) -> JSONResponse:
"""
处理未捕获的异常
Args:
request: 请求对象
exc: 异常对象
Returns:
JSON 响应
"""
logger.error(
"unhandled_exception",
error=str(exc),
error_type=type(exc).__name__,
path=request.url.path,
exc_info=True,
)
return create_error_response(
code="INTERNAL_ERROR",
message="服务器内部错误",
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
)
+65
View File
@@ -0,0 +1,65 @@
from typing import Any, Optional
class S2FException(Exception):
"""
财务 AI 助手基础异常类
Attributes:
message: 错误消息
code: 错误代码
details: 错误详情
"""
def __init__(
self,
message: str,
code: str = "INTERNAL_ERROR",
details: Optional[dict[str, Any]] = None,
) -> None:
self.message = message
self.code = code
self.details = details or {}
super().__init__(self.message)
class NotFoundException(S2FException):
"""资源未找到异常"""
def __init__(self, message: str = "资源未找到", details: Optional[dict[str, Any]] = None) -> None:
super().__init__(message=message, code="NOT_FOUND", details=details)
class UnauthorizedException(S2FException):
"""未授权异常"""
def __init__(self, message: str = "未授权访问", details: Optional[dict[str, Any]] = None) -> None:
super().__init__(message=message, code="UNAUTHORIZED", details=details)
class ForbiddenException(S2FException):
"""禁止访问异常"""
def __init__(self, message: str = "禁止访问", details: Optional[dict[str, Any]] = None) -> None:
super().__init__(message=message, code="FORBIDDEN", details=details)
class ValidationException(S2FException):
"""数据验证异常"""
def __init__(self, message: str = "数据验证失败", details: Optional[dict[str, Any]] = None) -> None:
super().__init__(message=message, code="VALIDATION_ERROR", details=details)
class BusinessException(S2FException):
"""业务逻辑异常"""
def __init__(self, message: str, code: str = "BUSINESS_ERROR", details: Optional[dict[str, Any]] = None) -> None:
super().__init__(message=message, code=code, details=details)
class ExternalServiceException(S2FException):
"""外部服务异常"""
def __init__(self, message: str = "外部服务调用失败", details: Optional[dict[str, Any]] = None) -> None:
super().__init__(message=message, code="EXTERNAL_SERVICE_ERROR", details=details)
+61
View File
@@ -0,0 +1,61 @@
# 权限常量定义
# 基于 PRD 第7章权限设计
# 任务管理权限
PERM_TASK_CREATE = "task:create"
PERM_TASK_VIEW = "task:view"
PERM_TASK_UPDATE = "task:update"
PERM_TASK_DELETE = "task:delete"
PERM_TASK_EXPORT = "task:export"
# 用户管理权限
PERM_USER_CREATE = "user:create"
PERM_USER_VIEW = "user:view"
PERM_USER_UPDATE = "user:update"
PERM_USER_DELETE = "user:delete"
# 企业管理权限
PERM_COMPANY_VIEW = "company:view"
PERM_COMPANY_UPDATE = "company:update"
# 审计日志权限
PERM_AUDIT_VIEW = "audit:view"
# 角色默认权限映射
ROLE_PERMISSIONS = {
"管理员": [
PERM_TASK_CREATE, PERM_TASK_VIEW, PERM_TASK_UPDATE, PERM_TASK_DELETE, PERM_TASK_EXPORT,
PERM_USER_CREATE, PERM_USER_VIEW, PERM_USER_UPDATE, PERM_USER_DELETE,
PERM_COMPANY_VIEW, PERM_COMPANY_UPDATE,
PERM_AUDIT_VIEW,
],
"财务主管": [
PERM_TASK_CREATE, PERM_TASK_VIEW, PERM_TASK_UPDATE, PERM_TASK_DELETE, PERM_TASK_EXPORT,
PERM_USER_CREATE, PERM_USER_VIEW, PERM_USER_UPDATE,
PERM_COMPANY_VIEW,
PERM_AUDIT_VIEW,
],
"会计": [
PERM_TASK_CREATE, PERM_TASK_VIEW, PERM_TASK_UPDATE, PERM_TASK_EXPORT,
PERM_USER_VIEW,
],
"出纳": [
PERM_TASK_VIEW, PERM_TASK_EXPORT,
],
"人事": [
PERM_TASK_VIEW,
],
}
def get_role_permissions(role: str) -> list[str]:
"""
获取角色的默认权限列表
Args:
role: 角色名称
Returns:
权限列表
"""
return ROLE_PERMISSIONS.get(role, [])
+76
View File
@@ -0,0 +1,76 @@
from datetime import datetime, timedelta
from typing import Any
import bcrypt
from jose import JWTError, jwt
from app.core.config import get_settings
settings = get_settings()
def hash_password(password: str) -> str:
"""
加密密码
Args:
password: 明文密码
Returns:
加密后的密码
"""
return bcrypt.hashpw(password.encode(), bcrypt.gensalt()).decode()
def verify_password(plain_password: str, hashed_password: str) -> bool:
"""
验证密码
Args:
plain_password: 明文密码
hashed_password: 加密后的密码
Returns:
密码是否正确
"""
return bcrypt.checkpw(plain_password.encode(), hashed_password.encode())
def create_access_token(data: dict[str, Any], expires_delta: timedelta | None = None) -> str:
"""
创建访问令牌
Args:
data: 要编码的数据
expires_delta: 过期时间(可选)
Returns:
JWT token
"""
to_encode = data.copy()
if expires_delta:
expire = datetime.utcnow() + expires_delta
else:
expire = datetime.utcnow() + timedelta(minutes=settings.jwt_access_token_expire_minutes)
to_encode.update({"exp": expire})
encoded_jwt = jwt.encode(to_encode, settings.jwt_secret_key, algorithm=settings.jwt_algorithm)
return encoded_jwt
def decode_access_token(token: str) -> dict[str, Any] | None:
"""
解码访问令牌
Args:
token: JWT token
Returns:
解码后的数据,解码失败返回 None
"""
try:
payload = jwt.decode(token, settings.jwt_secret_key, algorithms=[settings.jwt_algorithm])
return payload
except JWTError:
return None
+108
View File
@@ -0,0 +1,108 @@
from contextvars import ContextVar
from typing import Optional
from fastapi import Depends, Request
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.database import get_db
from app.models.company import Company
from app.services.company import company_service
# 当前租户的上下文变量
_current_company_id: ContextVar[Optional[int]] = ContextVar("current_company_id", default=None)
def get_current_company_id(
request: Request,
) -> Optional[int]:
"""
获取当前租户 ID
从请求头 X-Company-ID 获取租户 ID
Args:
request: FastAPI 请求对象
Returns:
当前租户 ID,未设置则返回 None
"""
company_id_str = request.headers.get("X-Company-ID")
if not company_id_str:
return None
try:
return int(company_id_str)
except ValueError:
return None
def set_current_company_id(company_id: Optional[int]) -> None:
"""
设置当前租户 ID
Args:
company_id: 租户 ID
"""
_current_company_id.set(company_id)
async def get_current_company(
request: Request,
db: AsyncSession = Depends(get_db),
) -> Optional[Company]:
"""
从请求头获取当前企业
支持两种方式:
1. HTTP Header: X-Company-ID
2. JWT Token 中的 company_id(后续实现认证后)
Args:
request: FastAPI 请求对象
db: 数据库会话
Returns:
当前企业对象,不存在则返回 None
"""
# 从请求头获取 company_id
company_id_str = request.headers.get("X-Company-ID")
if not company_id_str:
return None
try:
company_id = int(company_id_str)
except ValueError:
return None
# 设置上下文变量
set_current_company_id(company_id)
# 查询企业
company = await company_service.get_company(db, company_id)
return company
def require_company(
company: Optional[Company] = Depends(get_current_company),
) -> Company:
"""
要求必须有企业上下文的依赖
Args:
company: 当前企业
Returns:
企业对象
Raises:
HTTPException: 未提供企业 ID 或企业不存在
"""
if not company:
from fastapi import HTTPException, status
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="未提供企业 ID 或企业不存在",
)
return company
+25
View File
@@ -1,7 +1,18 @@
from fastapi import FastAPI
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.core.config import get_settings
from app.core.error_handlers import (
generic_exception_handler,
http_exception_handler,
integrity_error_handler,
s2f_exception_handler,
validation_exception_handler,
)
from app.core.exceptions import S2FException
from app.core.logging import configure_logging
settings = get_settings()
@@ -21,6 +32,20 @@ app.add_middleware(
allow_headers=["*"],
)
# 注册异常处理器
app.add_exception_handler(S2FException, s2f_exception_handler)
app.add_exception_handler(RequestValidationError, validation_exception_handler)
app.add_exception_handler(IntegrityError, integrity_error_handler)
app.add_exception_handler(Exception, generic_exception_handler)
# 注册路由
app.include_router(auth.router)
app.include_router(mappings.router)
app.include_router(exceptions.router)
app.include_router(exports.router)
app.include_router(reconciliation.router)
app.include_router(tasks.router)
@app.get("/api/health", tags=["health"])
async def health_check() -> dict[str, str]:
+53
View File
@@ -0,0 +1,53 @@
from enum import Enum
from sqlalchemy import Integer, JSON, String, Text
from sqlalchemy.orm import Mapped, mapped_column
from app.models.base import BaseModel
class AuditAction(str, Enum):
"""审计操作类型"""
CREATE = "CREATE"
UPDATE = "UPDATE"
DELETE = "DELETE"
VIEW = "VIEW"
EXPORT = "EXPORT"
LOGIN = "LOGIN"
LOGOUT = "LOGOUT"
UPLOAD = "UPLOAD"
DOWNLOAD = "DOWNLOAD"
class AuditLog(BaseModel):
"""
审计日志模型
记录所有关键业务操作
Attributes:
company_id: 企业 ID
user_id: 用户 ID
action: 操作类型
resource_type: 资源类型 (如 company, user, task, file)
resource_id: 资源 ID
details: 操作详情 (JSON)
ip_address: IP 地址
user_agent: User-Agent
"""
__tablename__ = "audit_logs"
company_id: Mapped[int] = mapped_column(Integer, nullable=False, index=True)
user_id: Mapped[int | None] = mapped_column(Integer, nullable=True, index=True)
action: Mapped[str] = mapped_column(String(50), nullable=False, index=True)
resource_type: Mapped[str] = mapped_column(String(100), nullable=False, index=True)
resource_id: Mapped[str | None] = mapped_column(String(200), nullable=True, index=True)
details: Mapped[dict | None] = mapped_column(JSON, nullable=True)
ip_address: Mapped[str | None] = mapped_column(String(50), nullable=True)
user_agent: Mapped[str | None] = mapped_column(Text, nullable=True)
def __repr__(self) -> str:
return (
f"<AuditLog(id={self.id}, company_id={self.company_id}, "
f"action='{self.action}', resource_type='{self.resource_type}')>"
)
+39
View File
@@ -0,0 +1,39 @@
from datetime import datetime
from typing import Any
from sqlalchemy import DateTime, Integer
from sqlalchemy.orm import Mapped, mapped_column
from sqlalchemy.sql import func
from app.core.database import Base
class BaseModel(Base):
"""
所有数据库模型的基类
包含通用字段:id, created_at, updated_at
"""
__abstract__ = True
id: Mapped[int] = mapped_column(Integer, primary_key=True, index=True, autoincrement=True)
created_at: Mapped[datetime] = mapped_column(
DateTime(timezone=True),
server_default=func.now(),
nullable=False,
)
updated_at: Mapped[datetime] = mapped_column(
DateTime(timezone=True),
server_default=func.now(),
onupdate=func.now(),
nullable=False,
)
def dict(self) -> dict[str, Any]:
"""
将模型转换为字典
"""
return {
column.name: getattr(self, column.name)
for column in self.__table__.columns
}
+63
View File
@@ -0,0 +1,63 @@
from enum import Enum
from sqlalchemy import Boolean, Integer, String
from sqlalchemy.orm import Mapped, mapped_column, relationship
from app.models.base import BaseModel
class CompanyPlan(str, Enum):
"""企业套餐类型"""
FREE = "免费版"
BASIC = "基础版"
PROFESSIONAL = "专业版"
ENTERPRISE = "企业版"
class CompanyStatus(str, Enum):
"""企业状态"""
ACTIVE = "正常"
SUSPENDED = "暂停"
EXPIRED = "已过期"
DELETED = "已删除"
class Company(BaseModel):
"""
企业模型 - 多租户隔离的核心实体
Attributes:
name: 企业名称
plan: 套餐类型
data_retention_months: 数据保留月数
status: 企业状态
max_users: 最大用户数
is_trial: 是否试用
"""
__tablename__ = "companies"
name: Mapped[str] = mapped_column(String(200), nullable=False, index=True)
plan: Mapped[str] = mapped_column(
String(50),
nullable=False,
default=CompanyPlan.FREE.value
)
data_retention_months: Mapped[int] = mapped_column(Integer, nullable=False, default=12)
status: Mapped[str] = mapped_column(
String(20),
nullable=False,
default=CompanyStatus.ACTIVE.value,
index=True,
)
max_users: Mapped[int] = mapped_column(Integer, nullable=False, default=5)
is_trial: Mapped[bool] = mapped_column(Boolean, nullable=False, default=True)
# 关系
users = relationship("User", back_populates="company")
field_mappings = relationship("FieldMapping", back_populates="company")
rules = relationship("CompanyRule", back_populates="company")
tasks = relationship("ReconciliationTask", back_populates="company")
def __repr__(self) -> str:
return f"<Company(id={self.id}, name='{self.name}', plan='{self.plan}')>"
+57
View File
@@ -0,0 +1,57 @@
"""
企业规则模型
用于存储企业自定义的映射规则,实现规则复用
"""
from datetime import datetime
from enum import Enum
from typing import Optional
from sqlalchemy import ForeignKey, Integer, String, Boolean, DateTime, JSON
from sqlalchemy.orm import Mapped, mapped_column, relationship
from app.models.base import BaseModel
class RuleType(str, Enum):
"""规则类型"""
FIELD_MAPPING = "FIELD_MAPPING" # 字段映射规则
ACCOUNT_MAPPING = "ACCOUNT_MAPPING" # 科目映射规则
DEPARTMENT_MAPPING = "DEPARTMENT_MAPPING" # 部门映射规则
class RuleStatus(str, Enum):
"""规则状态"""
ACTIVE = "ACTIVE" # 启用
INACTIVE = "INACTIVE" # 停用
class CompanyRule(BaseModel):
"""
企业规则模型
存储企业自定义的业务规则,用于自动应用
"""
__tablename__ = "company_rules"
company_id: Mapped[int] = mapped_column(Integer, ForeignKey("companies.id"), nullable=False, index=True)
rule_type: Mapped[str] = mapped_column(String(50), nullable=False, index=True, comment="规则类型")
match_condition: Mapped[dict] = mapped_column(JSON, nullable=False, comment="匹配条件")
target_value: Mapped[str] = mapped_column(String(100), nullable=False, comment="目标值")
priority: Mapped[int] = mapped_column(Integer, nullable=False, default=0, comment="优先级,数字越大优先级越高")
status: Mapped[str] = mapped_column(String(20), nullable=False, default=RuleStatus.ACTIVE.value, comment="状态")
description: Mapped[Optional[str]] = mapped_column(String(500), nullable=True, comment="规则描述")
match_count: Mapped[int] = mapped_column(Integer, nullable=False, default=0, comment="匹配次数")
last_used_at: Mapped[Optional[datetime]] = mapped_column(DateTime, nullable=True)
created_by: Mapped[Optional[int]] = mapped_column(Integer, ForeignKey("users.id"), nullable=True)
updated_by: Mapped[Optional[int]] = mapped_column(Integer, ForeignKey("users.id"), nullable=True)
# 关系
company = relationship("Company", back_populates="rules")
creator = relationship("User", foreign_keys=[created_by])
updater = relationship("User", foreign_keys=[updated_by])
def __repr__(self) -> str:
return f"<CompanyRule(type='{self.rule_type}', target='{self.target_value}', priority={self.priority})>"
+98
View File
@@ -0,0 +1,98 @@
"""
异常项模型
存储对账过程中检测到的异常
"""
from datetime import datetime
from enum import Enum
from typing import Optional
from sqlalchemy import ForeignKey, Integer, String, Text, DateTime, JSON, Float
from sqlalchemy.orm import Mapped, mapped_column, relationship
from app.models.base import BaseModel
class ExceptionType(str, Enum):
"""异常类型"""
# 第一阶段:工资表内部异常
MISSING_EMPLOYEE = "MISSING_EMPLOYEE" # 员工缺失(三个表之一缺失该员工)
AMOUNT_MISMATCH = "AMOUNT_MISMATCH" # 金额不匹配
ZERO_AMOUNT = "ZERO_AMOUNT" # 金额为零
NEGATIVE_AMOUNT = "NEGATIVE_AMOUNT" # 负数金额
UNUSUAL_AMOUNT = "UNUSUAL_AMOUNT" # 金额异常(过大或过小)
DUPLICATE_EMPLOYEE = "DUPLICATE_EMPLOYEE" # 重复员工
# 第二阶段:银行实发对账(预留)
BANK_MISMATCH = "BANK_MISMATCH" # 银行实发与工资表不匹配
BANK_MISSING = "BANK_MISSING" # 银行记录缺失
BANK_EXTRA = "BANK_EXTRA" # 银行多余记录
class ExceptionStatus(str, Enum):
"""异常状态"""
PENDING = "PENDING" # 待处理
CONFIRMED = "CONFIRMED" # 已确认
IGNORED = "IGNORED" # 已忽略
RESOLVED = "RESOLVED" # 已解决
class ExceptionSeverity(str, Enum):
"""异常严重程度"""
LOW = "LOW" # 低
MEDIUM = "MEDIUM" # 中
HIGH = "HIGH" # 高
CRITICAL = "CRITICAL" # 严重
class ExceptionItem(BaseModel):
"""
异常项模型
记录对账过程中检测到的各类异常
"""
__tablename__ = "exception_items"
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)
# 异常基本信息
exception_type: Mapped[str] = mapped_column(String(50), nullable=False, index=True, comment="异常类型")
severity: Mapped[str] = mapped_column(String(20), nullable=False, default=ExceptionSeverity.MEDIUM.value, comment="严重程度")
status: Mapped[str] = mapped_column(String(20), nullable=False, default=ExceptionStatus.PENDING.value, index=True, comment="状态")
# 关联的员工信息
employee_id: Mapped[Optional[str]] = mapped_column(String(100), nullable=True, index=True, comment="员工ID")
employee_name: Mapped[Optional[str]] = mapped_column(String(100), nullable=True, comment="员工姓名")
# 异常详情
description: Mapped[str] = mapped_column(Text, nullable=False, comment="异常描述")
detail: Mapped[Optional[dict]] = mapped_column(JSON, nullable=True, comment="详细信息")
# 涉及的金额
salary_amount: Mapped[Optional[float]] = mapped_column(Float, nullable=True, comment="工资表金额")
social_security_amount: Mapped[Optional[float]] = mapped_column(Float, nullable=True, comment="社保表金额")
tax_amount: Mapped[Optional[float]] = mapped_column(Float, nullable=True, comment="个税表金额")
bank_amount: Mapped[Optional[float]] = mapped_column(Float, nullable=True, comment="银行实发金额")
# 差异金额
difference_amount: Mapped[Optional[float]] = mapped_column(Float, nullable=True, comment="差异金额")
# 建议处理方式
suggested_action: Mapped[Optional[str]] = mapped_column(String(200), nullable=True, comment="建议处理方式")
ai_suggestion: Mapped[Optional[str]] = mapped_column(Text, nullable=True, comment="AI 分析建议")
# 处理信息
handled_by: Mapped[Optional[int]] = mapped_column(Integer, ForeignKey("users.id"), nullable=True)
handled_at: Mapped[Optional[datetime]] = mapped_column(DateTime, nullable=True)
handling_note: Mapped[Optional[str]] = mapped_column(Text, nullable=True, comment="处理备注")
# 关系
company = relationship("Company", backref="exceptions")
task = relationship("ReconciliationTask", backref="exceptions")
handler = relationship("User", foreign_keys=[handled_by])
def __repr__(self) -> str:
return f"<ExceptionItem(id={self.id}, type='{self.exception_type}', status='{self.status}')>"
+43
View File
@@ -0,0 +1,43 @@
"""
字段映射模型
用于存储 AI 识别的字段映射结果
"""
from datetime import datetime
from typing import List, Optional
from sqlalchemy import ForeignKey, Integer, String, Float, Boolean, DateTime, JSON, Text
from sqlalchemy.orm import Mapped, mapped_column, relationship
from app.models.base import BaseModel
class FieldMapping(BaseModel):
"""
字段映射模型
记录源字段到标准字段的映射关系
"""
__tablename__ = "field_mappings"
company_id: Mapped[int] = mapped_column(Integer, ForeignKey("companies.id"), nullable=False, index=True)
file_id: Mapped[int] = mapped_column(Integer, ForeignKey("uploaded_files.id"), nullable=False, index=True)
source_field: Mapped[str] = mapped_column(String(100), nullable=False, comment="源字段名")
standard_field: Mapped[str] = mapped_column(String(100), nullable=False, comment="标准字段名")
confidence: Mapped[float] = mapped_column(Float, nullable=False, default=0.0, comment="置信度 0-1")
confirmed: Mapped[bool] = mapped_column(Boolean, nullable=False, default=False, comment="是否已确认")
confirmed_by: Mapped[Optional[int]] = mapped_column(Integer, ForeignKey("users.id"), nullable=True)
confirmed_at: Mapped[Optional[datetime]] = mapped_column(DateTime, nullable=True)
sample_values: Mapped[Optional[dict]] = mapped_column(JSON, nullable=True, comment="样例值")
reasoning: Mapped[Optional[str]] = mapped_column(Text, nullable=True, comment="AI 判断依据")
is_skipped: Mapped[bool] = mapped_column(Boolean, nullable=False, default=False, comment="是否跳过")
# 关系
company = relationship("Company", back_populates="field_mappings")
file = relationship("UploadedFile", back_populates="field_mappings")
confirmed_by_user = relationship("User", foreign_keys=[confirmed_by])
def __repr__(self) -> str:
return f"<FieldMapping(source='{self.source_field}' -> standard='{self.standard_field}', confidence={self.confidence})>"
+73
View File
@@ -0,0 +1,73 @@
"""
对账任务模型
用于管理工资对账任务的生命周期
"""
from datetime import datetime
from enum import Enum
from typing import List, Optional
from sqlalchemy import ForeignKey, Integer, String, DateTime, JSON, Text
from sqlalchemy.orm import Mapped, mapped_column, relationship
from app.models.base import BaseModel
class TaskStatus(str, Enum):
"""任务状态"""
CREATED = "CREATED" # 已创建
FILES_UPLOADED = "FILES_UPLOADED" # 文件已上传
PARSING = "PARSING" # 解析中
WAITING_MAPPING_CONFIRM = "WAITING_MAPPING_CONFIRM" # 等待字段映射确认
MAPPING_CONFIRMED = "MAPPING_CONFIRMED" # 字段映射已确认
RECONCILING = "RECONCILING" # 对账中
COMPLETED = "COMPLETED" # 完成
FAILED = "FAILED" # 失败
class ReconciliationTask(BaseModel):
"""
对账任务模型
管理一次完整的工资对账流程
"""
__tablename__ = "reconciliation_tasks"
company_id: Mapped[int] = mapped_column(Integer, ForeignKey("companies.id"), nullable=False, index=True)
name: Mapped[str] = mapped_column(String(200), nullable=False, comment="任务名称")
period: Mapped[str] = mapped_column(String(20), nullable=False, index=True, comment="对账期间,如 2024-01")
status: Mapped[str] = mapped_column(String(30), nullable=False, default=TaskStatus.CREATED.value, index=True)
# 文件信息 (JSON 格式存储)
file_ids: Mapped[Optional[dict]] = mapped_column(JSON, nullable=True, comment="关联的文件ID")
# 字段映射状态
mapping_completed: Mapped[bool] = mapped_column(Integer, nullable=False, default=False)
mapping_confirmed_at: Mapped[Optional[datetime]] = mapped_column(DateTime, nullable=True)
mapping_confirmed_by: Mapped[Optional[int]] = mapped_column(Integer, ForeignKey("users.id"), nullable=True)
# 对账结果摘要
total_employees: Mapped[Optional[int]] = mapped_column(Integer, nullable=True)
matched_count: Mapped[Optional[int]] = mapped_column(Integer, nullable=True)
exception_count: Mapped[Optional[int]] = mapped_column(Integer, nullable=True)
# 结果数据 (JSON)
reconciliation_result: Mapped[Optional[dict]] = mapped_column(JSON, nullable=True)
# 错误信息
error_message: Mapped[Optional[str]] = mapped_column(Text, nullable=True)
# 确认信息
created_by: Mapped[Optional[int]] = mapped_column(Integer, ForeignKey("users.id"), nullable=True)
created_at: Mapped[datetime] = mapped_column(DateTime, nullable=False, default=datetime.utcnow)
updated_at: Mapped[datetime] = mapped_column(DateTime, nullable=False, default=datetime.utcnow, onupdate=datetime.utcnow)
# 关系
company = relationship("Company", back_populates="tasks")
creator = relationship("User", foreign_keys=[created_by])
confirm_user = relationship("User", foreign_keys=[mapping_confirmed_by])
def __repr__(self) -> str:
return f"<ReconciliationTask(id={self.id}, name='{self.name}', status='{self.status}')>"
+130
View File
@@ -0,0 +1,130 @@
"""
标准字段定义
定义工资表、社保表、个税表的标准字段
"""
from enum import Enum
class StandardField(str, Enum):
"""标准字段枚举"""
# ===== 公共字段 =====
EMPLOYEE_NAME = "员工姓名" # 员工姓名
EMPLOYEE_ID = "工号" # 工号
DEPARTMENT = "部门" # 部门
POSITION = "岗位" # 岗位/职位
# ===== 工资表字段 =====
BASE_SALARY = "基本工资" # 基本工资
BONUS = "奖金" # 奖金
ALLOWANCE = "补贴" # 补贴
OVERTIME_PAY = "加班费" # 加班费
DEDUCTION = "扣款" # 扣款
GROSS_SALARY = "应发工资" # 应发工资
NET_SALARY = "实发工资" # 实发工资/净工资
BANK_CARD = "银行账号" # 银行账号
ID_CARD = "身份证号" # 身份证号
# ===== 社保表字段 =====
SOCIAL_SECURITY_BASE = "社保基数" # 社保基数
PENSION_INSURANCE = "养老保险" # 养老保险(个人)
MEDICAL_INSURANCE = "医疗保险" # 医疗保险(个人)
UNEMPLOYMENT_INSURANCE = "失业保险" # 失业保险(个人)
HOUSING_FUND = "公积金" # 住房公积金(个人)
PENSION_INSURANCE_COMPANY = "养老保险(公司)" # 养老保险(公司)
MEDICAL_INSURANCE_COMPANY = "医疗保险(公司)" # 医疗保险(公司)
UNEMPLOYMENT_INSURANCE_COMPANY = "失业保险(公司)" # 失业保险(公司)
HOUSING_FUND_COMPANY = "公积金(公司)" # 住房公积金(公司)
SOCIAL_SECURITY_TOTAL = "社保合计" # 社保合计
# ===== 个税表字段 =====
TAXABLE_INCOME = "应税收入" # 应税收入
PRE_TAX_DEDUCTION = "税前扣除" # 税前扣除(三险一金等)
TAX_FREE_INCOME = "免税收入" # 免税收入
TAX_EXEMPT_INCOME = "税前减免" # 税前减免
QUICK_DEDUCTION = "速算扣除数" # 速算扣除数
TAX_AMOUNT = "应缴个税" # 应缴个人所得税
TAX_PAID = "已缴个税" # 已缴个人所得税
AFTER_TAX_INCOME = "税后收入" # 税后收入
# 标准字段分组
FIELD_GROUPS = {
"公共字段": [
StandardField.EMPLOYEE_NAME,
StandardField.EMPLOYEE_ID,
StandardField.DEPARTMENT,
StandardField.POSITION,
],
"工资表字段": [
StandardField.BASE_SALARY,
StandardField.BONUS,
StandardField.ALLOWANCE,
StandardField.OVERTIME_PAY,
StandardField.DEDUCTION,
StandardField.GROSS_SALARY,
StandardField.NET_SALARY,
StandardField.BANK_CARD,
StandardField.ID_CARD,
],
"社保表字段": [
StandardField.SOCIAL_SECURITY_BASE,
StandardField.PENSION_INSURANCE,
StandardField.MEDICAL_INSURANCE,
StandardField.UNEMPLOYMENT_INSURANCE,
StandardField.HOUSING_FUND,
StandardField.PENSION_INSURANCE_COMPANY,
StandardField.MEDICAL_INSURANCE_COMPANY,
StandardField.UNEMPLOYMENT_INSURANCE_COMPANY,
StandardField.HOUSING_FUND_COMPANY,
StandardField.SOCIAL_SECURITY_TOTAL,
],
"个税表字段": [
StandardField.TAXABLE_INCOME,
StandardField.PRE_TAX_DEDUCTION,
StandardField.TAX_FREE_INCOME,
StandardField.TAX_EXEMPT_INCOME,
StandardField.QUICK_DEDUCTION,
StandardField.TAX_AMOUNT,
StandardField.TAX_PAID,
StandardField.AFTER_TAX_INCOME,
],
}
# 字段类型映射
FIELD_TYPES = {
StandardField.EMPLOYEE_NAME.value: "string",
StandardField.EMPLOYEE_ID.value: "string",
StandardField.DEPARTMENT.value: "string",
StandardField.POSITION.value: "string",
StandardField.BASE_SALARY.value: "number",
StandardField.BONUS.value: "number",
StandardField.ALLOWANCE.value: "number",
StandardField.OVERTIME_PAY.value: "number",
StandardField.DEDUCTION.value: "number",
StandardField.GROSS_SALARY.value: "number",
StandardField.NET_SALARY.value: "number",
StandardField.BANK_CARD.value: "string",
StandardField.ID_CARD.value: "string",
StandardField.SOCIAL_SECURITY_BASE.value: "number",
StandardField.PENSION_INSURANCE.value: "number",
StandardField.MEDICAL_INSURANCE.value: "number",
StandardField.UNEMPLOYMENT_INSURANCE.value: "number",
StandardField.HOUSING_FUND.value: "number",
StandardField.PENSION_INSURANCE_COMPANY.value: "number",
StandardField.MEDICAL_INSURANCE_COMPANY.value: "number",
StandardField.UNEMPLOYMENT_INSURANCE_COMPANY.value: "number",
StandardField.HOUSING_FUND_COMPANY.value: "number",
StandardField.SOCIAL_SECURITY_TOTAL.value: "number",
StandardField.TAXABLE_INCOME.value: "number",
StandardField.PRE_TAX_DEDUCTION.value: "number",
StandardField.TAX_FREE_INCOME.value: "number",
StandardField.TAX_EXEMPT_INCOME.value: "number",
StandardField.QUICK_DEDUCTION.value: "number",
StandardField.TAX_AMOUNT.value: "number",
StandardField.TAX_PAID.value: "number",
StandardField.AFTER_TAX_INCOME.value: "number",
}
+59
View File
@@ -0,0 +1,59 @@
from enum import Enum
from typing import List
from sqlalchemy import ForeignKey, Integer, String, BigInteger, Text
from sqlalchemy.orm import Mapped, mapped_column, relationship
from app.models.base import BaseModel
class FileType(str, Enum):
"""文件类型"""
SALARY = "工资表"
SOCIAL_SECURITY = "社保表"
TAX = "个税表"
class ParseStatus(str, Enum):
"""解析状态"""
PENDING = "待解析"
PARSING = "解析中"
SUCCESS = "解析成功"
FAILED = "解析失败"
class UploadedFile(BaseModel):
"""
上传文件模型
Attributes:
company_id: 企业ID
task_id: 对账任务ID(可选)
file_type: 文件类型
original_filename: 原始文件名
stored_filename: 存储文件名
file_size: 文件大小(字节)
mime_type: MIME类型
parse_status: 解析状态
parse_error: 解析错误信息
"""
__tablename__ = "uploaded_files"
company_id: Mapped[int] = mapped_column(Integer, ForeignKey("companies.id"), nullable=False, index=True)
task_id: Mapped[int | None] = mapped_column(Integer, nullable=True, index=True)
file_type: Mapped[str] = mapped_column(String(50), nullable=False, index=True)
original_filename: Mapped[str] = mapped_column(String(255), nullable=False)
stored_filename: Mapped[str] = mapped_column(String(255), nullable=False, unique=True)
file_size: Mapped[int] = mapped_column(BigInteger, nullable=False)
mime_type: Mapped[str] = mapped_column(String(100), nullable=False)
parse_status: Mapped[str] = mapped_column(String(20), nullable=False, default=ParseStatus.PENDING.value, index=True)
parse_error: Mapped[str | None] = mapped_column(Text, nullable=True)
# 关系
field_mappings: Mapped[List["FieldMapping"]] = relationship(
"FieldMapping", back_populates="file", lazy="selectin"
)
def __repr__(self) -> str:
return f"<UploadedFile(id={self.id}, filename='{self.original_filename}', type='{self.file_type}')>"
+62
View File
@@ -0,0 +1,62 @@
from enum import Enum
from sqlalchemy import Boolean, ForeignKey, Integer, JSON, String, DateTime
from sqlalchemy.orm import Mapped, mapped_column, relationship
from datetime import datetime
from app.models.base import BaseModel
class UserRole(str, Enum):
"""用户角色"""
ADMIN = "管理员"
FINANCE_MANAGER = "财务主管"
ACCOUNTANT = "会计"
CASHIER = "出纳"
HR = "人事"
class UserStatus(str, Enum):
"""用户状态"""
ACTIVE = "正常"
DISABLED = "禁用"
LOCKED = "锁定"
class User(BaseModel):
"""
用户模型
Attributes:
company_id: 企业ID
email: 邮箱(登录账号)
hashed_password: 加密后的密码
full_name: 姓名
role: 角色
permissions: 权限列表(JSONB
status: 状态
last_login_at: 最后登录时间
"""
__tablename__ = "users"
company_id: Mapped[int] = mapped_column(Integer, ForeignKey("companies.id"), nullable=False, index=True)
email: Mapped[str] = mapped_column(String(255), nullable=False, unique=True, index=True)
hashed_password: Mapped[str] = mapped_column(String(255), nullable=False)
full_name: Mapped[str] = mapped_column(String(100), nullable=False)
role: Mapped[str] = mapped_column(String(50), nullable=False, default=UserRole.ACCOUNTANT.value)
permissions: Mapped[dict | None] = mapped_column(JSON, nullable=True)
status: Mapped[str] = mapped_column(String(20), nullable=False, default=UserStatus.ACTIVE.value, index=True)
last_login_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
# 关系
company = relationship("Company", back_populates="users")
def __repr__(self) -> str:
return f"<User(id={self.id}, email='{self.email}', role='{self.role}')>"
def has_permission(self, permission: str) -> bool:
"""检查用户是否有指定权限"""
if not self.permissions:
return False
return permission in self.permissions.get("permissions", [])
+45
View File
@@ -0,0 +1,45 @@
from datetime import datetime
from typing import Any, Optional
from pydantic import BaseModel, ConfigDict, Field
class AuditLogCreate(BaseModel):
"""创建审计日志的 Schema"""
company_id: int = Field(..., description="企业 ID")
user_id: Optional[int] = Field(None, description="用户 ID")
action: str = Field(..., max_length=50, description="操作类型")
resource_type: str = Field(..., max_length=100, description="资源类型")
resource_id: Optional[str] = Field(None, max_length=200, description="资源 ID")
details: Optional[dict[str, Any]] = Field(None, description="操作详情")
ip_address: Optional[str] = Field(None, max_length=50, description="IP 地址")
user_agent: Optional[str] = Field(None, description="User-Agent")
class AuditLogResponse(BaseModel):
"""审计日志响应 Schema"""
model_config = ConfigDict(from_attributes=True)
id: int
company_id: int
user_id: Optional[int]
action: str
resource_type: str
resource_id: Optional[str]
details: Optional[dict[str, Any]]
ip_address: Optional[str]
user_agent: Optional[str]
created_at: datetime
class AuditLogQuery(BaseModel):
"""审计日志查询参数"""
company_id: Optional[int] = None
user_id: Optional[int] = None
action: Optional[str] = None
resource_type: Optional[str] = None
resource_id: Optional[str] = None
start_date: Optional[datetime] = None
end_date: Optional[datetime] = None
skip: int = Field(default=0, ge=0)
limit: int = Field(default=50, ge=1, le=1000)
+38
View File
@@ -0,0 +1,38 @@
from datetime import datetime
from typing import Optional
from pydantic import BaseModel, ConfigDict, Field
class CompanyBase(BaseModel):
"""Company 基础 Schema"""
name: str = Field(..., min_length=1, max_length=200, description="企业名称")
plan: str = Field(default="免费版", description="套餐类型")
data_retention_months: int = Field(default=12, ge=1, le=120, description="数据保留月数")
max_users: int = Field(default=5, ge=1, description="最大用户数")
class CompanyCreate(CompanyBase):
"""创建 Company 的 Schema"""
pass
class CompanyUpdate(BaseModel):
"""更新 Company 的 Schema"""
name: Optional[str] = Field(None, min_length=1, max_length=200)
plan: Optional[str] = None
data_retention_months: Optional[int] = Field(None, ge=1, le=120)
status: Optional[str] = None
max_users: Optional[int] = Field(None, ge=1)
is_trial: Optional[bool] = None
class CompanyResponse(CompanyBase):
"""Company 响应 Schema"""
model_config = ConfigDict(from_attributes=True)
id: int
status: str
is_trial: bool
created_at: datetime
updated_at: datetime
+147
View File
@@ -0,0 +1,147 @@
"""
字段映射相关 Pydantic Schema
"""
from datetime import datetime
from typing import List, Optional, Dict, Any
from pydantic import BaseModel, Field
# ===== 字段映射 Schema =====
class FieldMappingBase(BaseModel):
"""字段映射基础 schema"""
source_field: str = Field(..., description="源字段名")
standard_field: str = Field(..., description="标准字段名")
confidence: float = Field(..., ge=0, le=1, description="置信度 0-1")
reasoning: Optional[str] = Field(None, description="AI 判断依据")
sample_values: Optional[List[Any]] = Field(None, description="样例值")
is_skipped: bool = Field(False, description="是否跳过")
class FieldMappingCreate(FieldMappingBase):
"""创建字段映射"""
file_id: int = Field(..., description="文件ID")
class FieldMappingUpdate(BaseModel):
"""更新字段映射"""
standard_field: Optional[str] = Field(None, description="标准字段名")
is_skipped: Optional[bool] = Field(None, description="是否跳过")
class FieldMappingResponse(FieldMappingBase):
"""字段映射响应"""
id: int
company_id: int
file_id: int
confirmed: bool
confirmed_by: Optional[int] = None
confirmed_at: Optional[datetime] = None
class Config:
from_attributes = True
class FieldMappingWithFileResponse(FieldMappingResponse):
"""带文件信息的字段映射响应"""
file_type: Optional[str] = None
file_name: Optional[str] = None
# ===== 批量操作 Schema =====
class ConfirmMappingsRequest(BaseModel):
"""批量确认映射请求"""
mapping_ids: List[int] = Field(..., description="要确认的映射ID列表")
class ConfirmMappingsResponse(BaseModel):
"""批量确认映射响应"""
confirmed_count: int = Field(..., description="确认数量")
mappings: List[FieldMappingResponse]
class SaveAsRuleRequest(BaseModel):
"""保存为规则请求"""
mapping_id: int = Field(..., description="映射ID")
class SaveAsRuleResponse(BaseModel):
"""保存为规则响应"""
rule_id: int
source_field: str
target_field: str
message: str
# ===== 规则 Schema =====
class CompanyRuleBase(BaseModel):
"""企业规则基础 schema"""
rule_type: str = Field(..., description="规则类型")
match_condition: Dict[str, Any] = Field(..., description="匹配条件")
target_value: str = Field(..., description="目标值")
priority: int = Field(0, description="优先级")
description: Optional[str] = Field(None, description="规则描述")
class CompanyRuleResponse(CompanyRuleBase):
"""企业规则响应"""
id: int
company_id: int
status: str
match_count: int
last_used_at: Optional[datetime] = None
created_at: datetime
updated_at: datetime
class Config:
from_attributes = True
class UpdateRuleStatusRequest(BaseModel):
"""更新规则状态请求"""
status: str = Field(..., description="状态 ACTIVE/INACTIVE")
class RecognizeFieldsRequest(BaseModel):
"""识别字段请求"""
file_id: int = Field(..., description="文件ID")
file_type: str = Field("工资表", description="文件类型")
class RecognizeFieldsResponse(BaseModel):
"""识别字段响应"""
file_id: int
mappings: List[FieldMappingResponse]
high_confidence_count: int = Field(..., description="高置信度数量")
medium_confidence_count: int = Field(..., description="中等置信度数量")
low_confidence_count: int = Field(..., description="低置信度数量")
needs_review: bool = Field(..., description="是否需要人工确认")
# ===== 文件映射概览 =====
class FileMappingSummary(BaseModel):
"""文件映射概览"""
file_id: int
file_type: str
file_name: str
total_fields: int
confirmed_fields: int
high_confidence: int
medium_confidence: int
low_confidence: int
mappings: List[FieldMappingResponse]
class TaskMappingOverview(BaseModel):
"""任务映射概览"""
task_id: int
total_files: int
total_fields: int
confirmed_fields: int
needs_review: bool
files: List[FileMappingSummary]
+56
View File
@@ -0,0 +1,56 @@
from datetime import datetime
from typing import Optional
from pydantic import BaseModel, ConfigDict, EmailStr, Field
class UserBase(BaseModel):
"""User 基础 Schema"""
email: EmailStr = Field(..., description="邮箱")
full_name: str = Field(..., min_length=1, max_length=100, description="姓名")
role: str = Field(default="会计", description="角色")
class UserCreate(UserBase):
"""创建用户的 Schema"""
password: str = Field(..., min_length=6, max_length=50, description="密码")
company_id: int = Field(..., description="企业ID")
class UserUpdate(BaseModel):
"""更新用户的 Schema"""
full_name: Optional[str] = Field(None, min_length=1, max_length=100)
role: Optional[str] = None
permissions: Optional[dict] = None
status: Optional[str] = None
class UserResponse(UserBase):
"""用户响应 Schema"""
model_config = ConfigDict(from_attributes=True)
id: int
company_id: int
status: str
permissions: Optional[dict]
last_login_at: Optional[datetime]
created_at: datetime
updated_at: datetime
class UserLogin(BaseModel):
"""用户登录 Schema"""
email: EmailStr = Field(..., description="邮箱")
password: str = Field(..., description="密码")
class Token(BaseModel):
"""Token Schema"""
access_token: str
token_type: str = "bearer"
class TokenData(BaseModel):
"""Token 数据 Schema"""
user_id: Optional[int] = None
email: Optional[str] = None
+225
View File
@@ -0,0 +1,225 @@
"""
AI 字段识别服务
使用 AI 自动识别 Excel 表头对应的标准字段
"""
import json
import re
from typing import List, Dict, Optional, Tuple
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy import select
from app.models.standard_field import StandardField, FIELD_TYPES
class AIFieldRecognizer:
"""AI 字段识别器"""
# 字段关键词映射(用于规则匹配兜底)
FIELD_KEYWORDS = {
StandardField.EMPLOYEE_NAME.value: ["姓名", "名字", "员工名", "name", "员工姓名"],
StandardField.EMPLOYEE_ID.value: ["工号", "员工号", "编号", "id", "员工编号"],
StandardField.DEPARTMENT.value: ["部门", "科室", "事业部", "department", "所属部门"],
StandardField.POSITION.value: ["岗位", "职位", "职务", "position", "job"],
StandardField.BASE_SALARY.value: ["基本工资", "岗位工资", "底薪", "base", "基本薪资"],
StandardField.BONUS.value: ["奖金", "绩效", "bonus", "绩效工资", "奖励"],
StandardField.ALLOWANCE.value: ["补贴", "津贴", "allowance", "餐补", "交通补贴"],
StandardField.OVERTIME_PAY.value: ["加班费", "加班工资", "overtime"],
StandardField.DEDUCTION.value: ["扣款", "扣除", "deduction", "罚款", "迟到扣款"],
StandardField.GROSS_SALARY.value: ["应发工资", "应发", "税前工资", "gross", "工资总额"],
StandardField.NET_SALARY.value: ["实发工资", "实发", "净工资", "net", "实发金额", "银行实发"],
StandardField.BANK_CARD.value: ["银行账号", "卡号", "账号", "bank", "银行卡"],
StandardField.ID_CARD.value: ["身份证", "证件号", "id_card", "身份证号"],
StandardField.SOCIAL_SECURITY_BASE.value: ["社保基数", "缴费基数", "基数"],
StandardField.PENSION_INSURANCE.value: ["养老保险", "养老", "pension", "养保"],
StandardField.MEDICAL_INSURANCE.value: ["医疗保险", "医保", "medical"],
StandardField.UNEMPLOYMENT_INSURANCE.value: ["失业保险", "失业", "unemployment"],
StandardField.HOUSING_FUND.value: ["公积金", "住房基金", "housing", "住房公金"],
StandardField.PENSION_INSURANCE_COMPANY.value: ["养老保险(公司)", "养老保险公司", "养保公司"],
StandardField.MEDICAL_INSURANCE_COMPANY.value: ["医疗保险(公司)", "医疗公司"],
StandardField.UNEMPLOYMENT_INSURANCE_COMPANY.value: ["失业保险(公司)", "失业公司"],
StandardField.HOUSING_FUND_COMPANY.value: ["公积金(公司)", "公积金公司"],
StandardField.SOCIAL_SECURITY_TOTAL.value: ["社保合计", "社保总计", "社保总额"],
StandardField.TAXABLE_INCOME.value: ["应税收入", "应税工资", "税前收入", "taxable"],
StandardField.PRE_TAX_DEDUCTION.value: ["税前扣除", "三险一金", "个人缴费"],
StandardField.TAX_FREE_INCOME.value: ["免税收入", "免税", "tax_free"],
StandardField.TAX_EXEMPT_INCOME.value: ["税前减免", "减免"],
StandardField.QUICK_DEDUCTION.value: ["速算扣除", "速算"],
StandardField.TAX_AMOUNT.value: ["应缴个税", "个人所得税", "个税", "tax"],
StandardField.TAX_PAID.value: ["已缴个税", "已扣税", "已缴税"],
StandardField.AFTER_TAX_INCOME.value: ["税后收入", "税后工资", "after_tax"],
}
def __init__(self, db: AsyncSession, company_id: int):
self.db = db
self.company_id = company_id
async def recognize_fields(
self,
headers: List[str],
sample_data: List[Dict],
file_type: str = "工资表"
) -> List[Dict]:
"""
识别字段映射
Args:
headers: 表头列表
sample_data: 样例数据(前10行)
file_type: 文件类型
Returns:
字段映射列表
"""
mappings = []
for header in headers:
# 获取该列的样例值
samples = [row.get(header) for row in sample_data if row.get(header)]
sample_values = samples[:5] if samples else []
# 先尝试规则匹配
matched_field, confidence, reasoning = self._rule_match(header, sample_values)
if matched_field:
mappings.append({
"source_field": header,
"standard_field": matched_field,
"confidence": confidence,
"reasoning": reasoning,
"sample_values": sample_values,
})
else:
# 无法匹配
mappings.append({
"source_field": header,
"standard_field": "",
"confidence": 0.0,
"reasoning": "无法识别字段类型",
"sample_values": sample_values,
})
return mappings
def _rule_match(
self,
header: str,
sample_values: List
) -> Tuple[Optional[str], float, str]:
"""
规则匹配
Args:
header: 字段名
sample_values: 样例值
Returns:
(标准字段, 置信度, 判断依据)
"""
header_lower = header.lower().strip()
header_normalized = header.strip()
best_match = None
best_confidence = 0.0
best_reasoning = ""
for standard_field, keywords in self.FIELD_KEYWORDS.items():
for keyword in keywords:
keyword_lower = keyword.lower()
# 精确匹配(完全相同)
if header_normalized == keyword or header_lower == keyword_lower:
return (
standard_field,
1.0,
f"字段名完全匹配:'{keyword}'"
)
# 包含匹配
if keyword_lower in header_lower or header_lower in keyword_lower:
confidence = 0.9
reasoning = f"字段名包含关键词:'{keyword}'"
# 数值字段检查样例值
field_type = FIELD_TYPES.get(standard_field, "string")
if field_type == "number" and sample_values:
if self._validate_numeric_samples(sample_values):
confidence = 0.95
reasoning += ",样例值验证为数值类型"
if confidence > best_confidence:
best_match = standard_field
best_confidence = confidence
best_reasoning = reasoning
# 模糊匹配(编辑距离)
distance = self._levenshtein_distance(header_lower, keyword_lower)
max_len = max(len(header_lower), len(keyword_lower))
similarity = 1 - (distance / max_len) if max_len > 0 else 0
if similarity > 0.7 and similarity > best_confidence:
best_match = standard_field
best_confidence = similarity * 0.8 # 模糊匹配降权
best_reasoning = f"字段名相似度:{similarity:.0%},参考词:'{keyword}'"
return best_match, best_confidence, best_reasoning
def _validate_numeric_samples(self, samples: List) -> bool:
"""验证样例值是否为数值"""
numeric_count = 0
for sample in samples[:5]:
if sample is None:
continue
sample_str = str(sample).strip()
# 移除常见的货币符号和逗号
sample_str = sample_str.replace("¥", "").replace(",", "").replace("", "")
try:
float(sample_str)
numeric_count += 1
except ValueError:
pass
return numeric_count >= len(samples) * 0.8
@staticmethod
def _levenshtein_distance(s1: str, s2: str) -> int:
"""计算编辑距离"""
if len(s1) < len(s2):
return AIFieldRecognizer._levenshtein_distance(s2, s1)
if len(s2) == 0:
return len(s1)
previous_row = range(len(s2) + 1)
for i, c1 in enumerate(s1):
current_row = [i + 1]
for j, c2 in enumerate(s2):
insertions = previous_row[j + 1] + 1
deletions = current_row[j] + 1
substitutions = previous_row[j] + (c1 != c2)
current_row.append(min(insertions, deletions, substitutions))
previous_row = current_row
return previous_row[-1]
async def recognize_fields_with_ai(
db: AsyncSession,
company_id: int,
headers: List[str],
sample_data: List[Dict],
file_type: str = "工资表"
) -> List[Dict]:
"""
使用 AI 识别字段(带重试的版本)
优先使用 OpenAI API,失败时降级到规则匹配
"""
recognizer = AIFieldRecognizer(db, company_id)
# 优先使用规则匹配(当前实现)
# TODO: 后续集成 OpenAI API
mappings = await recognizer.recognize_fields(headers, sample_data, file_type)
return mappings
+180
View File
@@ -0,0 +1,180 @@
from datetime import datetime
from functools import wraps
from typing import Any, Callable, Optional
from fastapi import Request
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.models.audit_log import AuditAction, AuditLog
from app.schemas.audit_log import AuditLogCreate, AuditLogQuery
class AuditService:
"""审计日志服务"""
@staticmethod
async def log_action(
db: AsyncSession,
company_id: int,
action: str,
resource_type: str,
user_id: Optional[int] = None,
resource_id: Optional[str] = None,
details: Optional[dict[str, Any]] = None,
ip_address: Optional[str] = None,
user_agent: Optional[str] = None,
) -> AuditLog:
"""
记录审计日志
Args:
db: 数据库会话
company_id: 企业 ID
action: 操作类型
resource_type: 资源类型
user_id: 用户 ID
resource_id: 资源 ID
details: 操作详情
ip_address: IP 地址
user_agent: User-Agent
Returns:
创建的审计日志对象
"""
log_data = AuditLogCreate(
company_id=company_id,
user_id=user_id,
action=action,
resource_type=resource_type,
resource_id=resource_id,
details=details,
ip_address=ip_address,
user_agent=user_agent,
)
audit_log = AuditLog(**log_data.model_dump())
db.add(audit_log)
await db.commit()
await db.refresh(audit_log)
return audit_log
@staticmethod
async def get_logs(
db: AsyncSession,
query_params: AuditLogQuery,
) -> list[AuditLog]:
"""
查询审计日志
Args:
db: 数据库会话
query_params: 查询参数
Returns:
审计日志列表
"""
stmt = select(AuditLog)
# 添加过滤条件
if query_params.company_id:
stmt = stmt.where(AuditLog.company_id == query_params.company_id)
if query_params.user_id:
stmt = stmt.where(AuditLog.user_id == query_params.user_id)
if query_params.action:
stmt = stmt.where(AuditLog.action == query_params.action)
if query_params.resource_type:
stmt = stmt.where(AuditLog.resource_type == query_params.resource_type)
if query_params.resource_id:
stmt = stmt.where(AuditLog.resource_id == query_params.resource_id)
if query_params.start_date:
stmt = stmt.where(AuditLog.created_at >= query_params.start_date)
if query_params.end_date:
stmt = stmt.where(AuditLog.created_at <= query_params.end_date)
# 排序和分页
stmt = stmt.order_by(AuditLog.created_at.desc())
stmt = stmt.offset(query_params.skip).limit(query_params.limit)
result = await db.execute(stmt)
return list(result.scalars().all())
@staticmethod
async def log_from_request(
db: AsyncSession,
request: Request,
company_id: int,
action: str,
resource_type: str,
user_id: Optional[int] = None,
resource_id: Optional[str] = None,
details: Optional[dict[str, Any]] = None,
) -> AuditLog:
"""
从 Request 对象中提取信息并记录审计日志
Args:
db: 数据库会话
request: FastAPI 请求对象
company_id: 企业 ID
action: 操作类型
resource_type: 资源类型
user_id: 用户 ID
resource_id: 资源 ID
details: 操作详情
Returns:
创建的审计日志对象
"""
# 提取 IP 地址
ip_address = request.client.host if request.client else None
# 提取 User-Agent
user_agent = request.headers.get("user-agent")
return await AuditService.log_action(
db=db,
company_id=company_id,
action=action,
resource_type=resource_type,
user_id=user_id,
resource_id=resource_id,
details=details,
ip_address=ip_address,
user_agent=user_agent,
)
# 单例实例
audit_service = AuditService()
def audit_log(
action: str,
resource_type: str,
get_resource_id: Optional[Callable] = None,
):
"""
审计日志装饰器
Args:
action: 操作类型
resource_type: 资源类型
get_resource_id: 从函数返回值中获取 resource_id 的函数
Example:
@audit_log(action=AuditAction.CREATE, resource_type="company")
async def create_company(...):
...
"""
def decorator(func: Callable):
@wraps(func)
async def wrapper(*args, **kwargs):
result = await func(*args, **kwargs)
# TODO: 在实现认证后,从上下文获取 company_id 和 user_id
# 目前暂时跳过实际记录
return result
return wrapper
return decorator
+158
View File
@@ -0,0 +1,158 @@
from datetime import datetime
from typing import Optional
from fastapi import Depends, HTTPException, status
from fastapi.security import OAuth2PasswordBearer
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.database import get_db
from app.core.security import create_access_token, hash_password, verify_password, decode_access_token
from app.models.user import User, UserStatus
from app.schemas.user import UserCreate, UserLogin, Token
oauth2_scheme = OAuth2PasswordBearer(tokenUrl="/api/auth/login")
class AuthService:
"""认证服务"""
@staticmethod
async def register_user(
db: AsyncSession,
user_data: UserCreate,
) -> User:
"""
注册用户
Args:
db: 数据库会话
user_data: 用户数据
Returns:
创建的用户对象
"""
# 检查邮箱是否已存在
result = await db.execute(select(User).where(User.email == user_data.email))
if result.scalar_one_or_none():
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="邮箱已被注册"
)
# 创建用户
hashed_password = hash_password(user_data.password)
user = User(
company_id=user_data.company_id,
email=user_data.email,
hashed_password=hashed_password,
full_name=user_data.full_name,
role=user_data.role,
)
db.add(user)
await db.commit()
await db.refresh(user)
return user
@staticmethod
async def authenticate_user(
db: AsyncSession,
login_data: UserLogin,
) -> tuple[User, str]:
"""
认证用户并生成 Token
Args:
db: 数据库会话
login_data: 登录数据
Returns:
用户对象和访问令牌
Raises:
HTTPException: 认证失败
"""
# 查找用户
result = await db.execute(select(User).where(User.email == login_data.email))
user = result.scalar_one_or_none()
if not user:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="邮箱或密码错误"
)
# 验证密码
if not verify_password(login_data.password, user.hashed_password):
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="邮箱或密码错误"
)
# 检查用户状态
if user.status != UserStatus.ACTIVE.value:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail=f"用户已被{user.status}"
)
# 更新最后登录时间
user.last_login_at = datetime.utcnow()
await db.commit()
await db.refresh(user)
# 生成 Token
access_token = create_access_token(data={"sub": str(user.id), "email": user.email})
return user, access_token
@staticmethod
async def get_current_user(
token: str = Depends(oauth2_scheme),
db: AsyncSession = Depends(get_db),
) -> User:
"""
获取当前用户
Args:
token: JWT token
db: 数据库会话
Returns:
当前用户对象
Raises:
HTTPException: Token 无效或用户不存在
"""
credentials_exception = HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="无法验证凭据",
headers={"WWW-Authenticate": "Bearer"},
)
payload = decode_access_token(token)
if not payload:
raise credentials_exception
user_id: str = payload.get("sub")
if not user_id:
raise credentials_exception
result = await db.execute(select(User).where(User.id == int(user_id)))
user = result.scalar_one_or_none()
if not user:
raise credentials_exception
if user.status != UserStatus.ACTIVE.value:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="用户已被禁用"
)
return user
# 单例实例
auth_service = AuthService()
+113
View File
@@ -0,0 +1,113 @@
from typing import Optional
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.models.company import Company, CompanyStatus
from app.schemas.company import CompanyCreate, CompanyUpdate
class CompanyService:
"""企业服务"""
@staticmethod
async def create_company(
db: AsyncSession,
company_data: CompanyCreate,
) -> Company:
"""
创建企业
Args:
db: 数据库会话
company_data: 企业数据
Returns:
创建的企业对象
"""
company = Company(**company_data.model_dump())
db.add(company)
await db.commit()
await db.refresh(company)
return company
@staticmethod
async def get_company(
db: AsyncSession,
company_id: int,
) -> Optional[Company]:
"""
获取企业信息
Args:
db: 数据库会话
company_id: 企业 ID
Returns:
企业对象,不存在则返回 None
"""
result = await db.execute(
select(Company).where(
Company.id == company_id,
Company.status != CompanyStatus.DELETED.value,
)
)
return result.scalar_one_or_none()
@staticmethod
async def update_company(
db: AsyncSession,
company_id: int,
company_data: CompanyUpdate,
) -> Optional[Company]:
"""
更新企业信息
Args:
db: 数据库会话
company_id: 企业 ID
company_data: 更新数据
Returns:
更新后的企业对象,不存在则返回 None
"""
company = await CompanyService.get_company(db, company_id)
if not company:
return None
update_data = company_data.model_dump(exclude_unset=True)
for field, value in update_data.items():
setattr(company, field, value)
await db.commit()
await db.refresh(company)
return company
@staticmethod
async def list_companies(
db: AsyncSession,
skip: int = 0,
limit: int = 100,
) -> list[Company]:
"""
获取企业列表
Args:
db: 数据库会话
skip: 跳过数量
limit: 返回数量
Returns:
企业列表
"""
result = await db.execute(
select(Company)
.where(Company.status != CompanyStatus.DELETED.value)
.offset(skip)
.limit(limit)
)
return list(result.scalars().all())
# 单例实例
company_service = CompanyService()
+399
View File
@@ -0,0 +1,399 @@
"""
数据清洗服务
清洗和标准化工资表数据,为对账做准备
"""
import re
from datetime import datetime
from decimal import Decimal, InvalidOperation
from typing import Any, Callable, Dict, List, Optional, Tuple, Union
from dataclasses import dataclass
@dataclass
class CleaningResult:
"""清洗结果"""
original_value: Any
cleaned_value: Any
is_valid: bool
error_message: Optional[str] = None
@dataclass
class ValidationRule:
"""验证规则"""
name: str
validate: Callable[[Any], bool]
error_message: str
class DataCleaner:
"""
数据清洗器
用于清洗和验证工资表数据
"""
def __init__(self):
self.validation_rules: Dict[str, List[ValidationRule]] = {}
self._init_default_rules()
def _init_default_rules(self):
"""初始化默认验证规则"""
# 姓名验证:2-20个中文字符或英文字母
self.add_rule(
"employee_name",
ValidationRule(
name="name_length",
validate=lambda v: bool(v) and 1 < len(str(v).strip()) <= 50,
error_message="姓名长度应在2-50个字符之间"
)
)
# 工号验证:数字或字母组合
self.add_rule(
"employee_id",
ValidationRule(
name="id_format",
validate=lambda v: bool(v) and bool(re.match(r'^[\w\-]+$', str(v))),
error_message="工号格式不正确"
)
)
# 金额验证:正数且在合理范围内
self.add_rule(
"amount",
ValidationRule(
name="amount_positive",
validate=lambda v: self._parse_number(v) is not None and self._parse_number(v) >= 0,
error_message="金额必须为非负数"
)
)
# 身份证号验证
self.add_rule(
"id_card",
ValidationRule(
name="id_card_length",
validate=lambda v: bool(v) and len(str(v).strip()) in [15, 18],
error_message="身份证号长度应为15或18位"
)
)
# 银行账号验证
self.add_rule(
"bank_card",
ValidationRule(
name="bank_card_digits",
validate=lambda v: bool(v) and str(v).isdigit() and len(str(v)) >= 10,
error_message="银行账号应为至少10位数字"
)
)
def add_rule(self, field_type: str, rule: ValidationRule):
"""添加验证规则"""
if field_type not in self.validation_rules:
self.validation_rules[field_type] = []
self.validation_rules[field_type].append(rule)
def _parse_number(self, value: Any) -> Optional[Decimal]:
"""解析数字"""
if value is None:
return None
if isinstance(value, (int, float, Decimal)):
return Decimal(str(value))
if isinstance(value, str):
# 移除常见的货币符号、逗号和空格
cleaned = value.strip()
cleaned = re.sub(r'[¥$,\s]', '', cleaned)
# 处理负数
is_negative = cleaned.startswith('-')
cleaned = cleaned.lstrip('-')
# 检查是否包含有效数字
if not cleaned or not re.match(r'^\d+(\.\d+)?$', cleaned):
return None
try:
result = Decimal(cleaned)
return -result if is_negative else result
except InvalidOperation:
return None
return None
def _parse_date(self, value: Any) -> Optional[datetime]:
"""解析日期"""
if value is None:
return None
if isinstance(value, datetime):
return value
if isinstance(value, str):
# 尝试多种日期格式
formats = [
'%Y-%m-%d',
'%Y/%m/%d',
'%Y年%m月%d',
'%Y%m%d',
'%Y-%m-%d %H:%M:%S',
]
for fmt in formats:
try:
return datetime.strptime(value.strip(), fmt)
except ValueError:
continue
return None
def clean_string(self, value: Any, strip: bool = True) -> str:
"""清洗字符串"""
if value is None:
return ""
result = str(value)
if strip:
result = result.strip()
return result
def clean_amount(self, value: Any) -> CleaningResult:
"""清洗金额字段"""
original = value
parsed = self._parse_number(value)
if parsed is None:
return CleaningResult(
original_value=original,
cleaned_value=None,
is_valid=False,
error_message="无法解析为数字"
)
return CleaningResult(
original_value=original,
cleaned_value=float(parsed),
is_valid=True
)
def clean_employee_name(self, value: Any) -> CleaningResult:
"""清洗员工姓名"""
cleaned = self.clean_string(value)
if not cleaned:
return CleaningResult(
original_value=value,
cleaned_value=None,
is_valid=False,
error_message="姓名为空"
)
# 移除多余空格
cleaned = re.sub(r'\s+', ' ', cleaned)
return CleaningResult(
original_value=value,
cleaned_value=cleaned,
is_valid=True
)
def clean_employee_id(self, value: Any) -> CleaningResult:
"""清洗工号"""
cleaned = self.clean_string(value)
if not cleaned:
return CleaningResult(
original_value=value,
cleaned_value=None,
is_valid=False,
error_message="工号为空"
)
return CleaningResult(
original_value=value,
cleaned_value=cleaned.upper(), # 统一大写
is_valid=True
)
def clean_bank_card(self, value: Any) -> CleaningResult:
"""清洗银行账号"""
original = value
cleaned = self.clean_string(value)
# 只保留数字
cleaned = re.sub(r'\D', '', cleaned)
if len(cleaned) < 10:
return CleaningResult(
original_value=original,
cleaned_value=None,
is_valid=False,
error_message="银行账号长度不足"
)
return CleaningResult(
original_value=original,
cleaned_value=cleaned,
is_valid=True
)
def clean_id_card(self, value: Any) -> CleaningResult:
"""清洗身份证号"""
original = value
cleaned = self.clean_string(value)
# 统一大写
cleaned = cleaned.upper()
# 只保留数字和X
cleaned = re.sub(r'[^0-9X]', '', cleaned)
if len(cleaned) not in [15, 18]:
return CleaningResult(
original_value=original,
cleaned_value=None,
is_valid=False,
error_message="身份证号长度应为15或18位"
)
return CleaningResult(
original_value=original,
cleaned_value=cleaned,
is_valid=True
)
def validate_field(self, field_type: str, value: Any) -> Tuple[bool, Optional[str]]:
"""
验证字段
Returns:
(是否通过, 错误信息)
"""
rules = self.validation_rules.get(field_type, [])
for rule in rules:
if not rule.validate(value):
return False, rule.error_message
return True, None
def clean_row(self, row: Dict[str, Any], field_mappings: Dict[str, str]) -> Dict[str, Any]:
"""
清洗一行数据
Args:
row: 原始行数据 {源字段名: 值}
field_mappings: 字段映射 {标准字段名: 源字段名}
Returns:
清洗后的数据 {标准字段名: 清洗后的值}
"""
cleaned = {}
errors = {}
# 反转映射:标准字段 -> 源字段
standard_to_source = {v: k for k, v in field_mappings.items()}
for standard_field, source_field in standard_to_source.items():
raw_value = row.get(source_field)
# 根据字段类型选择清洗方法
if standard_field in ["员工姓名", "姓名"]:
result = self.clean_employee_name(raw_value)
elif standard_field in ["工号", "员工号"]:
result = self.clean_employee_id(raw_value)
elif standard_field in ["银行账号", "卡号"]:
result = self.clean_bank_card(raw_value)
elif standard_field in ["身份证号", "证件号"]:
result = self.clean_id_card(raw_value)
elif standard_field in ["基本工资", "奖金", "补贴", "加班费", "扣款",
"应发工资", "实发工资", "社保基数",
"养老保险", "医疗保险", "失业保险", "公积金",
"个税", "应税收入", "税后收入"]:
result = self.clean_amount(raw_value)
else:
result = CleaningResult(
original_value=raw_value,
cleaned_value=self.clean_string(raw_value),
is_valid=bool(raw_value)
)
cleaned[standard_field] = result.cleaned_value
if not result.is_valid:
errors[standard_field] = result.error_message
return cleaned
def batch_clean(
self,
data: List[Dict[str, Any]],
field_mappings: Dict[str, str]
) -> Tuple[List[Dict[str, Any]], List[Dict[str, Any]]]:
"""
批量清洗数据
Returns:
(有效数据, 无效数据)
"""
valid_data = []
invalid_data = []
for row in data:
cleaned_row = self.clean_row(row, field_mappings)
# 检查是否有无效字段
has_invalid = any(v is None for v in cleaned_row.values())
if has_invalid:
invalid_data.append({
"original": row,
"cleaned": cleaned_row
})
else:
valid_data.append(cleaned_row)
return valid_data, invalid_data
# 单例实例
_cleaner_instance: Optional[DataCleaner] = None
def get_cleaner() -> DataCleaner:
"""获取数据清洗器单例"""
global _cleaner_instance
if _cleaner_instance is None:
_cleaner_instance = DataCleaner()
return _cleaner_instance
async def clean_salary_data(
data: List[Dict[str, Any]],
field_mappings: Dict[str, str]
) -> Dict[str, Any]:
"""
清洗工资数据
Args:
data: 原始工资数据
field_mappings: 字段映射
Returns:
清洗结果
"""
cleaner = get_cleaner()
valid_data, invalid_data = cleaner.batch_clean(data, field_mappings)
return {
"total_count": len(data),
"valid_count": len(valid_data),
"invalid_count": len(invalid_data),
"valid_data": valid_data,
"invalid_data": invalid_data,
"cleaning_time": datetime.utcnow().isoformat(),
}
+308
View File
@@ -0,0 +1,308 @@
"""
异常处理服务
处理异常的查询、更新、批量操作等
"""
from datetime import datetime
from typing import Any, Dict, List, Optional, Tuple
from sqlalchemy import select, update, func, and_, or_
from sqlalchemy.ext.asyncio import AsyncSession
from app.models.exception_item import ExceptionItem, ExceptionStatus
from app.models.reconciliation_task import ReconciliationTask
class ExceptionService:
"""异常处理服务"""
def __init__(self, db: AsyncSession):
self.db = db
async def get_exceptions(
self,
task_id: Optional[int] = None,
company_id: Optional[int] = None,
status: Optional[str] = None,
severity: Optional[str] = None,
exception_type: Optional[str] = None,
page: int = 1,
page_size: int = 20,
) -> Tuple[List[ExceptionItem], int]:
"""
查询异常列表
Args:
task_id: 任务ID
company_id: 企业ID
status: 状态
severity: 严重程度
exception_type: 异常类型
page: 页码
page_size: 每页数量
Returns:
(异常列表, 总数)
"""
from sqlalchemy.orm import selectinload
query = select(ExceptionItem).options(selectinload(ExceptionItem.handler))
count_query = select(func.count(ExceptionItem.id))
# 条件
filters = []
if task_id:
filters.append(ExceptionItem.task_id == task_id)
if company_id:
filters.append(ExceptionItem.company_id == company_id)
if status:
filters.append(ExceptionItem.status == status)
if severity:
filters.append(ExceptionItem.severity == severity)
if exception_type:
filters.append(ExceptionItem.exception_type == exception_type)
if filters:
query = query.where(and_(*filters))
count_query = count_query.where(and_(*filters))
# 总数
total_result = await self.db.execute(count_query)
total = total_result.scalar() or 0
# 分页
query = query.order_by(
ExceptionItem.created_at.desc()
).offset((page - 1) * page_size).limit(page_size)
result = await self.db.execute(query)
exceptions = result.scalars().all()
return list(exceptions), total
async def get_exception(self, exception_id: int) -> Optional[ExceptionItem]:
"""获取单个异常"""
return await self.db.get(ExceptionItem, exception_id)
async def update_exception(
self,
exception_id: int,
status: Optional[str] = None,
handling_note: Optional[str] = None,
handled_by: Optional[int] = None,
) -> Optional[ExceptionItem]:
"""
更新异常
Args:
exception_id: 异常ID
status: 新状态
handling_note: 处理备注
handled_by: 处理人ID
Returns:
更新后的异常
"""
exception = await self.db.get(ExceptionItem, exception_id)
if not exception:
return None
if status:
exception.status = status
if handling_note:
exception.handling_note = handling_note
if handled_by:
exception.handled_by = handled_by
exception.updated_at = datetime.utcnow()
await self.db.commit()
await self.db.refresh(exception)
return exception
async def batch_update_status(
self,
exception_ids: List[int],
status: str,
) -> int:
"""
批量更新状态
Args:
exception_ids: 异常ID列表
status: 新状态
Returns:
更新的数量
"""
if not exception_ids:
return 0
stmt = (
update(ExceptionItem)
.where(ExceptionItem.id.in_(exception_ids))
.values(
status=status,
updated_at=datetime.utcnow(),
)
)
result = await self.db.execute(stmt)
await self.db.commit()
return result.rowcount or 0
async def batch_resolve(
self,
exception_ids: List[int],
handling_note: str,
handled_by: Optional[int] = None,
) -> int:
"""
批量处理异常
Args:
exception_ids: 异常ID列表
handling_note: 处理备注
handled_by: 处理人ID
Returns:
处理的数量
"""
if not exception_ids:
return 0
update_data = {
"status": ExceptionStatus.RESOLVED.value,
"handling_note": handling_note,
"updated_at": datetime.utcnow(),
}
if handled_by:
update_data["handled_by"] = handled_by
stmt = (
update(ExceptionItem)
.where(ExceptionItem.id.in_(exception_ids))
.values(**update_data)
)
result = await self.db.execute(stmt)
await self.db.commit()
return result.rowcount or 0
async def get_exception_summary(
self,
task_id: Optional[int] = None,
company_id: Optional[int] = None,
) -> Dict[str, Any]:
"""
获取异常汇总统计
Args:
task_id: 任务ID
company_id: 企业ID
Returns:
汇总统计
"""
# 按状态统计
status_query = select(
ExceptionItem.status,
func.count(ExceptionItem.id).label("count")
).group_by(ExceptionItem.status)
# 按严重程度统计
severity_query = select(
ExceptionItem.severity,
func.count(ExceptionItem.id).label("count")
).group_by(ExceptionItem.severity)
# 按类型统计
type_query = select(
ExceptionItem.exception_type,
func.count(ExceptionItem.id).label("count")
).group_by(ExceptionItem.exception_type)
# 条件
filters = []
if task_id:
filters.append(ExceptionItem.task_id == task_id)
if company_id:
filters.append(ExceptionItem.company_id == company_id)
if filters:
status_query = status_query.where(and_(*filters))
severity_query = severity_query.where(and_(*filters))
type_query = type_query.where(and_(*filters))
# 执行查询
status_result = await self.db.execute(status_query)
severity_result = await self.db.execute(severity_query)
type_result = await self.db.execute(type_query)
return {
"by_status": {
row.status: row.count for row in status_result.all()
},
"by_severity": {
row.severity: row.count for row in severity_result.all()
},
"by_type": {
row.exception_type: row.count for row in type_result.all()
},
}
async def get_pending_exceptions(
self,
company_id: int,
limit: int = 10,
) -> List[ExceptionItem]:
"""获取待处理的异常"""
query = (
select(ExceptionItem)
.where(
and_(
ExceptionItem.company_id == company_id,
ExceptionItem.status == ExceptionStatus.PENDING.value,
)
)
.order_by(ExceptionItem.created_at.desc())
.limit(limit)
)
result = await self.db.execute(query)
return list(result.scalars().all())
# 便捷函数
async def get_exceptions(
db: AsyncSession,
**kwargs,
) -> Tuple[List[ExceptionItem], int]:
"""查询异常列表"""
service = ExceptionService(db)
return await service.get_exceptions(**kwargs)
async def update_exception(
db: AsyncSession,
exception_id: int,
**kwargs,
) -> Optional[ExceptionItem]:
"""更新异常"""
service = ExceptionService(db)
return await service.update_exception(exception_id, **kwargs)
async def batch_resolve_exceptions(
db: AsyncSession,
exception_ids: List[int],
handling_note: str,
handled_by: Optional[int] = None,
) -> int:
"""批量处理异常"""
service = ExceptionService(db)
return await service.batch_resolve(exception_ids, handling_note, handled_by)
+353
View File
@@ -0,0 +1,353 @@
"""
导出服务
导出对账结果为 Excel 和金蝶凭证格式
"""
import io
from datetime import datetime
from typing import Any, Dict, List, Optional
from sqlalchemy.ext.asyncio import AsyncSession
from app.models.reconciliation_task import ReconciliationTask
from app.models.exception_item import ExceptionItem
try:
from openpyxl import Workbook
from openpyxl.styles import Font, Alignment, Border, Side, PatternFill
from openpyxl.utils import get_column_letter
OPENPYXL_AVAILABLE = True
except ImportError:
OPENPYXL_AVAILABLE = False
class ExportService:
"""导出服务"""
def __init__(self, db: AsyncSession):
self.db = db
async def export_task_result(self, task_id: int) -> bytes:
"""
导出任务对账结果
Args:
task_id: 任务ID
Returns:
Excel 文件二进制数据
"""
if not OPENPYXL_AVAILABLE:
raise ImportError("openpyxl 未安装")
# 获取任务信息
task = await self.db.get(ReconciliationTask, task_id)
if not task:
raise ValueError("任务不存在")
# 获取异常列表
from sqlalchemy import select
query = select(ExceptionItem).where(ExceptionItem.task_id == task_id)
result = await self.db.execute(query)
exceptions = result.scalars().all()
# 创建工作簿
wb = Workbook()
# Sheet 1: 概览
ws_summary = wb.active
ws_summary.title = "概览"
self._fill_summary_sheet(ws_summary, task, exceptions)
# Sheet 2: 异常列表
ws_exceptions = wb.create_sheet("异常列表")
self._fill_exceptions_sheet(ws_exceptions, exceptions)
# Sheet 3: 异常类型统计
ws_stats = wb.create_sheet("异常统计")
self._fill_stats_sheet(ws_stats, exceptions)
# 保存到字节流
output = io.BytesIO()
wb.save(output)
output.seek(0)
return output.read()
def _fill_summary_sheet(
self,
ws,
task: ReconciliationTask,
exceptions: List[ExceptionItem],
):
"""填充概览Sheet"""
# 标题样式
title_font = Font(bold=True, size=14)
header_font = Font(bold=True)
header_fill = PatternFill(start_color="4472C4", end_color="4472C4", fill_type="solid")
header_font_white = Font(bold=True, color="FFFFFF")
thin_border = Border(
left=Side(style="thin"),
right=Side(style="thin"),
top=Side(style="thin"),
bottom=Side(style="thin"),
)
# 标题
ws["A1"] = f"对账结果汇总 - {task.name}"
ws["A1"].font = title_font
ws.merge_cells("A1:D1")
# 基本信息
ws["A3"] = "任务名称"
ws["B3"] = task.name
ws["A4"] = "对账期间"
ws["B4"] = task.period
ws["A5"] = "创建时间"
ws["B5"] = task.created_at.strftime("%Y-%m-%d %H:%M:%S") if task.created_at else ""
ws["A6"] = "完成时间"
ws["B6"] = task.completed_at.strftime("%Y-%m-%d %H:%M:%S") if task.completed_at else ""
# 统计数据
ws["A8"] = "统计数据"
ws["A8"].font = header_font
ws["A9"] = "总人数"
ws["B9"] = task.total_employees
ws["A10"] = "匹配人数"
ws["B10"] = task.matched_count
ws["A11"] = "异常数量"
ws["B11"] = task.exception_count
ws["A12"] = "匹配率"
ws["B12"] = f"{(task.matched_count / task.total_employees * 100):.1f}%" if task.total_employees > 0 else "0%"
# 设置列宽
ws.column_dimensions["A"].width = 15
ws.column_dimensions["B"].width = 30
def _fill_exceptions_sheet(self, ws, exceptions: List[ExceptionItem]):
"""填充异常列表Sheet"""
# 样式
header_fill = PatternFill(start_color="4472C4", end_color="4472C4", fill_type="solid")
header_font = Font(bold=True, color="FFFFFF")
thin_border = Border(
left=Side(style="thin"),
right=Side(style="thin"),
top=Side(style="thin"),
bottom=Side(style="thin"),
)
# 表头
headers = [
"异常ID", "员工ID", "员工姓名", "异常类型", "严重程度",
"状态", "描述", "工资", "社保", "个税", "银行", "差异金额",
"解决方案", "处理人", "创建时间"
]
for col, header in enumerate(headers, 1):
cell = ws.cell(row=1, column=col, value=header)
cell.fill = header_fill
cell.font = header_font
cell.border = thin_border
cell.alignment = Alignment(horizontal="center")
# 数据
for row_idx, exc in enumerate(exceptions, 2):
ws.cell(row=row_idx, column=1, value=exc.id).border = thin_border
ws.cell(row=row_idx, column=2, value=exc.employee_id).border = thin_border
ws.cell(row=row_idx, column=3, value=exc.employee_name).border = thin_border
ws.cell(row=row_idx, column=4, value=exc.exception_type).border = thin_border
ws.cell(row=row_idx, column=5, value=exc.severity).border = thin_border
ws.cell(row=row_idx, column=6, value=exc.status).border = thin_border
ws.cell(row=row_idx, column=7, value=exc.description).border = thin_border
# 金额(保留两位小数)
ws.cell(row=row_idx, column=8, value=exc.salary_amount).border = thin_border
ws.cell(row=row_idx, column=9, value=exc.social_security_amount).border = thin_border
ws.cell(row=row_idx, column=10, value=exc.tax_amount).border = thin_border
ws.cell(row=row_idx, column=11, value=exc.bank_amount).border = thin_border
ws.cell(row=row_idx, column=12, value=exc.difference_amount).border = thin_border
ws.cell(row=row_idx, column=13, value=exc.resolution).border = thin_border
ws.cell(row=row_idx, column=14, value=exc.handler).border = thin_border
ws.cell(row=row_idx, column=15, value=exc.created_at.strftime("%Y-%m-%d %H:%M:%S") if exc.created_at else "").border = thin_border
# 设置列宽
col_widths = [10, 15, 15, 20, 10, 10, 40, 12, 12, 12, 12, 12, 20, 15, 20]
for col, width in enumerate(col_widths, 1):
ws.column_dimensions[get_column_letter(col)].width = width
def _fill_stats_sheet(self, ws, exceptions: List[ExceptionItem]):
"""填充统计Sheet"""
from collections import Counter
header_fill = PatternFill(start_color="4472C4", end_color="4472C4", fill_type="solid")
header_font = Font(bold=True, color="FFFFFF")
thin_border = Border(
left=Side(style="thin"),
right=Side(style="thin"),
top=Side(style="thin"),
bottom=Side(style="thin"),
)
# 按类型统计
ws["A1"] = "按异常类型统计"
ws["A1"].font = Font(bold=True)
ws["A2"] = "类型"
ws["B2"] = "数量"
ws["A2"].fill = header_fill
ws["B2"].fill = header_fill
ws["A2"].font = header_font
ws["B2"].font = header_font
ws["A2"].border = thin_border
ws["B2"].border = thin_border
type_counts = Counter(e.exception_type for e in exceptions)
for row_idx, (exc_type, count) in enumerate(type_counts.items(), 3):
ws.cell(row=row_idx, column=1, value=exc_type).border = thin_border
ws.cell(row=row_idx, column=2, value=count).border = thin_border
# 按严重程度统计
start_row = len(type_counts) + 5
ws.cell(row=start_row, column=1, value="按严重程度统计").font = Font(bold=True)
ws.cell(row=start_row + 1, column=1, value="严重程度")
ws.cell(row=start_row + 1, column=2, value="数量")
ws.cell(row=start_row + 1, column=1).fill = header_fill
ws.cell(row=start_row + 1, column=2).fill = header_fill
ws.cell(row=start_row + 1, column=1).font = header_font
ws.cell(row=start_row + 1, column=2).font = header_font
severity_counts = Counter(e.severity for e in exceptions)
for row_idx, (severity, count) in enumerate(severity_counts.items(), start_row + 2):
ws.cell(row=row_idx, column=1, value=severity).border = thin_border
ws.cell(row=row_idx, column=2, value=count).border = thin_border
# 按状态统计
start_row = start_row + len(severity_counts) + 3
ws.cell(row=start_row, column=1, value="按状态统计").font = Font(bold=True)
ws.cell(row=start_row + 1, column=1, value="状态")
ws.cell(row=start_row + 1, column=2, value="数量")
ws.cell(row=start_row + 1, column=1).fill = header_fill
ws.cell(row=start_row + 1, column=2).fill = header_fill
ws.cell(row=start_row + 1, column=1).font = header_font
ws.cell(row=start_row + 1, column=2).font = header_font
status_counts = Counter(e.status for e in exceptions)
for row_idx, (status, count) in enumerate(status_counts.items(), start_row + 2):
ws.cell(row=row_idx, column=1, value=status).border = thin_border
ws.cell(row=row_idx, column=2, value=count).border = thin_border
async def export_kingdee_voucher(self, task_id: int) -> bytes:
"""
导出金蝶凭证
Args:
task_id: 任务ID
Returns:
Excel 文件二进制数据
"""
if not OPENPYXL_AVAILABLE:
raise ImportError("openpyxl 未安装")
# 获取任务信息
task = await self.db.get(ReconciliationTask, task_id)
if not task:
raise ValueError("任务不存在")
# 获取已解决的异常(排除已忽略的)
from sqlalchemy import select
query = select(ExceptionItem).where(
ExceptionItem.task_id == task_id,
ExceptionItem.status == "resolved"
)
result = await self.db.execute(query)
resolved_exceptions = result.scalars().all()
# 创建工作簿
wb = Workbook()
ws = wb.active
ws.title = "凭证数据"
# 样式
header_fill = PatternFill(start_color="4472C4", end_color="4472C4", fill_type="solid")
header_font = Font(bold=True, color="FFFFFF")
thin_border = Border(
left=Side(style="thin"),
right=Side(style="thin"),
top=Side(style="thin"),
bottom=Side(style="thin"),
)
# 金蝶凭证格式表头
headers = [
"凭证字", "凭证号", "凭证日期", "附单据数",
"摘要", "科目代码", "科目名称",
"借方金额", "贷方金额", "币种", "汇率"
]
for col, header in enumerate(headers, 1):
cell = ws.cell(row=1, column=col, value=header)
cell.fill = header_fill
cell.font = header_font
cell.border = thin_border
cell.alignment = Alignment(horizontal="center")
# 填充数据(根据金蝶凭证格式要求)
row_idx = 2
for exc in resolved_exceptions:
# 生成凭证摘要
summary = f"工资对账调整 - {exc.employee_name} ({exc.employee_id})"
# 如果有差异金额,生成借贷分录
if exc.difference_amount and exc.difference_amount != 0:
# 借方分录
ws.cell(row=row_idx, column=1, value="").border = thin_border
ws.cell(row=row_idx, column=2, value="").border = thin_border # 凭证号自动生成
ws.cell(row=row_idx, column=3, value=task.period + "-01").border = thin_border
ws.cell(row=row_idx, column=4, value=1).border = thin_border
ws.cell(row=row_idx, column=5, value=summary).border = thin_border
ws.cell(row=row_idx, column=6, value="").border = thin_border # 科目代码
ws.cell(row=row_idx, column=7, value="待处理").border = thin_border # 科目名称
if exc.difference_amount > 0:
ws.cell(row=row_idx, column=8, value=abs(exc.difference_amount)).border = thin_border
ws.cell(row=row_idx, column=9, value=0).border = thin_border
else:
ws.cell(row=row_idx, column=8, value=0).border = thin_border
ws.cell(row=row_idx, column=9, value=abs(exc.difference_amount)).border = thin_border
ws.cell(row=row_idx, column=10, value="人民币").border = thin_border
ws.cell(row=row_idx, column=11, value=1).border = thin_border
row_idx += 1
# 设置列宽
col_widths = [10, 12, 15, 10, 40, 15, 20, 15, 15, 10, 10]
for col, width in enumerate(col_widths, 1):
ws.column_dimensions[get_column_letter(col)].width = width
# 保存到字节流
output = io.BytesIO()
wb.save(output)
output.seek(0)
return output.read()
# 便捷函数
async def export_task_result(db: AsyncSession, task_id: int) -> bytes:
"""导出任务对账结果"""
service = ExportService(db)
return await service.export_task_result(task_id)
async def export_kingdee_voucher(db: AsyncSession, task_id: int) -> bytes:
"""导出金蝶凭证"""
service = ExportService(db)
return await service.export_kingdee_voucher(task_id)
+114
View File
@@ -0,0 +1,114 @@
import csv
from pathlib import Path
from typing import Any
import openpyxl
import pandas as pd
class FileParserService:
"""文件解析服务"""
@staticmethod
def parse_excel(file_path: str | Path) -> dict[str, Any]:
"""
解析 Excel 文件
Args:
file_path: 文件路径
Returns:
解析结果:headers, sample_rows, total_rows
"""
file_path = Path(file_path)
# 使用 openpyxl 解析 .xlsx
if file_path.suffix.lower() == ".xlsx":
wb = openpyxl.load_workbook(file_path, read_only=True, data_only=True)
ws = wb.active
# 获取表头(第一行)
headers = [cell.value for cell in ws[1]]
# 获取样例数据(前5行,不含表头)
sample_rows = []
for row_idx, row in enumerate(ws.iter_rows(min_row=2, max_row=6, values_only=True), start=2):
row_data = dict(zip(headers, row))
sample_rows.append(row_data)
# 获取总行数
total_rows = ws.max_row - 1 # 减去表头行
wb.close()
return {
"headers": headers,
"sample_rows": sample_rows,
"total_rows": total_rows,
}
# 使用 pandas 解析其他格式
else:
df = pd.read_excel(file_path)
return {
"headers": df.columns.tolist(),
"sample_rows": df.head(5).to_dict(orient="records"),
"total_rows": len(df),
}
@staticmethod
def parse_csv(file_path: str | Path) -> dict[str, Any]:
"""
解析 CSV 文件
Args:
file_path: 文件路径
Returns:
解析结果
"""
file_path = Path(file_path)
# 自动检测编码
encodings = ["utf-8", "gbk", "gb2312", "utf-8-sig"]
df = None
for encoding in encodings:
try:
df = pd.read_csv(file_path, encoding=encoding)
break
except (UnicodeDecodeError, Exception):
continue
if df is None:
raise ValueError("无法解析CSV文件,编码格式不支持")
return {
"headers": df.columns.tolist(),
"sample_rows": df.head(5).to_dict(orient="records"),
"total_rows": len(df),
}
@staticmethod
def detect_file_type(file_content: bytes) -> str:
"""
检测文件类型
Args:
file_content: 文件内容
Returns:
文件类型
"""
# 简单的文件类型检测
if file_content[:2] == b"PK":
return "xlsx"
elif file_content[:8] == b"\xd0\xcf\x11\xe0\xa1\xb1\x1a\xe1":
return "xls"
else:
return "csv"
# 单例实例
file_parser_service = FileParserService()
+87
View File
@@ -0,0 +1,87 @@
import os
import uuid
from pathlib import Path
from fastapi import UploadFile
from app.core.config import get_settings
settings = get_settings()
class FileStorageService:
"""文件存储服务"""
@staticmethod
def get_upload_dir() -> Path:
"""获取上传目录"""
upload_dir = Path(settings.upload_dir)
upload_dir.mkdir(parents=True, exist_ok=True)
return upload_dir
@staticmethod
async def save_file(file: UploadFile, company_id: int) -> tuple[str, int]:
"""
保存上传的文件
Args:
file: 上传的文件
company_id: 企业ID
Returns:
(存储文件名, 文件大小)
"""
# 创建企业专属目录
company_dir = FileStorageService.get_upload_dir() / str(company_id)
company_dir.mkdir(parents=True, exist_ok=True)
# 生成唯一文件名
file_ext = Path(file.filename or "file").suffix
stored_filename = f"{uuid.uuid4().hex}{file_ext}"
file_path = company_dir / stored_filename
# 保存文件
content = await file.read()
file_size = len(content)
with open(file_path, "wb") as f:
f.write(content)
return f"{company_id}/{stored_filename}", file_size
@staticmethod
def get_file_path(stored_filename: str) -> Path:
"""
获取文件完整路径
Args:
stored_filename: 存储文件名(格式: company_id/filename
Returns:
文件路径
"""
return FileStorageService.get_upload_dir() / stored_filename
@staticmethod
def delete_file(stored_filename: str) -> bool:
"""
删除文件
Args:
stored_filename: 存储文件名
Returns:
是否删除成功
"""
try:
file_path = FileStorageService.get_file_path(stored_filename)
if file_path.exists():
file_path.unlink()
return True
return False
except Exception:
return False
# 单例实例
file_storage_service = FileStorageService()
+314
View File
@@ -0,0 +1,314 @@
"""
字段映射服务
管理字段映射的创建、确认和规则沉淀
"""
from datetime import datetime
from typing import List, Optional, Dict
from sqlalchemy import select, and_, or_
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.orm import selectinload
from app.models.field_mapping import FieldMapping
from app.models.company_rule import CompanyRule, RuleType, RuleStatus
from app.models.uploaded_file import UploadedFile
from app.services.ai_recognizer import recognize_fields_with_ai
class MappingService:
"""字段映射服务"""
def __init__(self, db: AsyncSession):
self.db = db
async def recognize_and_save(
self,
company_id: int,
file_id: int,
file_type: str = "工资表"
) -> List[FieldMapping]:
"""
识别并保存字段映射
Args:
company_id: 企业ID
file_id: 文件ID
file_type: 文件类型
Returns:
字段映射列表
"""
# 1. 获取文件信息
file = await self.db.get(UploadedFile, file_id)
if not file:
raise ValueError(f"文件不存在: {file_id}")
# 2. 解析文件获取表头和样例数据
# TODO: 调用文件解析服务获取实际数据
# 暂时使用空数据
headers = []
sample_data = []
# 3. 调用 AI 识别
mappings_data = await recognize_fields_with_ai(
db=self.db,
company_id=company_id,
headers=headers,
sample_data=sample_data,
file_type=file_type
)
# 4. 保存映射
mappings = []
for data in mappings_data:
mapping = FieldMapping(
company_id=company_id,
file_id=file_id,
source_field=data["source_field"],
standard_field=data["standard_field"],
confidence=data["confidence"],
reasoning=data.get("reasoning"),
sample_values=data.get("sample_values"),
)
self.db.add(mapping)
mappings.append(mapping)
await self.db.commit()
for mapping in mappings:
await self.db.refresh(mapping)
return mappings
async def get_mappings_by_file(self, file_id: int) -> List[FieldMapping]:
"""获取文件的所有字段映射"""
result = await self.db.execute(
select(FieldMapping)
.where(FieldMapping.file_id == file_id)
.order_by(FieldMapping.id)
)
return list(result.scalars().all())
async def get_mappings_by_task(self, task_id: int, company_id: int) -> Dict[int, List[FieldMapping]]:
"""获取任务的所有字段映射(按文件分组)"""
result = await self.db.execute(
select(FieldMapping)
.where(FieldMapping.company_id == company_id)
.options(selectinload(FieldMapping.file))
.order_by(FieldMapping.file_id, FieldMapping.id)
)
mappings = list(result.scalars().all())
# 按文件ID分组
grouped: Dict[int, List[FieldMapping]] = {}
for mapping in mappings:
if mapping.file_id not in grouped:
grouped[mapping.file_id] = []
grouped[mapping.file_id].append(mapping)
return grouped
async def update_mapping(
self,
mapping_id: int,
standard_field: Optional[str] = None,
is_skipped: Optional[bool] = None
) -> Optional[FieldMapping]:
"""更新字段映射"""
mapping = await self.db.get(FieldMapping, mapping_id)
if not mapping:
return None
if standard_field is not None:
mapping.standard_field = standard_field
# 用户手动修改,重置置信度为 1.0
mapping.confidence = 1.0
if is_skipped is not None:
mapping.is_skipped = is_skipped
await self.db.commit()
await self.db.refresh(mapping)
return mapping
async def confirm_mapping(
self,
mapping_id: int,
user_id: int
) -> Optional[FieldMapping]:
"""确认字段映射"""
mapping = await self.db.get(FieldMapping, mapping_id)
if not mapping:
return None
mapping.confirmed = True
mapping.confirmed_by = user_id
mapping.confirmed_at = datetime.utcnow()
await self.db.commit()
await self.db.refresh(mapping)
return mapping
async def confirm_mappings(
self,
mapping_ids: List[int],
user_id: int
) -> List[FieldMapping]:
"""批量确认字段映射"""
confirmed = []
for mapping_id in mapping_ids:
mapping = await self.confirm_mapping(mapping_id, user_id)
if mapping:
confirmed.append(mapping)
return confirmed
async def save_as_rule(
self,
mapping: FieldMapping,
user_id: Optional[int] = None
) -> CompanyRule:
"""
将字段映射保存为企业规则
Args:
mapping: 字段映射
user_id: 用户ID
Returns:
创建的规则
"""
rule = CompanyRule(
company_id=mapping.company_id,
rule_type=RuleType.FIELD_MAPPING.value,
match_condition={
"source_field": mapping.source_field,
"file_type": mapping.standard_field.split("_")[0] if "_" in mapping.standard_field else "",
},
target_value=mapping.standard_field,
priority=0,
status=RuleStatus.ACTIVE.value,
description=f"字段映射规则:'{mapping.source_field}' -> '{mapping.standard_field}'",
created_by=user_id,
)
self.db.add(rule)
await self.db.commit()
await self.db.refresh(rule)
return rule
async def apply_rules(
self,
company_id: int,
headers: List[str]
) -> List[Dict]:
"""
应用企业规则进行字段匹配
Args:
company_id: 企业ID
headers: 表头列表
Returns:
匹配的字段映射列表
"""
# 1. 获取企业的所有字段映射规则
result = await self.db.execute(
select(CompanyRule)
.where(
and_(
CompanyRule.company_id == company_id,
CompanyRule.rule_type == RuleType.FIELD_MAPPING.value,
CompanyRule.status == RuleStatus.ACTIVE.value
)
)
.order_by(CompanyRule.priority.desc())
)
rules = list(result.scalars().all())
# 2. 构建规则索引
field_rules: Dict[str, CompanyRule] = {}
for rule in rules:
source_field = rule.match_condition.get("source_field", "")
if source_field:
field_rules[source_field] = rule
# 3. 匹配规则
mappings = []
for header in headers:
matched_rule = None
confidence = 0.0
# 精确匹配
if header in field_rules:
matched_rule = field_rules[header]
confidence = 1.0
# 模糊匹配
if not matched_rule:
for source_field, rule in field_rules.items():
if source_field in header or header in source_field:
if confidence < 0.9:
matched_rule = rule
confidence = 0.9
break
if matched_rule:
mappings.append({
"source_field": header,
"standard_field": matched_rule.target_value,
"confidence": confidence,
"reasoning": "规则命中" if confidence == 1.0 else "规则模糊匹配",
"is_rule_based": True,
})
else:
mappings.append({
"source_field": header,
"standard_field": "",
"confidence": 0.0,
"reasoning": "无匹配规则",
"is_rule_based": False,
})
return mappings
async def get_company_rules(
self,
company_id: int,
rule_type: Optional[str] = None
) -> List[CompanyRule]:
"""获取企业的规则列表"""
query = select(CompanyRule).where(CompanyRule.company_id == company_id)
if rule_type:
query = query.where(CompanyRule.rule_type == rule_type)
query = query.order_by(CompanyRule.priority.desc(), CompanyRule.created_at.desc())
result = await self.db.execute(query)
return list(result.scalars().all())
async def update_rule_status(
self,
rule_id: int,
status: str
) -> Optional[CompanyRule]:
"""更新规则状态"""
rule = await self.db.get(CompanyRule, rule_id)
if not rule:
return None
rule.status = status
await self.db.commit()
await self.db.refresh(rule)
return rule
async def delete_rule(self, rule_id: int) -> bool:
"""删除规则"""
rule = await self.db.get(CompanyRule, rule_id)
if not rule:
return False
await self.db.delete(rule)
await self.db.commit()
return True
@@ -0,0 +1,335 @@
"""
对账引擎
执行对账逻辑,协调各个规则
"""
import time
from dataclasses import dataclass, field
from datetime import datetime
from typing import Any, Dict, List, Optional, Tuple
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy import select
from app.models.exception_item import ExceptionItem, ExceptionStatus
from app.services.reconciliation.rules import (
RuleContext,
RuleResult,
RuleFactory,
ReconciliationRule,
)
@dataclass
class ReconciliationResult:
"""对账结果"""
task_id: int
period: str
total_employees: int
matched_employees: int
exception_count: int
# 统计
exceptions_by_type: Dict[str, int] = field(default_factory=dict)
exceptions_by_severity: Dict[str, int] = field(default_factory=dict)
# 详细结果
matched_employees_list: List[Dict[str, Any]] = field(default_factory=list)
exception_items: List[Dict[str, Any]] = field(default_factory=list)
# 性能指标
execution_time_ms: float = 0
rules_executed: int = 0
# 时间戳
completed_at: str = ""
class ReconciliationEngine:
"""
对账引擎
协调数据加载、规则执行、结果存储
"""
def __init__(
self,
db: AsyncSession,
company_id: int,
task_id: int,
period: str,
enable_bank_rules: bool = False,
):
self.db = db
self.company_id = company_id
self.task_id = task_id
self.period = period
self.enable_bank_rules = enable_bank_rules
# 创建规则
self.rules = RuleFactory.create_all_rules(enable_bank_rules)
# 数据存储
self.salary_data: Dict[str, Dict[str, Any]] = {}
self.social_security_data: Dict[str, Dict[str, Any]] = {}
self.tax_data: Dict[str, Dict[str, Any]] = {}
self.bank_data: Optional[Dict[str, Dict[str, Any]]] = None
# 结果
self.exceptions: List[ExceptionItem] = []
self.matched_count = 0
self.exception_count = 0
def load_data(
self,
salary_records: List[Dict[str, Any]],
social_security_records: List[Dict[str, Any]],
tax_records: List[Dict[str, Any]],
bank_records: Optional[List[Dict[str, Any]]] = None,
employee_id_field: str = "employee_id",
employee_name_field: str = "employee_name",
):
"""
加载数据
Args:
salary_records: 工资表数据
social_security_records: 社保表数据
tax_records: 个税表数据
bank_records: 银行数据(可选)
employee_id_field: 员工ID字段名
employee_name_field: 员工姓名字段名
"""
# 加载工资数据
for record in salary_records:
emp_id = str(record.get(employee_id_field, ""))
if emp_id:
self.salary_data[emp_id] = record
# 加载社保数据
for record in social_security_records:
emp_id = str(record.get(employee_id_field, ""))
if emp_id:
self.social_security_data[emp_id] = record
# 加载个税数据
for record in tax_records:
emp_id = str(record.get(employee_id_field, ""))
if emp_id:
self.tax_data[emp_id] = record
# 加载银行数据
if bank_records:
self.bank_data = {}
for record in bank_records:
emp_id = str(record.get(employee_id_field, ""))
if emp_id:
self.bank_data[emp_id] = record
async def execute(self) -> ReconciliationResult:
"""
执行对账
Returns:
对账结果
"""
start_time = time.time()
# 获取所有员工ID(合并三个表)
all_employee_ids = set()
all_employee_ids.update(self.salary_data.keys())
all_employee_ids.update(self.social_security_data.keys())
all_employee_ids.update(self.tax_data.keys())
if self.bank_data:
all_employee_ids.update(self.bank_data.keys())
total_employees = len(all_employee_ids)
matched_employees = 0
exception_items: List[Dict[str, Any]] = []
# 统计
exceptions_by_type: Dict[str, int] = {}
exceptions_by_severity: Dict[str, int] = {}
matched_list: List[Dict[str, Any]] = []
# 重置需要状态的规则
for rule in self.rules:
if hasattr(rule, 'reset'):
rule.reset()
# 遍历每个员工
for employee_id in all_employee_ids:
# 获取员工姓名
employee_name = (
self.salary_data.get(employee_id, {}).get("employee_name") or
self.social_security_data.get(employee_id, {}).get("employee_name") or
self.tax_data.get(employee_id, {}).get("employee_name") or
employee_id
)
# 构建规则上下文
context = RuleContext(
company_id=self.company_id,
task_id=self.task_id,
period=self.period,
employee_id=employee_id,
employee_name=employee_name,
salary_data=self.salary_data,
social_security_data=self.social_security_data,
tax_data=self.tax_data,
bank_data=self.bank_data,
)
# 执行所有规则
employee_exceptions = []
for rule in self.rules:
if not rule.enabled:
continue
try:
result = rule.check(context)
if result.is_exception:
employee_exceptions.append((rule, result))
except Exception as e:
# 规则执行出错,记录但不中断
pass
# 处理该员工的异常
if employee_exceptions:
for rule, result in employee_exceptions:
# 保存异常到数据库
exception_item = await self._save_exception(rule, result, context)
self.exceptions.append(exception_item)
exception_items.append({
"id": exception_item.id,
"type": result.exception_type.value if result.exception_type else rule.exception_type.value,
"severity": result.severity.value,
"description": result.description,
"employee_id": employee_id,
"employee_name": employee_name,
})
# 统计
type_key = result.exception_type.value if result.exception_type else rule.exception_type.value
exceptions_by_type[type_key] = exceptions_by_type.get(type_key, 0) + 1
exceptions_by_severity[result.severity.value] = (
exceptions_by_severity.get(result.severity.value, 0) + 1
)
else:
matched_employees += 1
matched_list.append({
"employee_id": employee_id,
"employee_name": employee_name,
})
# 更新任务统计
await self._update_task_stats(matched_employees, len(exception_items))
execution_time = (time.time() - start_time) * 1000
return ReconciliationResult(
task_id=self.task_id,
period=self.period,
total_employees=total_employees,
matched_employees=matched_employees,
exception_count=len(exception_items),
exceptions_by_type=exceptions_by_type,
exceptions_by_severity=exceptions_by_severity,
matched_employees_list=matched_list,
exception_items=exception_items,
execution_time_ms=execution_time,
rules_executed=len([r for r in self.rules if r.enabled]),
completed_at=datetime.utcnow().isoformat(),
)
async def _save_exception(
self,
rule: ReconciliationRule,
result: RuleResult,
context: RuleContext,
) -> ExceptionItem:
"""保存异常到数据库"""
exception_item = ExceptionItem(
company_id=self.company_id,
task_id=self.task_id,
exception_type=result.exception_type.value if result.exception_type else rule.exception_type.value,
severity=result.severity.value,
status=ExceptionStatus.PENDING.value,
employee_id=context.employee_id,
employee_name=context.employee_name,
description=result.description,
detail=result.detail,
salary_amount=result.salary_amount,
social_security_amount=result.social_security_amount,
tax_amount=result.tax_amount,
bank_amount=result.bank_amount,
difference_amount=result.difference_amount,
suggested_action=result.suggested_action,
)
self.db.add(exception_item)
await self.db.commit()
await self.db.refresh(exception_item)
return exception_item
async def _update_task_stats(self, matched: int, exceptions: int):
"""更新任务统计"""
from app.models.reconciliation_task import ReconciliationTask
task = await self.db.get(ReconciliationTask, self.task_id)
if task:
task.matched_count = matched
task.exception_count = exceptions
task.total_employees = matched + exceptions
task.status = "COMPLETED"
await self.db.commit()
# 便捷函数
async def run_reconciliation(
db: AsyncSession,
company_id: int,
task_id: int,
period: str,
salary_records: List[Dict[str, Any]],
social_security_records: List[Dict[str, Any]],
tax_records: List[Dict[str, Any]],
bank_records: Optional[List[Dict[str, Any]]] = None,
enable_bank_rules: bool = False,
) -> ReconciliationResult:
"""
运行对账
Args:
db: 数据库会话
company_id: 企业ID
task_id: 任务ID
period: 对账期间
salary_records: 工资表数据
social_security_records: 社保表数据
tax_records: 个税表数据
bank_records: 银行数据(可选)
enable_bank_rules: 是否启用银行规则
Returns:
对账结果
"""
engine = ReconciliationEngine(
db=db,
company_id=company_id,
task_id=task_id,
period=period,
enable_bank_rules=enable_bank_rules,
)
engine.load_data(
salary_records=salary_records,
social_security_records=social_security_records,
tax_records=tax_records,
bank_records=bank_records,
)
return await engine.execute()
@@ -0,0 +1,459 @@
"""
对账规则定义
实现 7 类异常检测规则(第一阶段)+ 预留银行实发对账规则(第二阶段)
"""
from abc import ABC, abstractmethod
from dataclasses import dataclass, field
from datetime import datetime
from typing import Any, Callable, Dict, List, Optional, Tuple
from app.models.exception_item import ExceptionType, ExceptionSeverity
@dataclass
class RuleContext:
"""规则上下文"""
company_id: int
task_id: int
period: str
employee_id: str
employee_name: str
salary_data: Dict[str, Any]
social_security_data: Dict[str, Any]
tax_data: Dict[str, Any]
bank_data: Optional[Dict[str, Any]] = None
@dataclass
class RuleResult:
"""规则检测结果"""
is_exception: bool
exception_type: Optional[ExceptionType] = None
severity: ExceptionSeverity = ExceptionSeverity.MEDIUM
description: str = ""
detail: Dict[str, Any] = field(default_factory=dict)
salary_amount: Optional[float] = None
social_security_amount: Optional[float] = None
tax_amount: Optional[float] = None
bank_amount: Optional[float] = None
difference_amount: Optional[float] = None
suggested_action: str = ""
class ReconciliationRule(ABC):
"""对账规则基类"""
def __init__(self, enabled: bool = True):
self.enabled = enabled
@property
@abstractmethod
def name(self) -> str:
"""规则名称"""
pass
@property
@abstractmethod
def exception_type(self) -> ExceptionType:
"""异常类型"""
pass
@property
def severity(self) -> ExceptionSeverity:
"""默认严重程度"""
return ExceptionSeverity.MEDIUM
@abstractmethod
def check(self, context: RuleContext) -> RuleResult:
"""执行规则检查"""
pass
class MissingEmployeeRule(ReconciliationRule):
"""员工缺失规则"""
@property
def name(self) -> str:
return "员工缺失检测"
@property
def exception_type(self) -> ExceptionType:
return ExceptionType.MISSING_EMPLOYEE
def check(self, context: RuleContext) -> RuleResult:
"""检测员工是否在所有相关表中都存在"""
has_salary = bool(context.employee_id in context.salary_data)
has_social = bool(context.employee_id in context.social_security_data)
has_tax = bool(context.employee_id in context.tax_data)
# 至少需要在两个表中出现
present_count = sum([has_salary, has_social, has_tax])
if present_count == 0:
return RuleResult(
is_exception=False, # 不在任何一个表中,可能是新增员工
description="员工不在任何表中"
)
if present_count < 3:
missing_tables = []
if not has_salary:
missing_tables.append("工资表")
if not has_social:
missing_tables.append("社保表")
if not has_tax:
missing_tables.append("个税表")
return RuleResult(
is_exception=True,
exception_type=self.exception_type,
severity=ExceptionSeverity.HIGH,
description=f"员工在 {', '.join(missing_tables)} 中缺失",
detail={
"has_salary": has_salary,
"has_social_security": has_social,
"has_tax": has_tax,
},
suggested_action="确认该员工是否应该存在于此期间"
)
return RuleResult(is_exception=False)
class AmountMismatchRule(ReconciliationRule):
"""金额不匹配规则"""
@property
def name(self) -> str:
return "金额不匹配检测"
@property
def exception_type(self) -> ExceptionType:
return ExceptionType.AMOUNT_MISMATCH
def check(self, context: RuleContext) -> RuleResult:
"""检测工资、社保、个税之间的金额一致性"""
salary_record = context.salary_data.get(context.employee_id, {})
social_record = context.social_security_data.get(context.employee_id, {})
tax_record = context.tax_data.get(context.employee_id, {})
salary_net = salary_record.get("实发工资", 0) or salary_record.get("net_salary", 0)
tax_amount = tax_record.get("应缴个税", 0) or tax_record.get("tax", 0)
social_total = self._calc_social_total(social_record)
# 应发工资 - 个税 - 社保 = 实发工资(理论上)
salary_gross = salary_record.get("应发工资", 0) or salary_record.get("gross_salary", 0)
expected_net = salary_gross - tax_amount - social_total
# 允许一定误差(0.01元)
tolerance = 0.01
diff = abs(salary_net - expected_net)
if diff > tolerance:
return RuleResult(
is_exception=True,
exception_type=self.exception_type,
severity=ExceptionSeverity.HIGH,
description=f"实发工资与计算值不一致(差异: {diff:.2f}元)",
detail={
"salary_net": salary_net,
"expected_net": expected_net,
"salary_gross": salary_gross,
"tax_amount": tax_amount,
"social_total": social_total,
},
difference_amount=diff,
salary_amount=salary_net,
tax_amount=tax_amount,
social_security_amount=social_total,
suggested_action="检查工资表、社保表、个税表数据是否正确"
)
return RuleResult(is_exception=False)
def _calc_social_total(self, record: Dict) -> float:
"""计算社保合计"""
keys = ["养老保险", "医疗保险", "失业保险", "公积金", "社保合计"]
for key in keys:
if key in record:
return float(record[key] or 0)
# 如果没有合计,手动计算
total = 0
for key in ["养老保险", "医疗保险", "失业保险", "公积金"]:
total += float(record.get(key, 0) or 0)
return total
class ZeroAmountRule(ReconciliationRule):
"""零金额规则"""
@property
def name(self) -> str:
return "零金额检测"
@property
def exception_type(self) -> ExceptionType:
return ExceptionType.ZERO_AMOUNT
def check(self, context: RuleContext) -> RuleResult:
"""检测实发工资为零的情况"""
salary_record = context.salary_data.get(context.employee_id, {})
net_salary = salary_record.get("实发工资", 0) or salary_record.get("net_salary", 0)
if net_salary == 0:
return RuleResult(
is_exception=True,
exception_type=self.exception_type,
severity=ExceptionSeverity.MEDIUM,
description="实发工资为零",
detail={"net_salary": 0},
salary_amount=0,
suggested_action="确认员工是否离职或暂停发放"
)
return RuleResult(is_exception=False)
class NegativeAmountRule(ReconciliationRule):
"""负数金额规则"""
@property
def name(self) -> str:
return "负数金额检测"
@property
def exception_type(self) -> ExceptionType:
return ExceptionType.NEGATIVE_AMOUNT
def check(self, context: RuleContext) -> RuleResult:
"""检测负数金额"""
salary_record = context.salary_data.get(context.employee_id, {})
negative_fields = []
for field_name in ["实发工资", "应发工资", "基本工资", "奖金", "补贴"]:
value = salary_record.get(field_name, 0) or 0
if value < 0:
negative_fields.append(f"{field_name}: {value}")
if negative_fields:
return RuleResult(
is_exception=True,
exception_type=self.exception_type,
severity=ExceptionSeverity.CRITICAL,
description=f"发现负数金额: {', '.join(negative_fields)}",
detail={"negative_fields": negative_fields},
suggested_action="检查数据录入是否有误"
)
return RuleResult(is_exception=False)
class UnusualAmountRule(ReconciliationRule):
"""金额异常规则"""
def __init__(self, max_salary: float = 1000000, min_salary: float = 0):
super().__init__()
self.max_salary = max_salary
self.min_salary = min_salary
@property
def name(self) -> str:
return "金额异常检测"
@property
def exception_type(self) -> ExceptionType:
return ExceptionType.UNUSUAL_AMOUNT
def check(self, context: RuleContext) -> RuleResult:
"""检测金额异常(过大或过小)"""
salary_record = context.salary_data.get(context.employee_id, {})
net_salary = salary_record.get("实发工资", 0) or salary_record.get("net_salary", 0)
if net_salary > self.max_salary:
return RuleResult(
is_exception=True,
exception_type=self.exception_type,
severity=ExceptionSeverity.HIGH,
description=f"实发工资金额异常偏高: {net_salary:.2f}",
detail={"net_salary": net_salary, "max_allowed": self.max_salary},
salary_amount=net_salary,
suggested_action="确认是否为高管或特殊人员"
)
if 0 < net_salary < self.min_salary:
return RuleResult(
is_exception=True,
exception_type=self.exception_type,
severity=ExceptionSeverity.LOW,
description=f"实发工资金额异常偏低: {net_salary:.2f}",
detail={"net_salary": net_salary, "min_expected": self.min_salary},
salary_amount=net_salary,
suggested_action="确认是否为基础工资人员"
)
return RuleResult(is_exception=False)
class DuplicateEmployeeRule(ReconciliationRule):
"""重复员工规则"""
def __init__(self):
super().__init__()
self._seen_employees: Dict[str, List[str]] = {} # employee_id -> [task_ids]
def reset(self):
"""重置状态"""
self._seen_employees = {}
@property
def name(self) -> str:
return "重复员工检测"
@property
def exception_type(self) -> ExceptionType:
return ExceptionType.DUPLICATE_EMPLOYEE
def check(self, context: RuleContext) -> RuleResult:
"""检测同一任务中是否有重复员工"""
employee_key = f"{context.task_id}:{context.employee_id}"
if employee_key in self._seen_employees:
return RuleResult(
is_exception=True,
exception_type=self.exception_type,
severity=ExceptionSeverity.CRITICAL,
description=f"员工 {context.employee_name} (ID: {context.employee_id}) 在本任务中重复出现",
detail={"occurrences": len(self._seen_employees[employee_key]) + 1},
suggested_action="检查并删除重复记录"
)
self._seen_employees[employee_key] = [context.task_id]
return RuleResult(is_exception=False)
# 第二阶段预留规则
class BankMismatchRule(ReconciliationRule):
"""银行实发不匹配规则"""
def __init__(self, enabled: bool = False):
super().__init__(enabled=enabled)
@property
def name(self) -> str:
return "银行实发不匹配检测"
@property
def exception_type(self) -> ExceptionType:
return ExceptionType.BANK_MISMATCH
def check(self, context: RuleContext) -> RuleResult:
"""检测银行实发与工资表实发是否一致"""
if not context.bank_data or context.employee_id not in context.bank_data:
return RuleResult(
is_exception=False,
description="无银行数据或员工不在银行记录中"
)
salary_record = context.salary_data.get(context.employee_id, {})
bank_record = context.bank_data.get(context.employee_id, {})
salary_net = salary_record.get("实发工资", 0) or salary_record.get("net_salary", 0)
bank_net = bank_record.get("实发金额", 0) or bank_record.get("bank_amount", 0)
diff = abs(salary_net - bank_net)
tolerance = 0.01
if diff > tolerance:
return RuleResult(
is_exception=True,
exception_type=self.exception_type,
severity=ExceptionSeverity.HIGH,
description=f"银行实发与工资表实发不一致(差异: {diff:.2f}元)",
detail={
"salary_net": salary_net,
"bank_net": bank_net,
},
salary_amount=salary_net,
bank_amount=bank_net,
difference_amount=diff,
suggested_action="联系银行确认交易状态"
)
return RuleResult(is_exception=False)
class BankMissingRule(ReconciliationRule):
"""银行记录缺失规则"""
def __init__(self, enabled: bool = False):
super().__init__(enabled=enabled)
@property
def name(self) -> str:
return "银行记录缺失检测"
@property
def exception_type(self) -> ExceptionType:
return ExceptionType.BANK_MISSING
def check(self, context: RuleContext) -> RuleResult:
"""检测有工资但无银行记录的员工"""
if not context.bank_data:
return RuleResult(
is_exception=False,
description="无银行数据"
)
has_salary = context.employee_id in context.salary_data
has_bank = context.employee_id in context.bank_data
if has_salary and not has_bank:
return RuleResult(
is_exception=True,
exception_type=self.exception_type,
severity=ExceptionSeverity.HIGH,
description="工资表有记录但银行无记录",
detail={"has_salary": True, "has_bank": False},
suggested_action="确认是否已发放或银行信息有误"
)
return RuleResult(is_exception=False)
# 规则工厂
class RuleFactory:
"""规则工厂"""
@staticmethod
def create_all_rules(enable_bank_rules: bool = False) -> List[ReconciliationRule]:
"""创建所有规则"""
rules = [
DuplicateEmployeeRule(), # 需要先检测重复
ZeroAmountRule(),
NegativeAmountRule(),
MissingEmployeeRule(),
AmountMismatchRule(),
UnusualAmountRule(),
]
if enable_bank_rules:
rules.extend([
BankMismatchRule(enabled=True),
BankMissingRule(enabled=True),
])
return rules
@staticmethod
def create_mvp_rules() -> List[ReconciliationRule]:
"""创建 MVP 阶段的规则(第一阶段 7 类异常检测)"""
return [
DuplicateEmployeeRule(),
ZeroAmountRule(),
NegativeAmountRule(),
MissingEmployeeRule(),
AmountMismatchRule(),
UnusualAmountRule(),
]
+1
View File
@@ -0,0 +1 @@
Generic single-database configuration.
+95
View File
@@ -0,0 +1,95 @@
import asyncio
from logging.config import fileConfig
from sqlalchemy import pool
from sqlalchemy.engine import Connection
from sqlalchemy.ext.asyncio import async_engine_from_config
from alembic import context
# 导入配置和模型
import sys
from pathlib import Path
sys.path.append(str(Path(__file__).parent.parent))
from app.core.config import get_settings
from app.core.database import Base
# this is the Alembic Config object, which provides
# access to the values within the .ini file in use.
config = context.config
# 从环境变量设置数据库 URL
settings = get_settings()
config.set_main_option("sqlalchemy.url", settings.database_url)
# Interpret the config file for Python logging.
# This line sets up loggers basically.
if config.config_file_name is not None:
fileConfig(config.config_file_name)
# add your model's MetaData object here
# for 'autogenerate' support
target_metadata = Base.metadata
# other values from the config, defined by the needs of env.py,
# can be acquired:
# my_important_option = config.get_main_option("my_important_option")
# ... etc.
def run_migrations_offline() -> None:
"""Run migrations in 'offline' mode.
This configures the context with just a URL
and not an Engine, though an Engine is acceptable
here as well. By skipping the Engine creation
we don't even need a DBAPI to be available.
Calls to context.execute() here emit the given string to the
script output.
"""
url = config.get_main_option("sqlalchemy.url")
context.configure(
url=url,
target_metadata=target_metadata,
literal_binds=True,
dialect_opts={"paramstyle": "named"},
)
with context.begin_transaction():
context.run_migrations()
def do_run_migrations(connection: Connection) -> None:
"""Execute migrations with the given connection."""
context.configure(connection=connection, target_metadata=target_metadata)
with context.begin_transaction():
context.run_migrations()
async def run_async_migrations() -> None:
"""Create async engine and run migrations."""
connectable = async_engine_from_config(
config.get_section(config.config_ini_section, {}),
prefix="sqlalchemy.",
poolclass=pool.NullPool,
)
async with connectable.connect() as connection:
await connection.run_sync(do_run_migrations)
await connectable.dispose()
def run_migrations_online() -> None:
"""Run migrations in 'online' mode (async)."""
asyncio.run(run_async_migrations())
if context.is_offline_mode():
run_migrations_offline()
else:
run_migrations_online()
+26
View File
@@ -0,0 +1,26 @@
"""${message}
Revision ID: ${up_revision}
Revises: ${down_revision | comma,n}
Create Date: ${create_date}
"""
from typing import Sequence, Union
from alembic import op
import sqlalchemy as sa
${imports if imports else ""}
# revision identifiers, used by Alembic.
revision: str = ${repr(up_revision)}
down_revision: Union[str, None] = ${repr(down_revision)}
branch_labels: Union[str, Sequence[str], None] = ${repr(branch_labels)}
depends_on: Union[str, Sequence[str], None] = ${repr(depends_on)}
def upgrade() -> None:
${upgrades if upgrades else "pass"}
def downgrade() -> None:
${downgrades if downgrades else "pass"}
@@ -0,0 +1,30 @@
"""添加文件上传模型
Revision ID: 251fa1f5d435
Revises: 2a96c11aa55c
Create Date: 2026-07-06 22:49:36.283202
"""
from typing import Sequence, Union
from alembic import op
import sqlalchemy as sa
# revision identifiers, used by Alembic.
revision: str = '251fa1f5d435'
down_revision: Union[str, None] = '2a96c11aa55c'
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
def upgrade() -> None:
# ### commands auto generated by Alembic - please adjust! ###
pass
# ### end Alembic commands ###
def downgrade() -> None:
# ### commands auto generated by Alembic - please adjust! ###
pass
# ### end Alembic commands ###
@@ -0,0 +1,30 @@
"""添加用户模型
Revision ID: 2a96c11aa55c
Revises: f6e668e45bff
Create Date: 2026-07-06 22:48:55.249510
"""
from typing import Sequence, Union
from alembic import op
import sqlalchemy as sa
# revision identifiers, used by Alembic.
revision: str = '2a96c11aa55c'
down_revision: Union[str, None] = 'f6e668e45bff'
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
def upgrade() -> None:
# ### commands auto generated by Alembic - please adjust! ###
pass
# ### end Alembic commands ###
def downgrade() -> None:
# ### commands auto generated by Alembic - please adjust! ###
pass
# ### end Alembic commands ###
@@ -0,0 +1,30 @@
"""添加字段映射、规则和任务模型
Revision ID: 52067074731d
Revises: 251fa1f5d435
Create Date: 2026-07-06 23:08:48.885971
"""
from typing import Sequence, Union
from alembic import op
import sqlalchemy as sa
# revision identifiers, used by Alembic.
revision: str = '52067074731d'
down_revision: Union[str, None] = '251fa1f5d435'
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
def upgrade() -> None:
# ### commands auto generated by Alembic - please adjust! ###
pass
# ### end Alembic commands ###
def downgrade() -> None:
# ### commands auto generated by Alembic - please adjust! ###
pass
# ### end Alembic commands ###
@@ -0,0 +1,30 @@
"""初始化数据库模型
Revision ID: f6e668e45bff
Revises:
Create Date: 2026-07-06 22:40:04.073613
"""
from typing import Sequence, Union
from alembic import op
import sqlalchemy as sa
# revision identifiers, used by Alembic.
revision: str = 'f6e668e45bff'
down_revision: Union[str, None] = None
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
def upgrade() -> None:
# ### commands auto generated by Alembic - please adjust! ###
pass
# ### end Alembic commands ###
def downgrade() -> None:
# ### commands auto generated by Alembic - please adjust! ###
pass
# ### end Alembic commands ###
+1
View File
@@ -13,6 +13,7 @@ pandas==2.2.0
openai==1.10.0
python-dotenv==1.0.0
structlog==24.1.0
python-dateutil==2.8.2
pytest==7.4.4
pytest-asyncio==0.23.4
httpx==0.26.0
+556
View File
@@ -0,0 +1,556 @@
"""
示例数据初始化脚本
创建一套完整的真实业务场景示例数据:
- 企业:深圳星辰科技有限公司
- 用户:admin 管理员
- 3个月工资对账数据(2026年4月、5月、6月)
- 包含各类业务异常场景
"""
import sys
import os
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
from datetime import datetime, date
from decimal import Decimal
import random
from app.core.database import async_session_maker as AsyncSessionLocal
from app.core.security import hash_password
# 导入所有模型以初始化 SQLAlchemy 关系
from app.models.company import Company, CompanyPlan, CompanyStatus
from app.models.user import User, UserRole, UserStatus
from app.models.uploaded_file import UploadedFile, FileType, ParseStatus
from app.models.reconciliation_task import ReconciliationTask, TaskStatus
from app.models.exception_item import ExceptionItem, ExceptionType, ExceptionStatus, ExceptionSeverity
from app.models.field_mapping import FieldMapping
from app.models.company_rule import CompanyRule
from app.models.standard_field import StandardField
# ========== 基础数据定义 ==========
# 公司信息
COMPANY_DATA = {
"name": "深圳星辰科技有限公司",
"plan": CompanyPlan.PROFESSIONAL.value,
"status": CompanyStatus.ACTIVE.value,
"data_retention_months": 12,
"max_users": 20,
"is_trial": False,
}
# 用户列表
USERS_DATA = [
{
"email": "admin@xingchen.com",
"password": "admin123",
"full_name": "张明华",
"role": UserRole.ADMIN.value,
},
{
"email": "finance@xingchen.com",
"password": "finance123",
"full_name": "李晓燕",
"role": UserRole.FINANCE_MANAGER.value,
},
{
"email": "accountant@xingchen.com",
"password": "account123",
"full_name": "王建国",
"role": UserRole.ACCOUNTANT.value,
},
]
# 员工数据(25名员工)
EMPLOYEES = [
# 管理层
{"id": "EMP001", "name": "张明华", "department": "管理层", "position": "总经理", "base_salary": 35000},
{"id": "EMP002", "name": "李晓燕", "department": "财务部", "position": "财务总监", "base_salary": 28000},
{"id": "EMP003", "name": "王建国", "department": "技术部", "position": "技术总监", "base_salary": 30000},
# 技术部
{"id": "EMP004", "name": "刘伟强", "department": "技术部", "position": "高级工程师", "base_salary": 22000},
{"id": "EMP005", "name": "陈志明", "department": "技术部", "position": "工程师", "base_salary": 18000},
{"id": "EMP006", "name": "周建军", "department": "技术部", "position": "工程师", "base_salary": 17000},
{"id": "EMP007", "name": "吴海涛", "department": "技术部", "position": "初级工程师", "base_salary": 12000},
{"id": "EMP008", "name": "郑晓峰", "department": "技术部", "position": "测试工程师", "base_salary": 15000},
# 产品部
{"id": "EMP009", "name": "孙丽娜", "department": "产品部", "position": "产品经理", "base_salary": 22000},
{"id": "EMP010", "name": "马俊杰", "department": "产品部", "position": "产品专员", "base_salary": 14000},
{"id": "EMP011", "name": "朱婷婷", "department": "产品部", "position": "UI设计师", "base_salary": 16000},
# 运营部
{"id": "EMP012", "name": "胡文静", "department": "运营部", "position": "运营总监", "base_salary": 25000},
{"id": "EMP013", "name": "林浩然", "department": "运营部", "position": "运营专员", "base_salary": 13000},
{"id": "EMP014", "name": "何雨晴", "department": "运营部", "position": "客服主管", "base_salary": 12000},
{"id": "EMP015", "name": "高建峰", "department": "运营部", "position": "客服专员", "base_salary": 9000},
# 市场部
{"id": "EMP016", "name": "罗晓东", "department": "市场部", "position": "市场总监", "base_salary": 26000},
{"id": "EMP017", "name": "宋志远", "department": "市场部", "position": "市场专员", "base_salary": 11000},
{"id": "EMP018", "name": "唐思远", "department": "市场部", "position": "销售经理", "base_salary": 15000},
{"id": "EMP019", "name": "韩冰冰", "department": "市场部", "position": "销售代表", "base_salary": 8000},
# 行政部
{"id": "EMP020", "name": "冯婉君", "department": "行政部", "position": "行政经理", "base_salary": 14000},
{"id": "EMP021", "name": "许志刚", "department": "行政部", "position": "行政专员", "base_salary": 8000},
{"id": "EMP022", "name": "邓小丽", "department": "行政部", "position": "前台接待", "base_salary": 6000},
# 人事部
{"id": "EMP023", "name": "曹丽华", "department": "人事部", "position": "人事经理", "base_salary": 15000},
{"id": "EMP024", "name": "彭海燕", "department": "人事部", "position": "招聘专员", "base_salary": 9000},
{"id": "EMP025", "name": "曾敏仪", "department": "人事部", "position": "薪酬专员", "base_salary": 8500},
]
def calculate_salary(employee: dict, month: int, year: int, has_bonus: bool = True) -> dict:
"""计算员工月工资"""
base = employee["base_salary"]
# 绩效工资(基本工资的10%-30%)
performance = base * random.uniform(0.1, 0.3)
# 奖金(部分员工有)
bonus = base * random.uniform(0.5, 1.5) if has_bonus and random.random() > 0.3 else 0
# 加班费(随机)
overtime = random.choice([0, 500, 1000, 1500, 2000])
# 应发工资
gross_salary = base + performance + bonus + overtime
# 社保个人部分(基本工资的10.5%)
social_security = base * 0.105
# 公积金(基本工资的12%
housing_fund = base * 0.12
# 个税专项附加扣除
tax_deduction = 2000 # 子女教育 + 继续教育 + 住房租金 等
# 应税工资
taxable = gross_salary - social_security - housing_fund - tax_deduction
# 个税(简化计算)
if taxable <= 3000:
tax = taxable * 0.03
elif taxable <= 12000:
tax = taxable * 0.10 - 210
elif taxable <= 25000:
tax = taxable * 0.20 - 1410
elif taxable <= 35000:
tax = taxable * 0.25 - 2660
elif taxable <= 55000:
tax = taxable * 0.30 - 4410
elif taxable <= 80000:
tax = taxable * 0.35 - 7160
else:
tax = taxable * 0.45 - 15160
tax = max(0, tax)
# 实发工资
net_salary = gross_salary - social_security - housing_fund - tax
return {
"employee_id": employee["id"],
"employee_name": employee["name"],
"department": employee["department"],
"base_salary": round(base, 2),
"performance_salary": round(performance, 2),
"bonus": round(bonus, 2),
"overtime_pay": round(overtime, 2),
"gross_salary": round(gross_salary, 2),
"social_security": round(social_security, 2),
"housing_fund": round(housing_fund, 2),
"personal_income_tax": round(tax, 2),
"net_salary": round(net_salary, 2),
}
def calculate_social_security(employee: dict) -> dict:
"""计算社保"""
base = employee["base_salary"]
# 养老保险(个人8%
pension = base * 0.08
# 医疗保险(个人2%
medical = base * 0.02
# 失业保险(个人0.5%
unemployment = base * 0.005
# 工伤保险(个人0%
work_injury = 0
# 生育保险(个人0%
maternity = 0
total = pension + medical + unemployment
return {
"employee_id": employee["id"],
"employee_name": employee["name"],
"department": employee["department"],
"pension": round(pension, 2),
"medical": round(medical, 2),
"unemployment": round(unemployment, 2),
"work_injury": round(work_injury, 2),
"maternity": round(maternity, 2),
"total": round(total, 2),
}
def calculate_tax(employee: dict, gross_salary: float) -> dict:
"""计算个税"""
base = employee["base_salary"]
# 社保个人部分
social = base * 0.105
# 公积金
housing = base * 0.12
# 专项附加扣除
deduction = 2000
# 应税所得
taxable = gross_salary - social - housing - deduction - 5000 # 起征点
if taxable <= 0:
tax = 0
tax_rate = 0
quick_deduction = 0
elif taxable <= 3000:
tax = taxable * 0.03
tax_rate = 0.03
quick_deduction = 0
elif taxable <= 12000:
tax = taxable * 0.10 - 210
tax_rate = 0.10
quick_deduction = 210
elif taxable <= 25000:
tax = taxable * 0.20 - 1410
tax_rate = 0.20
quick_deduction = 1410
elif taxable <= 35000:
tax = taxable * 0.25 - 2660
tax_rate = 0.25
quick_deduction = 2660
elif taxable <= 55000:
tax = taxable * 0.30 - 4410
tax_rate = 0.30
quick_deduction = 4410
elif taxable <= 80000:
tax = taxable * 0.35 - 7160
tax_rate = 0.35
quick_deduction = 7160
else:
tax = taxable * 0.45 - 15160
tax_rate = 0.45
quick_deduction = 15160
tax = max(0, tax)
return {
"employee_id": employee["id"],
"employee_name": employee["name"],
"department": employee["department"],
"gross_salary": round(gross_salary, 2),
"taxable_income": round(taxable, 2),
"tax_rate": tax_rate,
"quick_deduction": quick_deduction,
"tax_amount": round(tax, 2),
}
async def init_sample_data():
"""初始化示例数据"""
async with AsyncSessionLocal() as db:
try:
# 1. 创建公司
print("创建公司...")
company = Company(**COMPANY_DATA)
db.add(company)
await db.flush()
# 2. 创建用户
print("创建用户...")
admin_user = None
for user_data in USERS_DATA:
user = User(
company_id=company.id,
email=user_data["email"],
hashed_password=hash_password(user_data["password"]),
full_name=user_data["full_name"],
role=user_data["role"],
status=UserStatus.ACTIVE.value,
)
db.add(user)
if user_data["role"] == UserRole.ADMIN.value:
admin_user = user
await db.flush()
# 3. 创建3个月的对账任务
print("创建对账任务...")
tasks = []
months = [
(2026, 4, "4月"),
(2026, 5, "5月"),
(2026, 6, "6月"),
]
for year, month, month_name in months:
task = ReconciliationTask(
company_id=company.id,
name=f"{year}{month_name}工资对账",
period=f"{year}-{month:02d}",
status=TaskStatus.COMPLETED.value,
created_by=admin_user.id,
total_employees=len(EMPLOYEES),
)
db.add(task)
await db.flush()
tasks.append((task, year, month))
# 4. 为每个月创建异常
print("创建异常数据...")
# 10月异常
task_10, year_10, month_10 = tasks[0]
task_10.matched_count = 23
task_10.exception_count = 2
# 异常1:新员工入职(工资表有,社保表无)
exc1 = ExceptionItem(
company_id=company.id,
task_id=task_10.id,
exception_type=ExceptionType.MISSING_EMPLOYEE.value,
severity=ExceptionSeverity.MEDIUM.value,
status=ExceptionStatus.PENDING.value,
employee_id="EMP026",
employee_name="新入职员工",
description="员工【新入职员工】在工资表中存在,但在社保表中未找到记录",
detail={"source": "工资表", "missing_in": "社保表"},
ai_suggestion="请确认该员工是否已完成社保增员,如已增员请重新上传社保表",
)
db.add(exc1)
# 异常2:金额不匹配
salary_12 = calculate_salary(EMPLOYEES[11], month_10, year_10) # 胡文静
tax_12 = calculate_tax(EMPLOYEES[11], salary_12["gross_salary"])
exc2 = ExceptionItem(
company_id=company.id,
task_id=task_10.id,
exception_type=ExceptionType.AMOUNT_MISMATCH.value,
severity=ExceptionSeverity.HIGH.value,
status=ExceptionStatus.PENDING.value,
employee_id=EMPLOYEES[11]["id"],
employee_name=EMPLOYEES[11]["name"],
description=f"员工【{EMPLOYEES[11]['name']}】工资表实发金额与个税表申报金额不一致",
detail={
"salary_net": salary_12["net_salary"],
"tax_declared": tax_12["tax_amount"],
"difference": abs(salary_12["net_salary"] - tax_12["tax_amount"]),
},
salary_amount=salary_12["net_salary"],
tax_amount=tax_12["tax_amount"],
difference_amount=abs(salary_12["net_salary"] - tax_12["tax_amount"]),
ai_suggestion="请核对工资表和个税申报表的数据,确认是计算错误还是申报错误",
)
db.add(exc2)
# 11月异常
task_11, year_11, month_11 = tasks[1]
task_11.matched_count = 24
task_11.exception_count = 1
# 异常3:离职员工(工资表无,社保表有)
exc3 = ExceptionItem(
company_id=company.id,
task_id=task_11.id,
exception_type=ExceptionType.MISSING_EMPLOYEE.value,
severity=ExceptionSeverity.MEDIUM.value,
status=ExceptionStatus.CONFIRMED.value,
employee_id="EMP007",
employee_name="吴海涛",
description="员工【吴海涛】在社保表中存在,但在本月工资表中未找到",
detail={"source": "社保表", "missing_in": "工资表"},
handled_by=admin_user.id,
handled_at=datetime.utcnow(),
handling_note="该员工已于11月15日离职,工资表中已处理",
ai_suggestion="确认员工离职日期,如已办理社保减员可忽略此异常",
)
db.add(exc3)
# 12月异常(更多异常场景)
task_12, year_12, month_12 = tasks[2]
task_12.matched_count = 22
task_12.exception_count = 3
# 异常4:重复员工
exc4 = ExceptionItem(
company_id=company.id,
task_id=task_12.id,
exception_type=ExceptionType.DUPLICATE_EMPLOYEE.value,
severity=ExceptionSeverity.CRITICAL.value,
status=ExceptionStatus.PENDING.value,
employee_id="EMP019",
employee_name="韩冰冰",
description="员工【韩冰冰】在工资表中出现重复记录,请核实",
detail={"count": 2, "amounts": [8250.00, 8250.00]},
salary_amount=16500.00,
ai_suggestion="请检查是否是系统重复录入或是实际有两笔工资发放",
)
db.add(exc4)
# 异常5:金额为零
exc5 = ExceptionItem(
company_id=company.id,
task_id=task_12.id,
exception_type=ExceptionType.ZERO_AMOUNT.value,
severity=ExceptionSeverity.MEDIUM.value,
status=ExceptionStatus.PENDING.value,
employee_id="EMP022",
employee_name="邓小丽",
description="员工【邓小丽】本月实发工资为零,可能存在异常",
detail={"expected_amount": 6000.00, "actual_amount": 0},
salary_amount=0,
ai_suggestion="请确认该员工是否已离职、停薪留职或工资发放被扣押",
)
db.add(exc5)
# 异常6:金额异常(过高)
salary_3 = calculate_salary(EMPLOYEES[2], month_12, year_12)
exc6 = ExceptionItem(
company_id=company.id,
task_id=task_12.id,
exception_type=ExceptionType.UNUSUAL_AMOUNT.value,
severity=ExceptionSeverity.HIGH.value,
status=ExceptionStatus.PENDING.value,
employee_id=EMPLOYEES[2]["id"],
employee_name=EMPLOYEES[2]["name"],
description=f"员工【{EMPLOYEES[2]['name']}】本月工资较上月有较大波动",
detail={
"current": salary_3["net_salary"],
"previous_avg": 42000.00,
"change_ratio": 0.25,
},
salary_amount=salary_3["net_salary"],
difference_amount=8500.00,
ai_suggestion="该员工工资增长超过20%,请确认是否有大额奖金或提成发放",
)
db.add(exc6)
# 5. 为每个月创建上传文件
print("创建上传文件...")
files = []
for task, year, month in tasks:
# 工资表
salary_file = UploadedFile(
company_id=company.id,
task_id=task.id,
file_type=FileType.SALARY.value,
original_filename=f"{year}{month:02d}月工资表.xlsx",
stored_filename=f"salary_{task.id}_{year}{month:02d}.xlsx",
file_size=102400,
mime_type="application/vnd.openxmlformats-officedocument.spreadsheetml.sheet",
parse_status=ParseStatus.SUCCESS.value,
)
db.add(salary_file)
# 社保表
social_file = UploadedFile(
company_id=company.id,
task_id=task.id,
file_type=FileType.SOCIAL_SECURITY.value,
original_filename=f"{year}{month:02d}月社保表.xlsx",
stored_filename=f"social_{task.id}_{year}{month:02d}.xlsx",
file_size=92160,
mime_type="application/vnd.openxmlformats-officedocument.spreadsheetml.sheet",
parse_status=ParseStatus.SUCCESS.value,
)
db.add(social_file)
# 个税表
tax_file = UploadedFile(
company_id=company.id,
task_id=task.id,
file_type=FileType.TAX.value,
original_filename=f"{year}{month:02d}月个税表.xlsx",
stored_filename=f"tax_{task.id}_{year}{month:02d}.xlsx",
file_size=81920,
mime_type="application/vnd.openxmlformats-officedocument.spreadsheetml.sheet",
parse_status=ParseStatus.SUCCESS.value,
)
db.add(tax_file)
await db.flush()
files.append((salary_file.id, social_file.id, tax_file.id))
# 6. 创建字段映射规则(关联到第一个文件作为示例)
print("创建字段映射规则...")
first_file_id = files[0][0] # 使用10月工资表的文件ID
field_mappings = [
("员工编号", "employee_id", 0.98),
("工号", "employee_id", 0.95),
("姓名", "employee_name", 0.99),
("员工姓名", "employee_name", 0.99),
("部门", "department", 0.97),
("所属部门", "department", 0.95),
("岗位", "position", 0.90),
("职位", "position", 0.90),
("基本工资", "base_salary", 0.98),
("岗位工资", "base_salary", 0.92),
("应发工资", "gross_salary", 0.96),
("实发工资", "net_salary", 0.97),
("税前工资", "gross_salary", 0.93),
("养老保险", "pension", 0.95),
("医疗保险", "medical", 0.95),
("社保", "social_security", 0.90),
("公积金", "housing_fund", 0.95),
("个税", "personal_income_tax", 0.98),
("个人所得税", "personal_income_tax", 0.98),
]
for source, standard, confidence in field_mappings:
mapping = FieldMapping(
company_id=company.id,
file_id=first_file_id,
source_field=source,
standard_field=standard,
confidence=confidence,
confirmed=True,
confirmed_by=admin_user.id,
confirmed_at=datetime.utcnow(),
)
db.add(mapping)
await db.commit()
print("示例数据创建成功!")
print("\n" + "="*50)
print("登录信息:")
print("="*50)
for user_data in USERS_DATA:
print(f" 邮箱: {user_data['email']}")
print(f" 密码: {user_data['password']}")
print(f" 角色: {user_data['role']}")
print("-"*30)
print("\n包含 3 个月(4月、5月、6月)的对账数据")
print("共 25 名员工,多种异常场景")
except Exception as e:
await db.rollback()
print(f"创建示例数据失败: {e}")
raise
if __name__ == "__main__":
import asyncio
asyncio.run(init_sample_data())