From fc98e93ee368dbe57337eaaa65511b08e53be55a Mon Sep 17 00:00:00 2001
From: pekopoke <1135796875@qq.com>
Date: Wed, 6 Aug 2025 10:00:18 +0800
Subject: [PATCH 1/3] add three extractors
---
.github/workflows/test.yml | 2 +-
examples/basic_usage.py | 384 ++-
examples/magic_html_extract_demo.py | 68 +
examples/resiliparse_extract_demo.py | 79 +
examples/trafilatura_extract_demo.py | 2873 +++++++++++++++++
requirements.txt | 8 +
tests/test_extractors.py | 91 +
tests/test_metrics.py | 176 +-
webmainbench/extractors/__init__.py | 6 +
.../extractors/magic_html_extractor.py | 80 +
.../extractors/resiliparse_extractor.py | 123 +
.../extractors/trafilatura_extractor.py | 121 +
12 files changed, 3803 insertions(+), 208 deletions(-)
create mode 100644 examples/magic_html_extract_demo.py
create mode 100644 examples/resiliparse_extract_demo.py
create mode 100644 examples/trafilatura_extract_demo.py
create mode 100644 tests/test_extractors.py
create mode 100644 webmainbench/extractors/magic_html_extractor.py
create mode 100644 webmainbench/extractors/resiliparse_extractor.py
create mode 100644 webmainbench/extractors/trafilatura_extractor.py
diff --git a/.github/workflows/test.yml b/.github/workflows/test.yml
index e259983..525c3e6 100644
--- a/.github/workflows/test.yml
+++ b/.github/workflows/test.yml
@@ -35,7 +35,7 @@ jobs:
run: |
python -m pip install --upgrade pip
pip install -e .
- pip install -r requirements.txt
+ pip install --ignore-requires-python -r requirements.txt
pip install pytest pytest-cov pytest-xdist coverage
- name: Install optional dependencies (ignore failures)
diff --git a/examples/basic_usage.py b/examples/basic_usage.py
index c8f44a8..2e9ac34 100755
--- a/examples/basic_usage.py
+++ b/examples/basic_usage.py
@@ -9,14 +9,14 @@
# 导入 WebMainBench 模块
from webmainbench import (
DataLoader, DataSaver, BenchmarkDataset, DataSample,
- ExtractorFactory, Evaluator,
+ ExtractorFactory, Evaluator,
format_results, setup_logging
)
def create_sample_dataset():
"""创建示例数据集"""
-
+
# 创建示例数据 - 包含多种内容类型(代码、公式、表格等)
samples = [
{
@@ -51,8 +51,7 @@ def greet(name):
"groundtruth_content_list": [
{"type": "heading", "content": "Python编程教程", "level": 1},
{"type": "paragraph", "content": "这是一个Python基础教程,展示如何定义函数。"},
- {"type": "code",
- "content": 'def greet(name):\n """问候函数"""\n return f"Hello, {name}!"\n\n# 使用示例\nresult = greet("World")\nprint(result)'},
+ {"type": "code", "content": 'def greet(name):\n """问候函数"""\n return f"Hello, {name}!"\n\n# 使用示例\nresult = greet("World")\nprint(result)'},
{"type": "paragraph", "content": "这个函数可以用来问候任何人。"}
],
"url": "https://python-tutorial.example.com/functions",
@@ -181,11 +180,9 @@ def greet(name):
{"type": "heading", "content": "数据分析报告", "level": 1},
{"type": "paragraph", "content": "以下是2024年第一季度的销售数据分析。"},
{"type": "heading", "content": "数据处理代码", "level": 2},
- {"type": "code",
- "content": "import pandas as pd\nimport numpy as np\n\n# 读取数据\ndf = pd.read_csv('sales_q1_2024.csv')\n\n# 计算统计信息\nmonthly_avg = df.groupby('month')['sales'].mean()\nprint(f\"平均销售额: {monthly_avg}\")"},
+ {"type": "code", "content": "import pandas as pd\nimport numpy as np\n\n# 读取数据\ndf = pd.read_csv('sales_q1_2024.csv')\n\n# 计算统计信息\nmonthly_avg = df.groupby('month')['sales'].mean()\nprint(f\"平均销售额: {monthly_avg}\")"},
{"type": "heading", "content": "销售统计", "level": 2},
- {"type": "table",
- "content": "| 月份 | 销售额(万元) | 增长率 |\n|------|-------------|--------|\n| 1月 | 120.5 | +15.2% |\n| 2月 | 135.8 | +12.7% |\n| 3月 | 148.3 | +9.2% |"},
+ {"type": "table", "content": "| 月份 | 销售额(万元) | 增长率 |\n|------|-------------|--------|\n| 1月 | 120.5 | +15.2% |\n| 2月 | 135.8 | +12.7% |\n| 3月 | 148.3 | +9.2% |"},
{"type": "paragraph", "content": "标准差公式:σ = √(Σ(xi - μ)² / n)"},
{"type": "paragraph", "content": "总体来看,第一季度销售表现良好,呈现稳定增长趋势。"}
],
@@ -211,12 +208,12 @@ def greet(name):
def quicksort(arr):
if len(arr) <= 1:
return arr
-
+
pivot = arr[len(arr) // 2]
left = [x for x in arr if x < pivot]
middle = [x for x in arr if x == pivot]
right = [x for x in arr if x > pivot]
-
+
return quicksort(left) + middle + quicksort(right)
复杂度对比
@@ -238,12 +235,12 @@ def greet(name):
def quicksort(arr):
if len(arr) <= 1:
return arr
-
+
pivot = arr[len(arr) // 2]
left = [x for x in arr if x < pivot]
middle = [x for x in arr if x == pivot]
right = [x for x in arr if x > pivot]
-
+
return quicksort(left) + middle + quicksort(right)
```
@@ -262,11 +259,9 @@ def quicksort(arr):
{"type": "heading", "content": "算法复杂度分析", "level": 1},
{"type": "paragraph", "content": "这里介绍常见算法的时间复杂度。"},
{"type": "heading", "content": "快速排序实现", "level": 2},
- {"type": "code",
- "content": "def quicksort(arr):\n if len(arr) <= 1:\n return arr\n \n pivot = arr[len(arr) // 2]\n left = [x for x in arr if x < pivot]\n middle = [x for x in arr if x == pivot]\n right = [x for x in arr if x > pivot]\n \n return quicksort(left) + middle + quicksort(right)"},
+ {"type": "code", "content": "def quicksort(arr):\n if len(arr) <= 1:\n return arr\n \n pivot = arr[len(arr) // 2]\n left = [x for x in arr if x < pivot]\n middle = [x for x in arr if x == pivot]\n right = [x for x in arr if x > pivot]\n \n return quicksort(left) + middle + quicksort(right)"},
{"type": "heading", "content": "复杂度对比", "level": 2},
- {"type": "table",
- "content": "| 算法 | 最好情况 | 平均情况 | 最坏情况 |\n|------|----------|----------|----------|\n| 快速排序 | O(n log n) | O(n log n) | O(n²) |\n| 归并排序 | O(n log n) | O(n log n) | O(n log n) |\n| 冒泡排序 | O(n) | O(n²) | O(n²) |"},
+ {"type": "table", "content": "| 算法 | 最好情况 | 平均情况 | 最坏情况 |\n|------|----------|----------|----------|\n| 快速排序 | O(n log n) | O(n log n) | O(n²) |\n| 归并排序 | O(n log n) | O(n log n) | O(n log n) |\n| 冒泡排序 | O(n) | O(n²) | O(n²) |"},
{"type": "equation-inline", "content": "T(n) = aT(n/b) + f(n)"},
{"type": "paragraph", "content": "其中 a ≥ 1, b > 1 是常数,f(n) 是正函数。"}
],
@@ -284,67 +279,67 @@ def quicksort(arr):
"content_type": "computer_science"
}
]
-
+
# 创建数据集
dataset = BenchmarkDataset(name="sample_dataset", description="示例评测数据集")
-
+
for sample_data in samples:
sample = DataSample.from_dict(sample_data)
dataset.add_sample(sample)
-
+
return dataset
def demo_basic_mock_evaluation():
"""演示基本评测流程"""
-
+
print("=== WebMainBench 基本使用示例 ===\n")
-
+
# 设置日志
setup_logging(level="INFO")
-
+
# 1. 创建或加载数据集
print("1. 创建示例数据集...")
dataset = create_sample_dataset()
print(f"数据集包含 {len(dataset)} 个样本")
print(f"数据集统计: {dataset.get_statistics()}\n")
-
+
# 2. 保存数据集到文件
data_dir = Path("data")
data_dir.mkdir(exist_ok=True)
-
+
dataset_path = data_dir / "sample_dataset.jsonl"
DataSaver.save_jsonl(dataset, dataset_path, include_results=False)
print(f"数据集已保存到: {dataset_path}\n")
-
+
# 3. 重新加载数据集
print("2. 重新加载数据集...")
loaded_dataset = DataLoader.load_jsonl(dataset_path)
print(f"加载的数据集包含 {len(loaded_dataset)} 个样本\n")
-
+
# 4. 列出可用的抽取器
print("3. 可用的抽取器:")
available_extractors = ExtractorFactory.list_available()
for extractor_name in available_extractors:
print(f" - {extractor_name}")
print()
-
+
# 5. 创建评测器
print("4. 创建评测器...")
evaluator = Evaluator()
print(f"可用的评测指标: {evaluator.metric_calculator.list_available_metrics()}\n")
-
+
# 6. 创建一个模拟抽取器进行演示
print("5. 创建模拟抽取器...")
-
+
from webmainbench.extractors import BaseExtractor, ExtractionResult
-
+
class MockExtractor(BaseExtractor):
"""模拟抽取器,用于演示"""
-
+
def _setup(self):
pass
-
+
def _extract_content(self, html, url=None):
# 简单的模拟抽取逻辑
if "标题" in html:
@@ -356,19 +351,19 @@ def _extract_content(self, html, url=None):
else:
content = "提取的内容"
content_list = [{"type": "paragraph", "content": "提取的内容"}]
-
+
return ExtractionResult(
content=content,
content_list=content_list,
success=True,
confidence_score=0.85
)
-
+
# 注册模拟抽取器
ExtractorFactory.register("mock", MockExtractor)
mock_extractor = ExtractorFactory.create("mock")
print("模拟抽取器已创建\n")
-
+
# 7. 运行评测
print("6. 运行评测...")
result = evaluator.evaluate(
@@ -376,21 +371,21 @@ def _extract_content(self, html, url=None):
extractor=mock_extractor,
max_samples=2 # 限制样本数量用于演示
)
-
+
# 8. 显示结果
print("\n7. 评测结果:")
print("=" * 50)
formatted_results = format_results(result.to_dict())
print(formatted_results)
-
+
# 9. 保存结果
results_dir = Path("results")
results_dir.mkdir(exist_ok=True)
-
+
results_path = results_dir / "mock_evaluation_results.json"
DataSaver.save_evaluation_results(result, results_path)
print(f"\n结果已保存到: {results_path}")
-
+
# 10. 生成报告
report_path = results_dir / "mock_evaluation_report.csv"
DataSaver.save_summary_report(result, report_path)
@@ -399,19 +394,18 @@ def _extract_content(self, html, url=None):
def demo_extractor_comparison():
"""演示多抽取器对比"""
-
+
print("\n=== 多抽取器对比演示 ===\n")
-
+
# 创建数据集
dataset = create_sample_dataset()
-
+
# 创建多个模拟抽取器
from webmainbench.extractors import BaseExtractor, ExtractionResult
-
+
class ExtractorA(BaseExtractor):
def _setup(self):
pass
-
def _extract_content(self, html, url=None):
return ExtractionResult(
content="抽取器A的结果",
@@ -419,11 +413,10 @@ def _extract_content(self, html, url=None):
success=True,
confidence_score=0.9
)
-
+
class ExtractorB(BaseExtractor):
def _setup(self):
pass
-
def _extract_content(self, html, url=None):
return ExtractionResult(
content="抽取器B的结果",
@@ -431,33 +424,33 @@ def _extract_content(self, html, url=None):
success=True,
confidence_score=0.8
)
-
+
# 注册抽取器
ExtractorFactory.register("extractor_a", ExtractorA)
ExtractorFactory.register("extractor_b", ExtractorB)
-
+
# 运行对比
evaluator = Evaluator()
extractors = ["extractor_a", "extractor_b"]
-
+
results = evaluator.compare_extractors(
dataset=dataset,
extractors=extractors,
max_samples=2
)
-
+
# 显示对比结果
print("对比结果:")
print("-" * 40)
for extractor_name, result in results.items():
overall_score = result.overall_metrics.get('overall', 0)
print(f"{extractor_name}: {overall_score:.4f}")
-
+
# 保存多抽取器对比榜单
all_results = []
for extractor_name, result in results.items():
all_results.append(result.to_dict())
-
+
results_dir = Path("results")
results_dir.mkdir(exist_ok=True)
leaderboard_path = results_dir / "leaderboard.csv"
@@ -467,17 +460,17 @@ def _extract_content(self, html, url=None):
def demo_llm_webkit_evaluation():
"""演示LLM-WebKit抽取器的6项指标评测"""
-
+
print("=== LLM-WebKit Extractor 6项指标评测示例 ===\n")
-
+
# 设置日志
setup_logging(level="INFO")
-
+
# 1. 创建包含各种内容类型的测试数据集
print("1. 创建包含多种内容类型的测试数据集...")
-
+
samples = []
-
+
# 样本1: 包含文本和代码
samples.append(DataSample(
id="text_code_sample",
@@ -509,12 +502,11 @@ def hello_world():
groundtruth_content_list=[
{"type": "heading", "content": "Python编程示例", "level": 1},
{"type": "text", "content": "这是一段关于Python编程的介绍文本。"},
- {"type": "code", "content": "def hello_world():\n print(\"Hello, World!\")\n return True",
- "language": "python"},
+ {"type": "code", "content": "def hello_world():\n print(\"Hello, World!\")\n return True", "language": "python"},
{"type": "text", "content": "以上代码展示了一个简单的Python函数。"}
]
))
-
+
# 样本2: 包含表格
samples.append(DataSample(
id="table_sample",
@@ -554,11 +546,10 @@ def hello_world():
| 产品B | 200 | 3000 |""",
groundtruth_content_list=[
{"type": "heading", "content": "销售数据统计", "level": 2},
- {"type": "table",
- "content": "| 产品 | 销量 | 收入 |\n|------|------|------|\n| 产品A | 100 | 1000 |\n| 产品B | 200 | 3000 |"}
+ {"type": "table", "content": "| 产品 | 销量 | 收入 |\n|------|------|------|\n| 产品A | 100 | 1000 |\n| 产品B | 200 | 3000 |"}
]
))
-
+
# 样本3: 包含公式
samples.append(DataSample(
id="formula_sample",
@@ -586,74 +577,74 @@ def hello_world():
{"type": "formula", "content": "\\int_{-\\infty}^{\\infty} e^{-x^2} dx = \\sqrt{\\pi}"}
]
))
-
+
# 创建数据集并添加样本
dataset = BenchmarkDataset(name="llm_webkit_test", description="LLM-WebKit 6项指标测试数据集")
for sample in samples:
dataset.add_sample(sample)
-
+
print(f"测试数据集包含 {len(dataset)} 个样本")
print(f"样本类型: 文本+代码, 表格, 公式\n")
-
+
# 2. 创建LLM-WebKit抽取器
print("2. 创建LLM-WebKit抽取器...")
-
+
# 显示所有可用的抽取器
available_extractors = ExtractorFactory.list_available()
print(f"可用的抽取器: {available_extractors}")
-
+
# 直接创建LLM-WebKit抽取器,设置模型路径
config = {
"model_path": "/Users/chupei/model/checkpoint-3296"
}
extractor = ExtractorFactory.create("llm-webkit", config=config)
print(f"✅ LLM-WebKit抽取器创建成功,模型路径: {config['model_path']}")
-
+
print()
-
+
# 3. 创建评测器并显示所有可用指标
print("3. 创建评测器...")
evaluator = Evaluator()
available_metrics = evaluator.metric_calculator.list_available_metrics()
print(f"✅ 可用的评测指标 ({len(available_metrics)}项):")
-
+
# 按照6项指标分类显示
target_metrics = ["overall", "text_edit", "code_edit", "table_edit", "table_TEDS", "formula_edit"]
-
+
for metric in target_metrics:
if metric in available_metrics:
print(f" ✅ {metric}")
else:
print(f" ❌ {metric} (未注册)")
-
+
print()
-
+
# 4. 运行评测
print("4. 开始评测...")
print("=" * 60)
-
+
result = evaluator.evaluate(
dataset=dataset,
extractor=extractor,
max_samples=None # 评测所有样本
)
-
+
# 5. 显示详细的6项指标结果
print("\n5. 📊 6项指标详细评测结果:")
print("=" * 60)
-
+
results_dict = result.to_dict()
-
+
# 从overall_metrics中提取指标结果
metrics = results_dict.get('overall_metrics', {})
-
+
# 按照指标分类显示
print(f"\n🏆 综合指标:")
if 'overall' in metrics:
print(f" overall (综合得分): {metrics['overall']:.4f}")
else:
print(" overall: 未计算")
-
+
print(f"\n📝 文本相关指标:")
if 'text_edit' in metrics:
print(f" text_edit (文本编辑距离): {metrics['text_edit']:.4f}")
@@ -663,7 +654,7 @@ def hello_world():
print(f" code_edit (代码编辑距离): {metrics['code_edit']:.4f}")
else:
print(" code_edit: 未计算")
-
+
print(f"\n📊 表格相关指标:")
if 'table_edit' in metrics:
print(f" table_edit (表格编辑距离): {metrics['table_edit']:.4f}")
@@ -673,37 +664,37 @@ def hello_world():
print(f" table_TEDS (表格结构相似度): {metrics['table_TEDS']:.4f}")
else:
print(" table_TEDS: 未计算")
-
+
print(f"\n🧮 公式相关指标:")
if 'formula_edit' in metrics:
print(f" formula_edit (公式编辑距离): {metrics['formula_edit']:.4f}")
else:
print(" formula_edit: 未计算")
-
+
print(f"\n📈 详细统计:")
print(f" 总样本数: {len(dataset)}")
success_count = len([s for s in results_dict.get('sample_results', []) if s.get('extraction_success', False)])
failure_count = len(dataset) - success_count
print(f" 成功样本数: {success_count}")
print(f" 失败样本数: {failure_count}")
-
+
# 6. 保存结果到文件
print("\n" + "=" * 60)
print("6. 保存评测结果...")
-
+
results_dir = Path("results")
results_dir.mkdir(exist_ok=True)
-
+
# 保存详细结果
results_path = results_dir / "llm_webkit_evaluation_results.json"
DataSaver.save_evaluation_results(result, results_path) # 直接传递result对象
print(f"✅ 详细结果已保存到: {results_path}")
-
+
# 生成CSV报告
report_path = results_dir / "llm_webkit_evaluation_report.csv"
DataSaver.save_summary_report(result, report_path) # 直接传递result对象
print(f"✅ CSV报告已保存到: {report_path}")
-
+
print("\n" + "=" * 60)
print("✅ LLM-WebKit 6项指标评测完成!")
@@ -711,84 +702,84 @@ def hello_world():
def demo_dataset_with_extraction():
"""演示保存带有抽取内容的数据集"""
print("=== 演示:保存带有抽取内容的数据集 ===")
-
+
from webmainbench import DataLoader, DataSaver, Evaluator, ExtractorFactory
from pathlib import Path
-
+
# 配置文件路径
data_dir = Path("data")
dataset_path = data_dir / "sample_dataset.jsonl"
# dataset_path = "/Users/chupei/Downloads/WebMainBench_dataset_merge_2549.jsonl"
-
+
print(f"📂 数据集文件: {dataset_path}")
-
+
# 🔧 创建llm-webkit抽取器(统一使用)
extractor_config = {"model_path": "/Users/chupei/model/checkpoint-3296"}
extractor = ExtractorFactory.create("llm-webkit", config=extractor_config)
print(f"🤖 使用抽取器: {extractor.name}")
-
+
# 创建评测器
evaluator = Evaluator()
-
+
# 🔧 选择评测模式:内存模式 vs 批处理模式
USE_BATCHED_MODE = True # 设置为True使用批处理模式(适用于大数据集)
-
+
if USE_BATCHED_MODE:
print("🔄 使用批处理模式(内存优化)")
-
+
# 🚀 批处理评测(适用于大数据集)
result = evaluator.evaluate_batched(
jsonl_file_path=dataset_path,
extractor=extractor, # 直接传递extractor对象
- batch_size=10, # 小批次
- max_samples=20 # 演示用
+ batch_size=10, # 小批次
+ max_samples=20 # 演示用
)
print(f"✅ 批处理评测完成,总体得分: {result.overall_metrics.get('overall', 0):.4f}")
-
+
# 为了保存带有抽取内容的数据集,需要重新加载原始数据集
# 注:这里只是短暂加载用于保存,不影响前面的内存优化评测
dataset = DataLoader.load_jsonl(dataset_path, include_results=False)
dataset.name = result.dataset_name
-
+
else:
print("🔄 使用传统内存模式")
-
+
# 从文件加载数据集
print(f"📂 从文件加载数据集: {dataset_path}")
dataset = DataLoader.load_jsonl(dataset_path, include_results=False)
dataset.name = "WebMainBench_with_extraction"
dataset.description = "演示抽取内容保存的测试数据集"
-
+
print(f"📊 加载数据集完成,包含 {len(dataset.samples)} 个样本")
-
+
# 运行评测
result = evaluator.evaluate(dataset, extractor)
-
+
print(f"✅ 评测完成,总体得分: {result.overall_metrics.get('overall', 0):.4f}")
-
+
# 保存带有抽取内容的数据集
results_dir = Path("results")
enriched_dataset_path = results_dir / f"{dataset.name}_with_{extractor.name}_extraction.jsonl"
-
+
DataSaver.save_dataset_with_extraction(
results=result,
- dataset=dataset,
+ dataset=dataset,
file_path=enriched_dataset_path,
extractor_name=extractor.name
)
-
+
print(f"💾 已保存带有抽取内容的数据集到: {enriched_dataset_path}")
-
+
# 保存评测结果和摘要报告
evaluation_results_path = results_dir / f"{dataset.name}_{extractor.name}_evaluation_results.json"
summary_report_path = results_dir / f"{dataset.name}_{extractor.name}_evaluation_report.csv"
-
+
DataSaver.save_evaluation_results(result, evaluation_results_path)
DataSaver.save_summary_report(result, summary_report_path)
-
+
print(f"📊 已保存评测结果到: {evaluation_results_path}")
print(f"📈 已保存摘要报告到: {summary_report_path}")
-
+
# 显示保存的字段信息
print("\n📋 保存的新字段包括:")
print(f" - {extractor.name}_content: 抽取的内容")
@@ -798,16 +789,171 @@ def demo_dataset_with_extraction():
print(f" - {extractor.name}_*_score: 各项指标分数")
+def demo_multi_extraction():
+ """演示保存带有多个抽取器抽取内容的数据集(支持批处理模式)"""
+ print("=== 演示:保存带有多个抽取器抽取内容的数据集 ===")
+
+ from webmainbench import DataLoader, DataSaver, Evaluator, ExtractorFactory
+ from pathlib import Path
+ import time
+ # 设置日志
+ setup_logging(level="INFO")
+
+ # 配置文件路径
+ data_dir = Path("../data")
+ # dataset_path = data_dir / "sample_dataset.jsonl"
+ dataset_path = "/home/lulindong/Pycharm_projects/cc/test.jsonl"
+
+ print(f"📂 数据集文件: {dataset_path}")
+
+ # 🔧 定义要使用的抽取器列表及配置
+ extractors_info = [
+ {"name": "resiliparse", "config": {
+ "main_content": True,
+ "alt_texts": True,
+ "links": False,
+ "list_bullets": True,
+ "preserve_formatting": True
+ }},
+
+ # {"name": "trafilatura", "config": {}},
+ # {"name": "magic-html", "config": {}},
+ ]
+
+ # 🔧 选择评测模式:内存模式 vs 批处理模式
+ USE_BATCHED_MODE = True # 大数据集建议设为True
+ BATCH_SIZE = 10 # 批处理大小
+ MAX_SAMPLES = None # 演示用(全量评测可设为None)
+
+ # 创建结果目录
+ results_dir = Path("results")
+ results_dir.mkdir(exist_ok=True)
+
+ # 存储所有抽取器的评测结果和性能数据
+ all_results = []
+ extractor_performance = []
+
+ # 为每个抽取器运行评测
+ for info in extractors_info:
+ extractor_name = info["name"]
+ config = info["config"]
+
+ try:
+ # 创建抽取器实例
+ extractor = ExtractorFactory.create(extractor_name, config=config)
+ print(f"\n🤖 使用抽取器: {extractor.name}")
+ except Exception as e:
+ print(f"⚠️ {extractor_name} 抽取器创建失败: {e}")
+ continue
+
+ # 记录总耗时
+ start_time = time.time()
+
+ # 初始化评测器
+ evaluator = Evaluator()
+
+ # 选择批处理模式或传统模式
+ if USE_BATCHED_MODE:
+ print(f"🔄 使用批处理模式(批大小: {BATCH_SIZE},最大样本: {MAX_SAMPLES or '全部'})")
+ # 批处理评测(内存优化)
+ result = evaluator.evaluate_batched(
+ jsonl_file_path=dataset_path,
+ extractor=extractor,
+ batch_size=BATCH_SIZE,
+ max_samples=MAX_SAMPLES
+ )
+ # 为保存数据集,临时加载原始数据(不影响内存优化)
+ dataset = DataLoader.load_jsonl(dataset_path, include_results=False, max_samples=MAX_SAMPLES)
+ dataset.name = result.dataset_name
+ else:
+ print("🔄 使用传统内存模式")
+ # 加载完整数据集到内存
+ dataset = DataLoader.load_jsonl(dataset_path, include_results=False, max_samples=MAX_SAMPLES)
+ dataset.name = "WebMainBench_with_multi_extraction"
+ dataset.description = "多抽取器内容保存演示数据集"
+ print(f"📊 加载数据集完成,包含 {len(dataset.samples)} 个样本")
+
+ # 传统模式评测
+ result = evaluator.evaluate(dataset, extractor)
+
+ # 计算耗时指标
+ total_time = time.time() - start_time
+ total_samples = len(dataset.samples)
+ avg_time_per_sample = total_time / total_samples if total_samples else 0
+
+ # 保存性能数据
+ extractor_performance.append({
+ "name": extractor_name,
+ "total_samples": total_samples,
+ "total_time": total_time,
+ "avg_time_per_sample": avg_time_per_sample
+ })
+
+ # 输出评测结果
+ print(f"⏱️ 总耗时: {total_time:.4f}秒(单样本平均: {avg_time_per_sample:.4f}秒)")
+ print(f"📊 核心指标:")
+ print(f" code_edit: {result.overall_metrics.get('code_edit', 0):.4f}")
+ print(f" formula_edit: {result.overall_metrics.get('formula_edit', 0):.4f}")
+ print(f" table_TEDS: {result.overall_metrics.get('table_TEDS', 0):.4f}")
+ print(f" table_edit: {result.overall_metrics.get('table_edit', 0):.4f}")
+ print(f" text_edit: {result.overall_metrics.get('text_edit', 0):.4f}")
+ print(f"✅ 总体得分: {result.overall_metrics.get('overall', 0):.4f}")
+
+ all_results.append(result)
+
+ # 保存带有当前抽取器内容的数据集
+ enriched_dataset_path = results_dir / f"{dataset.name}_with_{extractor.name}_extraction.jsonl"
+ DataSaver.save_dataset_with_extraction(
+ results=result,
+ dataset=dataset,
+ file_path=enriched_dataset_path,
+ extractor_name=extractor.name
+ )
+ print(f"💾 已保存抽取内容到: {enriched_dataset_path}")
+
+ # 保存单个抽取器的评测结果
+ eval_results_path = results_dir / f"{dataset.name}_{extractor.name}_evaluation_results.json"
+ DataSaver.save_evaluation_results(result, eval_results_path)
+ print(f"📋 已保存评测结果到: {eval_results_path}")
+
+ # 保存所有抽取器的汇总报告
+ if all_results:
+ summary_path = results_dir / f"{dataset.name}_multi_extractors_summary_report.csv"
+ DataSaver.save_summary_report(all_results, summary_path)
+ print(f"\n📈 已保存汇总报告到: {summary_path}")
+
+ # 展示性能对比
+ if extractor_performance:
+ print("\n⚡ 抽取器性能对比:")
+ for perf in extractor_performance:
+ print(f" {perf['name']}:")
+ print(f" 样本数: {perf['total_samples']}")
+ print(f" 总耗时: {perf['total_time']:.4f}秒")
+ print(f" 单样本耗时: {perf['avg_time_per_sample']:.4f}秒")
+ print(f" 效率: {1 / perf['avg_time_per_sample']:.2f}样本/秒")
+
+ # 展示保存的字段信息
+ print("\n📋 保存的新字段说明:")
+ for info in extractors_info:
+ name = info["name"]
+ print(f" {name}相关字段:")
+ print(f" - {name}_content: 抽取的原始内容")
+ print(f" - {name}_content_list: 结构化内容列表(含type字段)")
+ print(f" - {name}_success: 抽取是否成功(布尔值)")
+ print(f" - {name}_time: 单样本抽取耗时(秒)")
+ print(f" - {name}_*_score: 各指标得分(如{name}_text_edit)")
+
if __name__ == "__main__":
try:
- demo_basic_mock_evaluation()
- demo_llm_webkit_evaluation() # 使用LLM-WebKit评测示例
- demo_extractor_comparison()
- demo_dataset_with_extraction() # 演示保存带有抽取内容的数据集
+ # demo_basic_mock_evaluation()
+ # demo_llm_webkit_evaluation() # 使用LLM-WebKit评测示例
+ # demo_extractor_comparison()
+ # demo_dataset_with_extraction() # 演示保存带有抽取内容的数据集
+ demo_multi_extraction() # 演示多个抽取器同时评测
+ # demo_lld_workers_extraction()
print("\n✅ 示例运行完成!")
-
+
except Exception as e:
print(f"\n❌ 运行出错: {e}")
import traceback
-
- traceback.print_exc()
\ No newline at end of file
+ traceback.print_exc()
\ No newline at end of file
diff --git a/examples/magic_html_extract_demo.py b/examples/magic_html_extract_demo.py
new file mode 100644
index 0000000..726c054
--- /dev/null
+++ b/examples/magic_html_extract_demo.py
@@ -0,0 +1,68 @@
+import time
+from webmainbench.extractors import ExtractorFactory
+
+# 配置 MagicHTML 抽取器(这里可根据需要添加更多配置)
+config = {}
+try:
+ # 创建 MagicHTML 抽取器实例
+ extractor = ExtractorFactory.create("magic-html", config=config)
+ print(f"✅ Extractor创建成功: {extractor.description}")
+ print(f"📋 版本: {extractor.version}")
+ print(f"⚙️ 配置: {extractor.get_config()}\n")
+except Exception as e:
+ print(f"❌ Extractor创建失败: {e}")
+
+# 测试 HTML
+test_html = """
+
+
+ Python编程教程
+ 这是一个Python基础教程,展示如何定义函数。
+ def greet(name):
+ ""问候函数""
+ return f"Hello, {name}!"
+
+# 使用示例
+result = greet("World")
+print(result)
+ 这个函数可以用来问候任何人。
+
+
+"""
+
+print("🔍 开始内容提取...")
+start_time = time.time()
+
+try:
+ result = extractor.extract(test_html)
+ end_time = time.time()
+
+ print(f"⏱️ 提取耗时: {end_time - start_time:.2f}秒\n")
+
+ # 显示提取结果
+ if result.success:
+ print("✅ 内容提取成功!\n")
+
+ print("📄 提取的主要内容:")
+ print("=" * 50)
+ print(result.content[:500] + "..." if len(result.content) > 500 else result.content)
+ print("=" * 50)
+
+ print(f"\n📊 提取统计:")
+ print(f" • 内容长度: {len(result.content)} 字符")
+ print(f" • 标题: {result.title}")
+ print(f" • 语言: {result.language}")
+ print(f" • 提取时间: {result.extraction_time:.3f}秒")
+
+ if result.content_list:
+ print(f" • 结构化内容块: {len(result.content_list)}个")
+ for i, item in enumerate(result.content_list[:3]): # 显示前3个
+ print(f" [{i + 1}] {item.get('type', 'unknown')}: {item.get('content', '')[:50]}...")
+ else:
+ print("❌ 内容提取失败")
+ print(f"错误信息: {result.error_message}")
+ if result.error_traceback:
+ print(f"错误详情:\n{result.error_traceback}")
+
+except Exception as e:
+ print(f"❌ 提取过程中发生异常: {e}")
\ No newline at end of file
diff --git a/examples/resiliparse_extract_demo.py b/examples/resiliparse_extract_demo.py
new file mode 100644
index 0000000..ba33a14
--- /dev/null
+++ b/examples/resiliparse_extract_demo.py
@@ -0,0 +1,79 @@
+import time
+from webmainbench.extractors import ExtractorFactory
+
+# 配置 Resiliparse 抽取器
+config = {
+ "main_content": True,
+ "alt_texts": True,
+ "links": False,
+ "form_fields": False,
+ "noscript": False,
+ "list_bullets": True,
+ "preserve_formatting": True,
+ "comments": True
+}
+
+try:
+ # 创建 Resiliparse 抽取器实例
+ extractor = ExtractorFactory.create("resiliparse", config=config)
+ print(f"✅ Extractor创建成功: {extractor.description}")
+ print(f"📋 版本: {extractor.version}")
+ print(f"⚙️ 配置: {extractor.get_config()}\n")
+except Exception as e:
+ print(f"❌ Extractor创建失败: {e}")
+
+
+# 测试 HTML
+test_html = """
+
+
+ Python编程教程
+ 这是一个Python基础教程,展示如何定义函数。
+ def greet(name):
+ ""问候函数""
+ return f"Hello, {name}!"
+
+# 使用示例
+result = greet("World")
+print(result)
+ 这个函数可以用来问候任何人。
+
+
+"""
+
+print("🔍 开始内容提取...")
+start_time = time.time()
+
+try:
+ result = extractor.extract(test_html)
+ end_time = time.time()
+
+ print(f"⏱️ 提取耗时: {end_time - start_time:.2f}秒\n")
+
+ # 显示提取结果
+ if result.success:
+ print("✅ 内容提取成功!\n")
+
+ print("📄 提取的主要内容:")
+ print("=" * 50)
+ print(result.content[:500] + "..." if len(result.content) > 500 else result.content)
+ print("=" * 50)
+
+ print(f"\n📊 提取统计:")
+ print(f" • 内容长度: {len(result.content)} 字符")
+ print(f" • 标题: {result.title}")
+ print(f" • 语言: {result.language}")
+ print(f" • 提取时间: {result.extraction_time:.3f}秒")
+
+ if result.content_list:
+ print(f" • 结构化内容块: {len(result.content_list)}个")
+ for i, item in enumerate(result.content_list[:3]): # 显示前3个
+ print(f" [{i + 1}] {item.get('type', 'unknown')}: {item.get('content', '')[:50]}...")
+ else:
+ print("❌ 内容提取失败")
+ print(f"错误信息: {result.error_message}")
+ if result.error_traceback:
+ print(f"错误详情:\n{result.error_traceback}")
+
+except Exception as e:
+ print(f"❌ 提取过程中发生异常: {e}")
diff --git a/examples/trafilatura_extract_demo.py b/examples/trafilatura_extract_demo.py
new file mode 100644
index 0000000..1ee9f3c
--- /dev/null
+++ b/examples/trafilatura_extract_demo.py
@@ -0,0 +1,2873 @@
+import time
+from webmainbench.extractors import ExtractorFactory
+
+# 配置 Trafilatura 抽取器(这里可根据需要添加更多配置)
+config = {}
+
+try:
+ # 创建 Trafilatura 抽取器实例
+ extractor = ExtractorFactory.create("trafilatura", config=config)
+ print(f"✅ Extractor创建成功: {extractor.description}")
+ print(f"📋 版本: {extractor.version}")
+ print(f"⚙️ 配置: {extractor.get_config()}\n")
+except Exception as e:
+ print(f"❌ Extractor创建失败: {e}")
+
+
+# 测试 HTML
+test_html = """
+
+
+Tracking Covid-19 vaccinations in the US
CNN's other Covid-19 trackers
Since vaccinations began in the United States, the federal government has deferred to states and territories on how, when and to whom they administer these shots. While no federal mandate exists, the Centers for Disease Control and Prevention has widely encouraged vaccinations.
+
Here is how many doses have been administered in each state.
+
In Alaska Alabama Arkansas Arizona California Colorado Connecticut Washington, DC Delaware Florida Georgia Hawaii Iowa Idaho Illinois Indiana Kansas Kentucky Louisiana Massachusetts Maryland Maine Michigan Minnesota Missouri Mississippi Montana North Carolina North Dakota Nebraska New Hampshire New Jersey New Mexico Nevada New York Ohio Oklahoma Oregon Pennsylvania Puerto Rico Rhode Island South Carolina South Dakota Tennessee Texas Utah Virginia Vermont Washington Wisconsin West Virginia Wyoming 7 million — or 56.99999999999999% — of the state’s 12.2 million doses doses have been administered. That’s about 143 doses for every hundred residents. Roughly 53.2% of residents are fully vaccinated.
Across the country, about 672.5 million doses have been administered. That translates to 203 doses per hundred people.
Percentage of residents who are fully vaccinated
Less than 30% 30 to 40% 40 to 50% 50 to 60% 60% or more
DC Compare vaccine rollouts by state. Immunization rates are based on state population.
Data from the CDC shows that minority populations are vaccinated at lower rates than their White peers. Only 171.7 million fully vaccinated people reported their ethnicity. Out of that group: More White people (55.7%) were fully vaccinated compared to Black (10.2%), Asian (7%) and Hispanic or Latino (19.8%) residents.
The Kaiser Family Foundation, a national health policy nonprofit, collects data on these breakdowns for states sharing the race and ethnicity of those receiving vaccines.
Location % Asian % Black % Hispanic % White Alabama 2 25 6 63 Alaska 7 3 6 56 Arizona 4 3 19 51 Arkansas 2 14 7 76 California 17 4 31 36 Colorado 3 4 13 75 Connecticut 5 8 15 65 Delaware 4 17 10 61 District of Columbia 6 45 14 49 Florida 9 33 53 Georgia 6 27 9 53 Hawaii 54 1 26 Idaho 2 1 11 82 Illinois 7 11 15 63 Indiana 3 7 7 80 Iowa 2 2 5 93 Kansas 3 4 12 75 Kentucky 2 7 81 Louisiana 3 31 7 58 Maine 2 2 2 82 Maryland 7 27 10 54 Massachusetts 8 7 10 73 Michigan 4 10 5 73 Minnesota 6 5 5 82 Mississippi 2 38 3 56 Missouri 3 11 5 86 Nevada 10 6 27 39 New Hampshire 3 1 3 88 New Jersey 11 9 18 51 New Mexico 3 2 40 42 New York 14 15 21 69 North Carolina 4 20 10 68 Ohio 3 10 4 78 Oklahoma 4 7 12 77 Oregon 6 3 10 74 Pennsylvania* 2 6 7 79 Rhode Island 4 5 16 75 South Carolina 22 6 59 South Dakota <0.01 1 <0.01 92 Tennessee 2 12 6 63 Texas 6 8 36 36 Utah 3 1 12 78 Vermont 2 1 2 95 Virginia 9 17 10 58 Washington 11 4 11 62 West Virginia 3 90 Wisconsin 4 5 6 90
Vaccination data for the line chart comes from Johns Hopkins University’s Centers for Civic Impact. More detailed information about JHU’s sourcing is available on its GitHub repository .
+
Kaiser Family Foundation data representing race and ethnicity comes from a variety of public sources. Persons of Hispanic origin may be of any race. However, states vary in whether they include or exclude Hispanic individuals in racial categories. Additionally, some states vary in whether they include or exclude Hispanic individuals in racial categories in their reporting of vaccination data. We have a made a note in our table which states do so.
+
+"""
+
+print("🔍 开始内容提取...")
+start_time = time.time()
+
+try:
+ result = extractor.extract(test_html)
+ end_time = time.time()
+
+ print(f"⏱️ 提取耗时: {end_time - start_time:.2f}秒\n")
+
+ # 显示提取结果
+ if result.success:
+ print("✅ 内容提取成功!\n")
+
+ print("📄 提取的主要内容:")
+ print("=" * 50)
+ print(result.content[:500] + "..." if len(result.content) > 500 else result.content)
+ print("=" * 50)
+
+ print(f"\n📊 提取统计:")
+ print(f" • 内容长度: {len(result.content)} 字符")
+ print(f" • 标题: {result.title}")
+ print(f" • 语言: {result.language}")
+ print(f" • 提取时间: {result.extraction_time:.3f}秒")
+
+ if result.content_list:
+ print(f" • 结构化内容块: {len(result.content_list)}个")
+ for i, item in enumerate(result.content_list[:3]): # 显示前3个
+ print(f" [{i + 1}] {item.get('type', 'unknown')}: {item.get('content', '')[:50]}...")
+ else:
+ print("❌ 内容提取失败")
+ print(f"错误信息: {result.error_message}")
+ if result.error_traceback:
+ print(f"错误详情:\n{result.error_traceback}")
+
+except Exception as e:
+ print(f"❌ 提取过程中发生异常: {e}")
\ No newline at end of file
diff --git a/requirements.txt b/requirements.txt
index 04708da..f392f6c 100644
--- a/requirements.txt
+++ b/requirements.txt
@@ -1 +1,9 @@
rapidFuzz
+setuptools
+jsonlines
+beautifulsoup4
+requests
+torch
+html2text
+resiliparse
+trafilatura
\ No newline at end of file
diff --git a/tests/test_extractors.py b/tests/test_extractors.py
new file mode 100644
index 0000000..9a86f0b
--- /dev/null
+++ b/tests/test_extractors.py
@@ -0,0 +1,91 @@
+import unittest
+from webmainbench.extractors.factory import ExtractorFactory
+from webmainbench.extractors.base import ExtractionResult
+
+
+class TestExtractors(unittest.TestCase):
+
+ def setUp(self):
+ # 自动发现抽取器
+ ExtractorFactory.auto_discover()
+
+ def test_trafilatura_extractor(self):
+ # 测试 Trafilatura 抽取器
+ extractor = ExtractorFactory.create("trafilatura")
+ html_content = """
+
+
+ Python编程教程
+ 这是一个Python基础教程,展示如何定义函数。
+ def greet(name):
+ ""问候函数""
+ return f"Hello, {name}!"
+
+# 使用示例
+result = greet("World")
+print(result)
+ 这个函数可以用来问候任何人。
+
+
+ """
+ result = extractor.extract(html_content)
+ self.assertEqual(isinstance(result, ExtractionResult), True)
+ self.assertEqual(result.success in [True, False], True)
+
+# def test_magic_html_extractor(self):
+# # 测试 Magic HTML 抽取器
+# try:
+# extractor = ExtractorFactory.create("magic-html")
+# html_content = """
+#
+#
+# Python编程教程
+# 这是一个Python基础教程,展示如何定义函数。
+# def greet(name):
+# ""问候函数""
+# return f"Hello, {name}!"
+#
+# # 使用示例
+# result = greet("World")
+# print(result)
+# 这个函数可以用来问候任何人。
+#
+#
+# """
+# result = extractor.extract(html_content)
+# self.assertEqual(isinstance(result, ExtractionResult), True)
+# self.assertEqual(result.success in [True, False], True)
+# except ValueError as e:
+# # 如果抽取器未注册,跳过测试
+# self.skipTest(f"Magic HTML 抽取器未注册: {e}")
+
+ def test_resiliparse_extractor(self):
+ # 测试 Resiliparse 抽取器
+ try:
+ extractor = ExtractorFactory.create("resiliparse")
+ html_content = """
+
+
+ Python编程教程
+ 这是一个Python基础教程,展示如何定义函数。
+ def greet(name):
+ ""问候函数""
+ return f"Hello, {name}!"
+
+# 使用示例
+result = greet("World")
+print(result)
+ 这个函数可以用来问候任何人。
+
+
+ """
+ result = extractor.extract(html_content)
+ self.assertEqual(isinstance(result, ExtractionResult), True)
+ self.assertEqual(result.success in [True, False], True)
+ except ValueError as e:
+ # 如果抽取器未注册,跳过测试
+ self.skipTest(f"Resiliparse 抽取器未注册: {e}")
+
+
+if __name__ == '__main__':
+ unittest.main()
\ No newline at end of file
diff --git a/tests/test_metrics.py b/tests/test_metrics.py
index 0e308b5..5c38b24 100644
--- a/tests/test_metrics.py
+++ b/tests/test_metrics.py
@@ -7,14 +7,14 @@
class TestContentMetrics(unittest.TestCase):
"""测试内容类型指标"""
-
+
def setUp(self):
"""测试前准备"""
self.calculator = MetricCalculator()
-
+
# 测试数据
self.predicted_content = """# 标题
-
+
这是一段文字内容。
```python
@@ -57,12 +57,10 @@ def hello():
最后是正确的文字内容。
"""
-
-
def test_available_metrics(self):
"""测试可用指标列表"""
metrics = self.calculator.list_available_metrics()
-
+
# 验证必要的指标都存在
expected_metrics = ['code_edit', 'formula_edit', 'table_edit', 'table_TEDS', 'text_edit']
for metric in expected_metrics:
@@ -74,12 +72,13 @@ def test_metric_calculation_success(self):
predicted_content=self.predicted_content,
groundtruth_content=self.groundtruth_content
)
-
+
# 验证所有指标都计算成功
expected_metrics = ['code_edit', 'formula_edit', 'table_edit', 'table_TEDS', 'text_edit', 'overall']
for metric_name in expected_metrics:
self.assertIn(metric_name, results, f"缺少指标结果: {metric_name}")
- self.assertTrue(results[metric_name].success, f"指标 {metric_name} 计算失败: {results[metric_name].error_message}")
+ self.assertTrue(results[metric_name].success,
+ f"指标 {metric_name} 计算失败: {results[metric_name].error_message}")
def test_code_edit_metric(self):
"""测试代码编辑距离指标"""
@@ -87,14 +86,14 @@ def test_code_edit_metric(self):
predicted_content=self.predicted_content,
groundtruth_content=self.groundtruth_content
)
-
+
code_result = results['code_edit']
self.assertTrue(code_result.success)
self.assertIsInstance(code_result.score, float)
# 验证固定内容的确定分数
self.assertAlmostEqual(code_result.score, 0.918367, places=5,
- msg=f"code_edit分数应该是0.918367,实际: {code_result.score}")
-
+ msg=f"code_edit分数应该是0.918367,实际: {code_result.score}")
+
# 验证详细信息
self.assertEqual(code_result.details['content_type'], 'code')
self.assertIn('distance', code_result.details)
@@ -107,14 +106,14 @@ def test_formula_edit_metric(self):
predicted_content=self.predicted_content,
groundtruth_content=self.groundtruth_content
)
-
+
formula_result = results['formula_edit']
self.assertTrue(formula_result.success)
self.assertIsInstance(formula_result.score, float)
# 验证固定内容的确定分数
self.assertAlmostEqual(formula_result.score, 1.000000, places=5,
- msg=f"formula_edit分数应该是1.000000,实际: {formula_result.score}")
-
+ msg=f"formula_edit分数应该是1.000000,实际: {formula_result.score}")
+
# 验证详细信息
self.assertEqual(formula_result.details['content_type'], 'formula')
self.assertIn('distance', formula_result.details)
@@ -125,14 +124,14 @@ def test_table_edit_metric(self):
predicted_content=self.predicted_content,
groundtruth_content=self.groundtruth_content
)
-
+
table_result = results['table_edit']
self.assertTrue(table_result.success)
self.assertIsInstance(table_result.score, float)
# 验证固定内容的确定分数
self.assertAlmostEqual(table_result.score, 0.868852, places=5,
- msg=f"table_edit分数应该是0.868852,实际: {table_result.score}")
-
+ msg=f"table_edit分数应该是0.868852,实际: {table_result.score}")
+
# 验证详细信息
self.assertEqual(table_result.details['content_type'], 'table')
self.assertIn('distance', table_result.details)
@@ -143,14 +142,14 @@ def test_table_teds_metric(self):
predicted_content=self.predicted_content,
groundtruth_content=self.groundtruth_content
)
-
+
teds_result = results['table_TEDS']
self.assertTrue(teds_result.success)
self.assertIsInstance(teds_result.score, float)
# 验证固定内容的确定分数
self.assertAlmostEqual(teds_result.score, 0.300000, places=5,
- msg=f"table_TEDS分数应该是0.300000,实际: {teds_result.score}")
-
+ msg=f"table_TEDS分数应该是0.300000,实际: {teds_result.score}")
+
# 验证详细信息
self.assertEqual(teds_result.details['content_type'], 'table')
@@ -160,14 +159,14 @@ def test_text_edit_metric(self):
predicted_content=self.predicted_content,
groundtruth_content=self.groundtruth_content
)
-
+
text_result = results['text_edit']
self.assertTrue(text_result.success)
self.assertIsInstance(text_result.score, float)
# 验证固定内容的确定分数
self.assertAlmostEqual(text_result.score, 0.769231, places=5,
- msg=f"text_edit分数应该是0.769231,实际: {text_result.score}")
-
+ msg=f"text_edit分数应该是0.769231,实际: {text_result.score}")
+
# 验证详细信息
self.assertEqual(text_result.details['content_type'], 'text')
self.assertIn('distance', text_result.details)
@@ -178,25 +177,25 @@ def test_overall_metric_calculation(self):
predicted_content=self.predicted_content,
groundtruth_content=self.groundtruth_content
)
-
+
# 获取individual指标分数
individual_metrics = ['code_edit', 'formula_edit', 'table_edit', 'table_TEDS', 'text_edit']
individual_scores = []
-
+
for metric_name in individual_metrics:
self.assertIn(metric_name, results)
self.assertTrue(results[metric_name].success)
individual_scores.append(results[metric_name].score)
-
+
# 计算期望的overall分数
expected_overall = sum(individual_scores) / len(individual_scores)
-
+
# 验证overall分数
overall_result = results['overall']
self.assertTrue(overall_result.success)
- self.assertAlmostEqual(overall_result.score, expected_overall, places=5,
- msg="overall分数应该是其他指标的平均值")
-
+ self.assertAlmostEqual(overall_result.score, expected_overall, places=5,
+ msg="overall分数应该是其他指标的平均值")
+
# 验证overall详细信息
self.assertEqual(overall_result.details['source'], 'average_of_all_metrics')
self.assertEqual(overall_result.details['successful_metrics'], len(individual_metrics))
@@ -208,12 +207,12 @@ def test_identical_content(self):
predicted_content=self.groundtruth_content,
groundtruth_content=self.groundtruth_content
)
-
+
# 完全相同的内容应该得到满分
for metric_name in ['code_edit', 'formula_edit', 'table_edit', 'text_edit']:
if metric_name in results and results[metric_name].success:
self.assertAlmostEqual(results[metric_name].score, 1.0, places=5,
- msg=f"相同内容的{metric_name}应该得到满分,实际: {results[metric_name].score}")
+ msg=f"相同内容的{metric_name}应该得到满分,实际: {results[metric_name].score}")
def test_empty_content(self):
"""测试空内容的情况"""
@@ -221,17 +220,17 @@ def test_empty_content(self):
predicted_content="",
groundtruth_content=""
)
-
+
# 空内容应该能正确处理,不应该出错
for metric_name, result in results.items():
if metric_name != 'overall': # overall可能会有特殊处理
- self.assertTrue(result.success or result.score == 0.0,
- f"空内容的{metric_name}应该正确处理")
+ self.assertTrue(result.success or result.score == 0.0,
+ f"空内容的{metric_name}应该正确处理")
class TestErrorHandling(unittest.TestCase):
"""测试错误处理"""
-
+
def setUp(self):
self.calculator = MetricCalculator()
@@ -242,7 +241,7 @@ def test_malformed_content(self):
predicted_content="test",
groundtruth_content="test"
)
-
+
# 不应该有未捕获的异常
self.assertIsInstance(results, dict)
@@ -252,18 +251,18 @@ def test_none_inputs(self):
predicted_content=None,
groundtruth_content=None
)
-
+
# 应该能处理None输入
self.assertIsInstance(results, dict)
class TestRealSampleMetrics(unittest.TestCase):
"""测试基于LLM-WebKit实际提取结果的指标计算"""
-
+
def setUp(self):
"""测试前准备"""
self.calculator = MetricCalculator()
-
+
def test_text_code_sample_edit_distance(self):
"""测试文本+代码样本的编辑距离"""
# 基于实际调试结果的数据
@@ -278,7 +277,7 @@ def hello_world():
```
以上代码展示了一个简单的Python函数。"""
-
+
predicted = """# Python编程示例
这是一段关于Python编程的介绍文本。
@@ -290,25 +289,25 @@ def hello_world():
```
以上代码展示了一个简单的Python函数。"""
-
+
# 计算编辑距离(基于实际调试结果)
results = self.calculator.calculate_all(
predicted_content=predicted,
groundtruth_content=groundtruth
)
-
+
# 验证文本编辑距离(固定内容应该有确定分数)
self.assertIn("text_edit", results)
self.assertTrue(results["text_edit"].success)
self.assertAlmostEqual(results["text_edit"].score, 1.000000, places=5,
- msg=f"text_edit分数应该是1.000000,实际: {results['text_edit'].score}")
-
+ msg=f"text_edit分数应该是1.000000,实际: {results['text_edit'].score}")
+
# 验证代码编辑距离(缺少python标识符导致轻微差异)
self.assertIn("code_edit", results)
self.assertTrue(results["code_edit"].success)
self.assertAlmostEqual(results["code_edit"].score, 0.905797, places=5,
- msg=f"code_edit分数应该是0.905797,实际: {results['code_edit'].score}")
-
+ msg=f"code_edit分数应该是0.905797,实际: {results['code_edit'].score}")
+
def test_table_sample_edit_distance(self):
"""测试表格样本的编辑距离"""
groundtruth = """## 销售数据统计
@@ -317,31 +316,31 @@ def test_table_sample_edit_distance(self):
|------|------|------|
| 产品A | 100 | 1000 |
| 产品B | 200 | 3000 |"""
-
+
predicted = """## 销售数据统计
| 产品 | 销量 | 收入 |
|---|---|---|
| 产品A | 100 | 1000 |
| 产品B | 200 | 3000 |"""
-
+
results = self.calculator.calculate_all(
predicted_content=predicted,
groundtruth_content=groundtruth
)
-
+
# 验证表格编辑距离(分隔符长度差异导致的固定分数)
self.assertIn("table_edit", results)
self.assertTrue(results["table_edit"].success)
self.assertAlmostEqual(results["table_edit"].score, 0.888889, places=5,
- msg=f"table_edit分数应该是0.888889,实际: {results['table_edit'].score}")
-
+ msg=f"table_edit分数应该是0.888889,实际: {results['table_edit'].score}")
+
# 验证TEDS指标(表格结构完全相同,满分)
self.assertIn("table_TEDS", results)
self.assertTrue(results["table_TEDS"].success)
self.assertAlmostEqual(results["table_TEDS"].score, 1.000000, places=5,
- msg=f"table_TEDS分数应该是1.000000,实际: {results['table_TEDS'].score}")
-
+ msg=f"table_TEDS分数应该是1.000000,实际: {results['table_TEDS'].score}")
+
def test_formula_sample_edit_distance(self):
"""测试公式样本的编辑距离"""
groundtruth = """## 数学公式示例
@@ -351,7 +350,7 @@ def test_formula_sample_edit_distance(self):
这是一个行间公式:
$$\\int_{-\\infty}^{\\infty} e^{-x^2} dx = \\sqrt{\\pi}$$"""
-
+
predicted = """## 数学公式示例
这是一个行内公式: \\$E = mc^2\\$
@@ -359,24 +358,24 @@ def test_formula_sample_edit_distance(self):
这是一个行间公式:
\\$\\$\\int_{-\\infty}^{\\infty} e^{-x^2} dx = \\sqrt{\\pi}\\$"""
-
+
results = self.calculator.calculate_all(
predicted_content=predicted,
groundtruth_content=groundtruth
)
-
+
# 验证公式编辑距离(符号转义导致的固定低分)
self.assertIn("formula_edit", results)
self.assertTrue(results["formula_edit"].success)
self.assertAlmostEqual(results["formula_edit"].score, 0.000000, places=5,
- msg=f"formula_edit分数应该是0.000000,实际: {results['formula_edit'].score}")
-
+ msg=f"formula_edit分数应该是0.000000,实际: {results['formula_edit'].score}")
+
# 验证文本编辑距离(去除公式后的纯文本,也受符号转义影响)
self.assertIn("text_edit", results)
self.assertTrue(results["text_edit"].success)
self.assertAlmostEqual(results["text_edit"].score, 0.320000, places=5,
- msg=f"text_edit分数应该是0.320000,实际: {results['text_edit'].score}")
-
+ msg=f"text_edit分数应该是0.320000,实际: {results['text_edit'].score}")
+
def test_overall_score_calculation(self):
"""测试综合分数计算"""
# 使用第一个样本测试综合分数
@@ -391,7 +390,7 @@ def hello_world():
```
以上代码展示了一个简单的Python函数。"""
-
+
predicted = """# Python编程示例
这是一段关于Python编程的介绍文本。
@@ -403,29 +402,29 @@ def hello_world():
```
以上代码展示了一个简单的Python函数。"""
-
+
results = self.calculator.calculate_all(
predicted_content=predicted,
groundtruth_content=groundtruth
)
-
+
# 验证overall分数存在且合理
self.assertIn("overall", results)
self.assertTrue(results["overall"].success)
-
+
# overall应该是所有成功指标的平均值
successful_scores = []
for metric_name, result in results.items():
if metric_name != "overall" and result.success:
successful_scores.append(result.score)
-
+
if successful_scores:
expected_overall = sum(successful_scores) / len(successful_scores)
actual_overall = results["overall"].score
-
+
# 允许小幅计算误差
self.assertAlmostEqual(actual_overall, expected_overall, places=3)
-
+
def test_all_metrics_coverage(self):
"""测试所有6项指标都被计算"""
groundtruth = """# 综合示例
@@ -444,50 +443,50 @@ def test():
| 1 | 2 |
更多文本。"""
-
+
predicted = groundtruth # 使用相同内容测试
-
+
results = self.calculator.calculate_all(
predicted_content=predicted,
groundtruth_content=groundtruth
)
-
+
# 验证所有6项指标都存在
expected_metrics = ["overall", "text_edit", "code_edit", "table_edit", "table_TEDS", "formula_edit"]
-
+
print(f"\n=== 完全相同内容的指标测试结果 ===")
-
+
for metric in expected_metrics:
self.assertIn(metric, results, f"指标 {metric} 缺失")
self.assertTrue(results[metric].success, f"指标 {metric} 计算失败")
-
+
score = results[metric].score
print(f"{metric}: {score:.6f}")
-
+
# 完全相同的内容应该得到满分 1.0
- self.assertAlmostEqual(score, 1.0,
- places=4,
- msg=f"完全相同内容的 {metric} 应该得到满分,实际得分: {score}")
-
+ self.assertAlmostEqual(score, 1.0,
+ places=4,
+ msg=f"完全相同内容的 {metric} 应该得到满分,实际得分: {score}")
+
print("✅ 所有指标都正确得到满分!")
def run_visual_test():
"""运行可视化测试(保留原有的打印功能)"""
print("=== 新指标功能测试 ===\n")
-
+
calculator = MetricCalculator()
-
+
# 显示可用指标
print("可用的指标:")
metrics = calculator.list_available_metrics()
for metric in metrics:
print(f" - {metric}")
print()
-
+
# 测试数据
predicted_content = """# 标题
-
+
这是一段文字内容。
```python
@@ -536,11 +535,11 @@ def hello():
predicted_content=predicted_content,
groundtruth_content=groundtruth_content
)
-
+
# 显示结果
print("\n=== 评测结果 ===")
print("-" * 60)
-
+
for metric_name, result in results.items():
if result.success:
print(f"{metric_name:15}: {result.score:.4f}")
@@ -551,7 +550,7 @@ def hello():
else:
print(f"{metric_name:15}: ERROR - {result.error_message}")
print()
-
+
# 显示详细信息
print("\n=== 详细信息 ===")
for metric_name in ["code_edit", "formula_edit", "table_edit", "text_edit"]:
@@ -565,7 +564,7 @@ def hello():
if __name__ == "__main__":
import sys
-
+
if len(sys.argv) > 1 and sys.argv[1] == "--visual":
# 运行可视化测试
try:
@@ -574,7 +573,8 @@ def hello():
except Exception as e:
print(f"\n❌ 测试失败: {e}")
import traceback
+
traceback.print_exc()
else:
# 运行单元测试
- unittest.main(verbosity=2)
\ No newline at end of file
+ unittest.main(verbosity=2)
\ No newline at end of file
diff --git a/webmainbench/extractors/__init__.py b/webmainbench/extractors/__init__.py
index e94c225..d71ccc9 100644
--- a/webmainbench/extractors/__init__.py
+++ b/webmainbench/extractors/__init__.py
@@ -8,6 +8,9 @@
from .factory import ExtractorFactory
from .llm_webkit_extractor import LlmWebkitExtractor
from .jina_extractor import JinaExtractor
+from .trafilatura_extractor import TrafilaturaExtractor
+from .resiliparse_extractor import ResiliparseExtractor
+from .magic_html_extractor import MagicHtmlExtractor
@@ -17,4 +20,7 @@
"ExtractorFactory",
"LlmWebkitExtractor",
"JinaExtractor",
+ "TrafilaturaExtractor",
+ "ResiliparseExtractor",
+ "MagicHtmlExtractor",
]
\ No newline at end of file
diff --git a/webmainbench/extractors/magic_html_extractor.py b/webmainbench/extractors/magic_html_extractor.py
new file mode 100644
index 0000000..7578de8
--- /dev/null
+++ b/webmainbench/extractors/magic_html_extractor.py
@@ -0,0 +1,80 @@
+# webmainbench/extractors/magic_html_extractor.py
+"""
+Magic HTML extractor implementation.
+"""
+
+from typing import Dict, Any, Optional, List
+from .base import BaseExtractor, ExtractionResult
+from .factory import extractor
+from magic_html import GeneralExtractor
+import re
+import html2text
+
+
+@extractor("magic-html")
+class MagicHtmlExtractor(BaseExtractor):
+ """Extractor using Magic HTML."""
+
+ version = "0.1.5"
+ description = "Magic HTML based content extractor"
+
+ def _setup(self) -> None:
+ """Set up the Magic HTML extractor."""
+ try:
+ self.extractor = GeneralExtractor()
+ except Exception as e:
+ raise RuntimeError(f"Failed to initialize Magic HTML extractor: {e}")
+
+ def _extract_content(self, html: str, url: str = None) -> ExtractionResult:
+ try:
+ # Use Magic HTML for extraction
+ data = self.extractor.extract(html)
+
+ # 从输出中提取所需信息
+ extracted_html = data.get('html', '')
+ markdown = html2text.html2text(extracted_html)
+ title = data.get('title', '')
+ # 简单地将提取的 HTML 作为内容
+ content = markdown
+ # 创建 content_list(简单分割段落)
+ content_list = []
+ if content:
+ paragraphs = content.split('\n\n')
+ for i, para in enumerate(paragraphs):
+ if para.strip():
+ content_list.append({
+ "type": "paragraph",
+ "content": para.strip(),
+ "index": i
+ })
+
+ return ExtractionResult(
+ content=content,
+ # content_list=content_list,
+ title=title,
+ language=self._detect_language(content),
+ success=True
+ )
+
+ except Exception as e:
+ return ExtractionResult.create_error_result(
+ f"Magic HTML extraction failed: {str(e)}"
+ )
+
+
+ def _detect_language(self, content: str) -> Optional[str]:
+ """检测内容语言."""
+ if not content:
+ return None
+
+ # 简单的语言检测逻辑
+ chinese_chars = len(re.findall(r'[\u4e00-\u9fff]', content))
+ english_chars = len(re.findall(r'[a-zA-Z]', content))
+
+ if chinese_chars > english_chars:
+ return "zh"
+ elif english_chars > 0:
+ return "en"
+ else:
+ return None
+
diff --git a/webmainbench/extractors/resiliparse_extractor.py b/webmainbench/extractors/resiliparse_extractor.py
new file mode 100644
index 0000000..e47c87d
--- /dev/null
+++ b/webmainbench/extractors/resiliparse_extractor.py
@@ -0,0 +1,123 @@
+"""
+resiliparse extractor implementation.
+"""
+from typing import Dict, Any, Optional, List
+from dataclasses import dataclass
+from .base import BaseExtractor, ExtractionResult
+from .factory import extractor
+from resiliparse.extract.html2text import extract_plain_text
+import re
+
+@dataclass
+class ResiliparseInferenceConfig:
+ """Configuration for Resiliparse extractor."""
+ main_content: bool = True
+ alt_texts: bool = True
+ links: bool = False
+ form_fields: bool = False
+ noscript: bool = False
+ list_bullets: bool = True
+ preserve_formatting: bool = True
+ comments: bool = True
+ # 可根据需要添加更多resiliparse支持的参数
+
+
+@extractor("resiliparse")
+class ResiliparseExtractor(BaseExtractor):
+ """Extractor using Resiliparse."""
+
+ version = "0.14.5"
+ description = "Resiliparse based content extractor"
+
+ def __init__(self, name: str, config: Optional[Dict[str, Any]] = None):
+ super().__init__(name, config)
+ self.inference_config = ResiliparseInferenceConfig()
+
+ # 应用用户配置
+ if config:
+ for key, value in config.items():
+ if hasattr(self.inference_config, key):
+ setattr(self.inference_config, key, value)
+
+ def _setup(self) -> None:
+ """Set up the Resiliparse extractor."""
+ # 初始化操作
+ pass
+
+ def _extract_content(self, html: str, url: str = None) -> ExtractionResult:
+ """
+ Extract content using Resiliparse.
+
+ Args:
+ html: HTML content to extract from
+ url: Optional URL of the page
+
+ Returns:
+ ExtractionResult instance
+ """
+ try:
+ # 使用配置参数进行内容抽取
+ content = extract_plain_text(
+ html,
+ main_content=self.inference_config.main_content,
+ alt_texts=self.inference_config.alt_texts,
+ links=self.inference_config.links,
+ form_fields=self.inference_config.form_fields,
+ noscript=self.inference_config.noscript,
+ list_bullets=self.inference_config.list_bullets,
+ preserve_formatting=self.inference_config.preserve_formatting,
+ comments=self.inference_config.comments
+ )
+
+ # 创建 content_list(简单分割段落)
+ content_list = []
+ if content:
+ paragraphs = content.split('\n\n')
+ for i, para in enumerate(paragraphs):
+ if para.strip():
+ content_list.append({
+ "type": "paragraph",
+ "content": para.strip(),
+ "index": i
+ })
+
+ return ExtractionResult(
+ content=content,
+ # content_list=content_list,
+ title=self._extract_title(html),
+ language=self._detect_language(content),
+ success=True
+ )
+
+ except Exception as e:
+ return ExtractionResult.create_error_result(
+ f"Resiliparse extraction failed: {str(e)}"
+ )
+
+ def _extract_title(self, html: str) -> Optional[str]:
+ """提取页面标题."""
+ try:
+ import re
+ title_match = re.search(r']*>(.*?) ', html, re.IGNORECASE | re.DOTALL)
+ if title_match:
+ return title_match.group(1).strip()
+ except:
+ pass
+ return None
+
+ def _detect_language(self, content: str) -> Optional[str]:
+ """检测内容语言."""
+ if not content:
+ return None
+
+ # 简单的语言检测逻辑
+ chinese_chars = len(re.findall(r'[\u4e00-\u9fff]', content))
+ english_chars = len(re.findall(r'[a-zA-Z]', content))
+
+ if chinese_chars > english_chars:
+ return "zh"
+ elif english_chars > 0:
+ return "en"
+ else:
+ return None
+
diff --git a/webmainbench/extractors/trafilatura_extractor.py b/webmainbench/extractors/trafilatura_extractor.py
new file mode 100644
index 0000000..f150da1
--- /dev/null
+++ b/webmainbench/extractors/trafilatura_extractor.py
@@ -0,0 +1,121 @@
+
+"""
+trafilatura extractor implementation.
+"""
+from typing import Dict, Any, Optional, List
+from dataclasses import dataclass
+from .base import BaseExtractor, ExtractionResult
+from .factory import extractor
+from trafilatura import extract
+import re
+
+
+@dataclass
+class TrafilaturaInferenceConfig:
+ """Configuration for Trafilatura extractor."""
+ favor_precision: bool = True
+ favor_recall: bool = True
+ include_comments: bool = False
+ include_tables: bool = False
+ # 可根据需要添加更多trafilatura支持的参数
+ include_images: bool = False
+ include_links: bool = False
+
+
+@extractor("trafilatura")
+class TrafilaturaExtractor(BaseExtractor):
+ """Extractor using Trafilatura."""
+
+ version = "2.0.0"
+ description = "Trafilatura based content extractor"
+
+ def __init__(self, name: str, config: Optional[Dict[str, Any]] = None):
+ super().__init__(name, config)
+ self.inference_config = TrafilaturaInferenceConfig()
+
+ # 应用用户配置
+ if config:
+ for key, value in config.items():
+ if hasattr(self.inference_config, key):
+ setattr(self.inference_config, key, value)
+
+ def _setup(self) -> None:
+ """Set up the Trafilatura extractor."""
+ # 初始化操作
+ pass
+
+ def _extract_content(self, html: str, url: str = None) -> ExtractionResult:
+ """
+ Extract content using Trafilatura.
+
+ Args:
+ html: HTML content to extract from
+ url: Optional URL of the page
+
+ Returns:
+ ExtractionResult instance
+ """
+ try:
+ # 使用配置参数进行内容抽取
+ content = extract(
+ html,
+ url=url,
+ favor_precision=self.inference_config.favor_precision,
+ favor_recall=self.inference_config.favor_recall,
+ include_comments=self.inference_config.include_comments,
+ include_tables=self.inference_config.include_tables,
+ include_images=self.inference_config.include_images,
+ include_links=self.inference_config.include_links
+ )
+
+ # 创建 content_list(简单分割段落)
+ content_list = []
+ if content:
+ paragraphs = content.split('\n\n')
+ for i, para in enumerate(paragraphs):
+ if para.strip():
+ content_list.append({
+ "type": "paragraph",
+ "content": para.strip(),
+ "index": i
+ })
+
+ return ExtractionResult(
+ content=content,
+ # content_list=content_list,
+ title=self._extract_title(html),
+ language=self._detect_language(content),
+ success=True
+ )
+
+ except Exception as e:
+ return ExtractionResult.create_error_result(
+ f"Trafilatura extraction failed: {str(e)}"
+ )
+
+ def _extract_title(self, html: str) -> Optional[str]:
+ """提取页面标题."""
+ try:
+ import re
+ title_match = re.search(r']*>(.*?) ', html, re.IGNORECASE | re.DOTALL)
+ if title_match:
+ return title_match.group(1).strip()
+ except:
+ pass
+ return None
+
+ def _detect_language(self, content: str) -> Optional[str]:
+ """检测内容语言."""
+ if not content:
+ return None
+
+ # 简单的语言检测逻辑
+ chinese_chars = len(re.findall(r'[\u4e00-\u9fff]', content))
+ english_chars = len(re.findall(r'[a-zA-Z]', content))
+
+ if chinese_chars > english_chars:
+ return "zh"
+ elif english_chars > 0:
+ return "en"
+ else:
+ return None
From 79140f703aa2e9d8bf058df5a05487a96dd603a3 Mon Sep 17 00:00:00 2001
From: pekopoke <1135796875@qq.com>
Date: Wed, 6 Aug 2025 10:24:10 +0800
Subject: [PATCH 2/3] update requirements
---
requirements.txt | 3 ++-
1 file changed, 2 insertions(+), 1 deletion(-)
diff --git a/requirements.txt b/requirements.txt
index f392f6c..04f79e8 100644
--- a/requirements.txt
+++ b/requirements.txt
@@ -6,4 +6,5 @@ requests
torch
html2text
resiliparse
-trafilatura
\ No newline at end of file
+trafilatura
+https://github.com/opendatalab/magic-html/releases/download/magic_html-0.1.5-released/magic_html-0.1.5-py3-none-any.whl
\ No newline at end of file
From 931e6dd284062e2e1c12e8d06c04044a2079f6ce Mon Sep 17 00:00:00 2001
From: pekopoke <1135796875@qq.com>
Date: Wed, 6 Aug 2025 10:52:58 +0800
Subject: [PATCH 3/3] update requirements
---
examples/basic_usage.py | 8 +++----
tests/test_extractors.py | 52 ++++++++++++++++++++--------------------
2 files changed, 30 insertions(+), 30 deletions(-)
diff --git a/examples/basic_usage.py b/examples/basic_usage.py
index 2e9ac34..2b3b22a 100755
--- a/examples/basic_usage.py
+++ b/examples/basic_usage.py
@@ -801,8 +801,8 @@ def demo_multi_extraction():
# 配置文件路径
data_dir = Path("../data")
- # dataset_path = data_dir / "sample_dataset.jsonl"
- dataset_path = "/home/lulindong/Pycharm_projects/cc/test.jsonl"
+ dataset_path = data_dir / "sample_dataset.jsonl"
+ # dataset_path = "/home/lulindong/Pycharm_projects/cc/test.jsonl"
print(f"📂 数据集文件: {dataset_path}")
@@ -816,8 +816,8 @@ def demo_multi_extraction():
"preserve_formatting": True
}},
- # {"name": "trafilatura", "config": {}},
- # {"name": "magic-html", "config": {}},
+ {"name": "trafilatura", "config": {}},
+ {"name": "magic-html", "config": {}},
]
# 🔧 选择评测模式:内存模式 vs 批处理模式
diff --git a/tests/test_extractors.py b/tests/test_extractors.py
index 9a86f0b..cc59feb 100644
--- a/tests/test_extractors.py
+++ b/tests/test_extractors.py
@@ -32,32 +32,32 @@ def test_trafilatura_extractor(self):
self.assertEqual(isinstance(result, ExtractionResult), True)
self.assertEqual(result.success in [True, False], True)
-# def test_magic_html_extractor(self):
-# # 测试 Magic HTML 抽取器
-# try:
-# extractor = ExtractorFactory.create("magic-html")
-# html_content = """
-#
-#
-# Python编程教程
-# 这是一个Python基础教程,展示如何定义函数。
-# def greet(name):
-# ""问候函数""
-# return f"Hello, {name}!"
-#
-# # 使用示例
-# result = greet("World")
-# print(result)
-# 这个函数可以用来问候任何人。
-#
-#
-# """
-# result = extractor.extract(html_content)
-# self.assertEqual(isinstance(result, ExtractionResult), True)
-# self.assertEqual(result.success in [True, False], True)
-# except ValueError as e:
-# # 如果抽取器未注册,跳过测试
-# self.skipTest(f"Magic HTML 抽取器未注册: {e}")
+ def test_magic_html_extractor(self):
+ # 测试 Magic HTML 抽取器
+ try:
+ extractor = ExtractorFactory.create("magic-html")
+ html_content = """
+
+
+ Python编程教程
+ 这是一个Python基础教程,展示如何定义函数。
+ def greet(name):
+ ""问候函数""
+ return f"Hello, {name}!"
+
+# 使用示例
+result = greet("World")
+print(result)
+ 这个函数可以用来问候任何人。
+
+
+ """
+ result = extractor.extract(html_content)
+ self.assertEqual(isinstance(result, ExtractionResult), True)
+ self.assertEqual(result.success in [True, False], True)
+ except ValueError as e:
+ # 如果抽取器未注册,跳过测试
+ self.skipTest(f"Magic HTML 抽取器未注册: {e}")
def test_resiliparse_extractor(self):
# 测试 Resiliparse 抽取器