""" 字段映射服务 管理字段映射的创建、确认和规则沉淀 """ 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