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

1"""Tests for agentos.validation.schema_enforcer.""" 

2 

3from __future__ import annotations 

4 

5import pytest 

6from pydantic import BaseModel, Field 

7from agentos.validation.schema_enforcer import ( 

8 SchemaEnforcer, 

9 EnforcerConfig, 

10 FixStrategy, 

11) 

12 

13 

14class SimpleOutput(BaseModel): 

15 """测试用简单输出 schema。""" 

16 

17 name: str 

18 score: float 

19 category: str = "general" 

20 

21 

22class NestedOutput(BaseModel): 

23 """测试用嵌套输出 schema。""" 

24 

25 title: str 

26 items: list[dict] = Field(default_factory=list) 

27 meta: dict = Field(default_factory=dict) 

28 

29 

30@pytest.fixture 

31def enforcer(): 

32 return SchemaEnforcer() 

33 

34 

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" 

43 

44 

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" 

52 

53 

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" 

61 

62 

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" 

70 

71 

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" 

79 

80 

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 

87 

88 

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 == "" 

96 

97 

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}] 

105 

106 

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 

115 

116 

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) 

128 

129 

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