33b4c734aa
后端: - 新增认证(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 数据库迁移配置 - 添加初始化示例数据脚本 - 更新项目文档
314 lines
9.4 KiB
Python
314 lines
9.4 KiB
Python
"""
|
|
字段映射服务
|
|
|
|
管理字段映射的创建、确认和规则沉淀
|
|
"""
|
|
|
|
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 |