# -*- coding: utf-8 -*- import re import os import importlib import inspect import traceback from lxml import etree from typing import Dict, List, Optional, Any from BiddingKG.dl.template_extract.abstract_template import AbstractTemplate from BiddingKG.dl.template_extract.out_line_extractor import OutlineExtractor from BiddingKG.dl.template_extract.table_extractor import get_html_table from BiddingKG.dl.common.Utils import log class TemplateManager: """ 模板管理与调用类(新增站源属性绑定+按站源筛选调用) 核心特性: 1. 自动注册时从模板文件的父文件夹名称读取“站源”属性 2. 调用模板时支持传入站源参数,仅执行对应站源的模板 3. 保留原有优先级规则、结果融合逻辑 """ def __init__(self): # 存储已注册的模板:key=模板ID,value=AbstractTemplate子类实例 self.registered_templates: Dict[str, AbstractTemplate] = {} # 缓存模板优先级映射(避免重复读取,提升筛选效率) self.template_priority_cache: Dict[str, int] = {} # 新增:缓存模板-站源映射(快速筛选站源对应的模板) self.template_site_map: Dict[str, List[str]] = {} # key=站源名,value=模板ID列表 def register_template(self, template: AbstractTemplate, site_source: str) -> None: """ 注册模板(新增站源参数,绑定模板与站源) :param template: AbstractTemplate子类实例 :param site_source: 站源名称(从模板文件父文件夹名称获取) """ if not isinstance(template, AbstractTemplate): raise TypeError("仅支持注册AbstractTemplate子类的模板实例") # 1. 处理模板覆盖逻辑 template_id = template.template_id if template_id in self.registered_templates: print(f"模板ID={template_id}已存在,将覆盖旧版本(原站源:{self._get_template_site(template_id)} → 新站源:{site_source})") # 移除旧模板在站源映射中的关联 old_site = self._get_template_site(template_id) if old_site and old_site in self.template_site_map: self.template_site_map[old_site].remove(template_id) if not self.template_site_map[old_site]: # 若站源下无模板,删除空列表 del self.template_site_map[old_site] # 2. 注册模板核心信息 self.registered_templates[template_id] = template self.template_priority_cache[template_id] = template.priority # 3. 绑定模板与站源(更新站源-模板ID映射) if site_source not in self.template_site_map: self.template_site_map[site_source] = [] if template_id not in self.template_site_map[site_source]: self.template_site_map[site_source].append(template_id) # 4. 扩展模板实例属性:动态添加站源属性(便于后续查看) setattr(template, "site_source", site_source) print(f"模板注册成功:ID={template_id},站源={site_source}") def unregister_template(self, template_id: str) -> bool: """ 注销模板(同步删除站源-模板映射) """ if template_id not in self.registered_templates: print(f"模板ID={template_id}不存在,注销失败") return False # 1. 移除站源-模板映射 site_source = self._get_template_site(template_id) if site_source and site_source in self.template_site_map: self.template_site_map[site_source].remove(template_id) if not self.template_site_map[site_source]: del self.template_site_map[site_source] # 2. 注销模板核心信息 del self.registered_templates[template_id] del self.template_priority_cache[template_id] print(f"模板注销成功:ID={template_id},站源={site_source}") return True def _get_template_site(self, template_id: str) -> Optional[str]: """ 辅助方法:通过模板ID获取对应的站源(从实例属性或站源映射反查) """ if template_id not in self.registered_templates: return None # 从模板实例的动态属性获取站源(优先) template = self.registered_templates[template_id] if hasattr(template, "site_source"): return template.site_source # 从站源映射反查(兜底) for site, tpl_ids in self.template_site_map.items(): if template_id in tpl_ids: return site return None def get_matched_templates(self, preprocessed_data: Dict[str, Any], site_source: Optional[str] = None) -> List[AbstractTemplate]: """ 筛选匹配模板(新增站源筛选逻辑) :param site_source: 目标站源(若传入,仅筛选该站源的模板;若为None,筛选所有站源模板) :return: 按“站源→优先级”筛选后的可用模板列表 """ # 1. 第一步:按站源筛选模板(核心新增逻辑) if site_source and site_source in self.template_site_map: # 仅保留目标站源的模板ID site_tpl_ids = self.template_site_map[site_source] candidate_templates = [ self.registered_templates[tpl_id] for tpl_id in site_tpl_ids ] else: # 无站源参数或站源不存在:候选模板为空 # candidate_templates = list(self.registered_templates.values()) candidate_templates = [] if site_source and site_source not in self.template_site_map: print(f"警告:站源{site_source}无对应模板") # 2. 第二步:按“调用时机”筛选(原有逻辑) matched_templates = [ tpl for tpl in candidate_templates if tpl.check_call_timing(preprocessed_data) ] # 3. 第三步:按优先级降序排序(原有逻辑) sorted_templates = sorted( matched_templates, key=lambda tpl: tpl.priority, reverse=True ) return sorted_templates def call_templates(self, preprocessed_data: Dict[str, Any], general_result: Dict[str, Any], site_source: Optional[str] = None, docid: Optional[str]='') -> Dict[str, Any]: """ 调用模板并融合结果(新增站源参数,仅调用对应站源模板) :param site_source: 目标站源(仅调用该站源的模板) """ # 1. 初始化结果:先纳入通用提取结果(基础值) final_result = general_result.copy() # 2. 获取“指定站源+符合调用时机”的模板,按优先级执行(核心修改) sorted_templates = self.get_matched_templates(preprocessed_data, site_source) if not sorted_templates: sorted_templates = self.get_matched_templates(preprocessed_data, site_source='shared_template') if not sorted_templates: print(f"站源{site_source}(或所有站源)无可用模板,直接使用通用结果+AI结果") else: for tpl in sorted_templates: tpl_result = tpl.extract(preprocessed_data) # 用模板结果补充/覆盖(非空值才更新) for key, value in tpl_result.items(): if value is not None and str(value).strip(): final_result[key] = value # print(f"调用模板:ID={tpl.template_id},站源={tpl.site_source},提取结果:{tpl_result}") log("调用模板:%s,站源:%s,docid:%s, 提取结果:%s"%(tpl.template_id, site_source, docid, str(tpl_result))) # 3. 融合AI提取结果(最高优先级,原有逻辑) ai_result = preprocessed_data.get("ai_extract_result", {}) for key, value in ai_result.items(): if value is not None and str(value).strip(): final_result[key] = value return final_result def batch_call(self, preprocessed_data_list: List[Dict[str, Any]], general_result_list: List[Dict[str, Any]], site_source: Optional[str] = None) -> List[Dict[str, Any]]: """ 批量调用模板(新增站源参数,批量筛选站源模板) :param site_source: 目标站源(所有批量数据共用一个站源) """ if len(preprocessed_data_list) != len(general_result_list): raise ValueError("预处理数据列表与通用结果列表长度不匹配") # 批量调用时传入统一的站源参数 batch_result = [ self.call_templates(pre_data, general_result, site_source) for pre_data, general_result in zip(preprocessed_data_list, general_result_list) ] return batch_result def get_template_info_list(self, site_source: Optional[str] = None) -> List[Dict[str, Any]]: """ 获取模板信息(支持按站源筛选) :return: 包含站源属性的模板信息列表 """ # 1. 按站源筛选模板 if site_source and site_source in self.template_site_map: target_templates = [ self.registered_templates[tpl_id] for tpl_id in self.template_site_map[site_source] ] else: target_templates = list(self.registered_templates.values()) # 2. 构造信息字典(新增站源字段) template_info_list = [] for tpl in target_templates: info = tpl.get_template_info() info["site_source"] = getattr(tpl, "site_source", "未知") # 新增站源信息 template_info_list.append(info) return template_info_list # ------------------------------ 核心修改:自动注册时绑定站源(父文件夹名称) ------------------------------ def auto_register_templates_from_dir(self, root_target_dir: str, exclude_files: Optional[List[str]] = None, exclude_dirs: Optional[List[str]] = None) -> None: """ 扫描文件夹并自动注册模板(核心修改:从父文件夹名称读取站源) 目录结构要求:root_target_dir/[站源文件夹]/[模板文件.py] 例:templates/zhihu/extract_tpl.py → 站源=zhihu;templates/toutiao/tpl_bidding.py → 站源=toutiao :param root_target_dir: 根目录(包含多个站源子文件夹) :param exclude_files: 需排除的文件名列表(默认排除__init__.py) :param exclude_dirs: 需排除的子文件夹列表(如["test_dir", "temp"]) """ # 1. 初始化排除列表 exclude_files = exclude_files or [] if "__init__.py" not in exclude_files: exclude_files.append("__init__.py") exclude_dirs = exclude_dirs or [] # 新增:排除不需要的子文件夹 # 2. 验证根目录有效性 root_dir_abs = os.path.abspath(root_target_dir) if not os.path.isdir(root_dir_abs): raise NotADirectoryError(f"根目录不是有效文件夹:{root_dir_abs}") print(f"\n开始扫描根目录:{root_dir_abs}") print(f"排除文件:{exclude_files},排除子文件夹:{exclude_dirs}") # 3. 遍历根目录下的所有子文件夹(每个子文件夹=一个站源) for dir_name in os.listdir(root_dir_abs): # 3.1 过滤非文件夹、排除文件夹 site_dir_path = os.path.join(root_dir_abs, dir_name) if not os.path.isdir(site_dir_path) or dir_name in exclude_dirs: continue site_source = dir_name # 核心:子文件夹名称 = 站源名称 print(f"\n=== 开始处理站源:{site_source}(文件夹:{dir_name})===") # 3.2 遍历当前站源文件夹下的所有.py文件 for filename in os.listdir(site_dir_path): # 过滤非.py文件、排除文件 if not filename.endswith(".py") or filename in exclude_files: continue file_path = os.path.join(site_dir_path, filename) module_name = f"{site_source}_{os.path.splitext(filename)[0]}" # 模块名:站源_文件名(避免重复) # 3.3 动态导入模块 try: spec = importlib.util.spec_from_file_location(module_name, file_path) if spec is None or spec.loader is None: print(f"警告:站源{site_source},无法创建模块规格,跳过文件:{filename}") continue module = importlib.util.module_from_spec(spec) spec.loader.exec_module(module) print(f"站源{site_source},成功加载模块:{module_name}(文件:{filename})") except Exception as e: print(f"警告:站源{site_source},加载模块{module_name}失败,跳过文件:{filename}") print(f"错误详情:{str(e)}") continue # 3.4 扫描模块中的模板并注册(绑定站源) self._scan_and_register_from_module(module, filename, site_source) # 4. 输出注册汇总 print(f"\n=== 文件夹扫描完成 ===") print(f"已识别站源数量:{len(self.template_site_map)}") for site, tpl_ids in self.template_site_map.items(): print(f" 站源{site}:{len(tpl_ids)}个模板(IDs:{tpl_ids})") print(f"总计注册模板数量:{len(self.registered_templates)}") def _scan_and_register_from_module(self, module: Any, filename: str, site_source: str) -> None: """ 辅助方法:从模块中识别模板并注册(新增站源参数) """ for member_name, member_obj in inspect.getmembers(module): # 筛选条件:AbstractTemplate子类 + 非基类本身 if (inspect.isclass(member_obj) and issubclass(member_obj, AbstractTemplate) and member_obj != AbstractTemplate): # 实例化并注册(传入站源参数) try: template_instance = member_obj() self.register_template(template_instance, site_source) # 调用修改后的注册方法 except Exception as e: print(f"警告:站源{site_source},实例化模板类{member_name}失败(文件:{filename})") print(f"错误详情:{str(e)}") continue # ------------------------------ 使用示例(按站源调用模板) ------------------------------ if __name__ == "__main__": # 1. 初始化模板管理器 manager = TemplateManager() # 2. 【关键】自动扫描根目录,按“子文件夹名称”绑定站源 # 目录结构示例: # ./templates_root/ # ├─ nanan/ (站源1:nanan) # │ ├─ bidding_place_tpl.py (模板:tpl_bidding_place_001) # │ └─ bidding_time_tpl.py (模板:tpl_bidding_time_001) # └─ xxs/ (站源2:xxs) # └─ bidding_place_tpl.py (模板:tpl_bidding_place_002) ROOT_TEMPLATE_DIR = "./templates" # 根目录(包含多个站源子文件夹) try: manager.auto_register_templates_from_dir( root_target_dir=ROOT_TEMPLATE_DIR, exclude_files=["test_tpl.py"], # 排除测试文件 exclude_dirs=["temp_dir"] # 排除临时文件夹 ) except Exception as e: traceback.print_exc() print(f"自动注册失败:{str(e)}") exit(1) with open('d:/html/2.html', encoding='utf-8') as f: html = f.read() tree = etree.HTML(html) # 3. 准备测试数据(2条数据,分别对应不同站源) test_pre_data = [ # 数据1:对应站源nanan { "text": """招标单位:南安市美林街道溪一村民委员会 代理单位:福建省协祥工程项目管理咨询有限公司 地区:南安市(福建省) 开标地点:其他场地 行业:市政 预算:323014.0元 项目状态:已公示""", "ai_extract_result": {"招标单位": "南安市美林街道溪一村民委员会"} }, # 数据2:对应站源xxs { "text": """招标单位:XX市住建局 代理单位:XX工程咨询公司 地区:XX市(浙江省) 开标地点:市公共资源交易中心3楼 行业:房建 预算:1500000.0元""", "ai_extract_result": {"地区": "XX市(浙江省)"} }, { "tree": tree, "html": html }, ] test_general_result = [ {"行业": "未知", "项目状态": "未知"}, {"代理单位": "未知", "预算": "未知"}, {} ] site_source = [ '', '15151', 'uni_jilin', '' ] # # 4. 批量调用模板并输出结果 # print("\n=== 所有注册模板信息 ===") # for idx, tpl_info in enumerate(manager.get_template_info_list(), 1): # print(f"{idx}. {tpl_info}") # # print("\n=== 批量提取结果 ===") # batch_result = manager.batch_call(test_pre_data, test_general_result, 'DX000002') # for idx, result in enumerate(batch_result, 1): # print(f"数据{idx}:{result}") import pandas as pd from collections import Counter import json from 分析表格内容 import 获取表格行列信息,获取行内文本,获取大纲文本信息,extract_table_headers # df = pd.read_csv('E:\实体识别数据/2023-08-24所有公告.csv') # df2 = pd.read_csv('E:\实体识别数据/2023-08-24所有公告html.csv') df = pd.read_csv('E:/模版提取/医院大学站源公告.csv') df2 = pd.read_csv('E:\模版提取/医院大学站源公告_html.csv') print('公告:', len(df), len(df2), len(set(df['docid'])&set(df2['docid']))) print(Counter(df['web_source_name']).most_common(100)) print(Counter(df['web_source_no']).most_common(100)) # df = df[df['web_source_name']=='哈尔滨工业大学'] # df = df[df['web_source_no'].str.contains('XX0182')] df = df.merge(df2, on='docid', how='inner') print('公告在:', set(df['docid'])&set(df2['docid'])) print('公告数量:', len(df), df['docid'].tolist()[:20],df.columns) # print(Counter(df['docchannel']).most_common()) # print(df['web_source_no'].tolist()[:20]) df.fillna('', inplace=True) # extractor = OutlineExtractor() # # df = pd.read_csv('E:/模版提取/中国南方电网-站源公告.csv') # df2 = pd.read_csv('E:/模版提取/中国南方电网-站源公告_html.csv') # df = df.merge(df2, on='docid', how='left') # # df = df[df['docid'].isin([673112494])] # print(Counter(df['docchannel']).most_common()) import time t1 = time.time() rs_list = [] n1 = n2 = 0 not_extract_heads = [] head_docid = {} for docid, html2,web_no in zip(df['docid'], df['dochtmlcon'], df['web_source_no']): tree = etree.HTML(html2) # outline = extractor.extract(html2) tables = get_html_table(html2) pre_data = { # "tree": tree, # "html": html2, # "outline": outline, "表格列表": tables } # with open('preprocessed_data.json', 'w', encoding='utf-8') as f: # text = json.dumps(pre_data, ensure_ascii=False, indent=2) # f.write(text) # print(text) if '-' in web_no: web_no = web_no.split('-')[0] # if web_no in ['Y00236', 'XX0182']: # continue result = manager.call_templates(pre_data, {}, web_no) # result = {} # for table in tables: # pre_data = { # "表格列表": [table] # } # if '-' in web_no: # web_no = web_no.split('-')[0] # result_tmp = manager.call_templates(pre_data, {}, web_no) # if result_tmp: # for k, v in result_tmp.items(): # if k not in result: # result[k] = v # else: # if isinstance(v, list): # result[k].extend(v) # else: # result[k] = v # print('公告数:%d, 耗时:%.4f'%(len(df), time.time()-t1)) # 公告数:714, 耗时:1.6382 1.5857 1.6712 # 公告数:714, 耗时:1.7684 1.7891 1.8029 # 耗时:1.9181 耗时:2.2976 if len(result) == 0: print('docid: ', docid) # 获取表格行列信息(html2) # 获取行内文本(html2) # 获取大纲文本信息(html2) headers = extract_table_headers(html2) if len(headers) > 2: print('未提取表头类型:', '_'.join(headers)) not_extract_heads.append('_'.join(headers)) head_docid['_'.join(headers)] = docid n1 += 1 else: print('提取结果:', docid, result) if len(result.get('招标信息', []))>1 and '标的' not in result.get('招标信息', [])[0] and '标包' not in result.get('招标信息', [])[0]: print('信息错误,缺少标包:', docid) if len(result.get('中标信息', []))>1 and '标的' not in result.get('中标信息', [])[0] and '标包' not in result.get('中标信息', [])[0]: print('信息错误,缺少标包:', docid) rs_list.append((docid, json.dumps(result, ensure_ascii=False, indent=2, sort_keys=True))) # print('docid: ', docid) n2 += 1 print() print('比例:', n1, n2, n2/(n1+n2)) # # # # df = pd.DataFrame(rs_list, columns=['docid', '结果']) # # df.to_excel('E:\待检查数据/中国南方电网-站源公告_模板提取结果.xlsx') # for it in Counter(not_extract_heads).most_common(50): # print(it, len(it[0].split('_')), head_docid[it[0]]) # print(Counter([it for l in not_extract_heads for it in l.split('_') ]).most_common(100))