Coverage for agentos/tests/test_schema_enforcer.py: 0%
83 statements
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-06 10:59 +0800
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-06 10:59 +0800
1"""Tests for agentos.validation.schema_enforcer."""
3from __future__ import annotations
5import pytest
6from pydantic import BaseModel, Field
7from agentos.validation.schema_enforcer import (
8 SchemaEnforcer,
9 EnforcerConfig,
10 FixStrategy,
11)
14class SimpleOutput(BaseModel):
15 """测试用简单输出 schema。"""
17 name: str
18 score: float
19 category: str = "general"
22class NestedOutput(BaseModel):
23 """测试用嵌套输出 schema。"""
25 title: str
26 items: list[dict] = Field(default_factory=list)
27 meta: dict = Field(default_factory=dict)
30@pytest.fixture
31def enforcer():
32 return SchemaEnforcer()
35@pytest.mark.asyncio
36async def test_valid_output_passes(enforcer):
37 """合法输出直接通过。"""
38 output = {"name": "task1", "score": 0.95, "category": "code"}
39 result = await enforcer.enforce(output, SimpleOutput)
40 assert result.is_valid
41 assert result.fix_attempts == 0
42 assert result.repaired_output.name == "task1"
45@pytest.mark.asyncio
46async def test_missing_field_fallback(enforcer):
47 """缺失字段使用默认值回退。"""
48 output = {"name": "task2", "score": 0.88}
49 result = await enforcer.enforce(output, SimpleOutput)
50 assert result.is_valid
51 assert result.repaired_output.category == "general"
54@pytest.mark.asyncio
55async def test_json_string_repair(enforcer):
56 """JSON 字符串格式自动修复。"""
57 output = '{"name": "task3", "score": 0.75,}'
58 result = await enforcer.enforce(output, SimpleOutput)
59 assert result.is_valid
60 assert result.repaired_output.name == "task3"
63@pytest.mark.asyncio
64async def test_json_markdown_codeblock_repair(enforcer):
65 """Markdown 代码块包裹的 JSON 自动修复。"""
66 output = '```json\n{"name": "task4", "score": 0.65}\n```'
67 result = await enforcer.enforce(output, SimpleOutput)
68 assert result.is_valid
69 assert result.repaired_output.name == "task4"
72@pytest.mark.asyncio
73async def test_single_quote_json_repair(enforcer):
74 """单引号 JSON 自动修复。"""
75 output = "{'name': 'task5', 'score': 0.55}"
76 result = await enforcer.enforce(output, SimpleOutput)
77 assert result.is_valid
78 assert result.repaired_output.name == "task5"
81@pytest.mark.asyncio
82async def test_extra_field_ok(enforcer):
83 """多余字段不影响校验。"""
84 output = {"name": "task6", "score": 0.45, "extra_field": "ignored"}
85 result = await enforcer.enforce(output, SimpleOutput)
86 assert result.is_valid
89@pytest.mark.asyncio
90async def test_completely_invalid_full_fallback(enforcer):
91 """完全无效时全默认值回退。"""
92 output = {"wrong": "oops"}
93 result = await enforcer.enforce(output, SimpleOutput)
94 assert result.is_valid
95 assert result.repaired_output.name == ""
98@pytest.mark.asyncio
99async def test_nested_output(enforcer):
100 """嵌套 schema 校验。"""
101 output = {"title": "report", "items": [{"a": 1}], "meta": {"page": 1}}
102 result = await enforcer.enforce(output, NestedOutput)
103 assert result.is_valid
104 assert result.repaired_output.items == [{"a": 1}]
107@pytest.mark.asyncio
108async def test_stats_tracking(enforcer):
109 """校验统计正确累加。"""
110 await enforcer.enforce({"name": "x", "score": 1.0}, SimpleOutput)
111 await enforcer.enforce({"bad": True}, SimpleOutput)
112 assert enforcer.stats.total_checks == 2
113 assert enforcer.stats.total_rejections == 1
114 assert enforcer.stats.total_repairs >= 1
117@pytest.mark.asyncio
118async def test_enforce_batch(enforcer):
119 """批量校验。"""
120 outputs = [
121 {"name": "b1", "score": 0.9},
122 {"name": "b2", "score": 0.8},
123 {"name": "b3", "score": 0.7},
124 ]
125 results = await enforcer.enforce_batch(outputs, SimpleOutput)
126 assert len(results) == 3
127 assert all(r.is_valid for r in results)
130@pytest.mark.asyncio
131async def test_fix_strategy_order_respected():
132 """自定义策略顺序生效。"""
133 config = EnforcerConfig(
134 strategy_order=[FixStrategy.FIELD_FALLBACK, FixStrategy.JSON_REPAIR],
135 max_retries=1,
136 )
137 enf = SchemaEnforcer(config)
138 output = "{'name': 's', 'score': 0.3,}"
139 result = await enf.enforce(output, SimpleOutput)
140 assert result.is_valid