fix(sampler): validate level param rejects invalid values
Add ValueError guard at the top of sample_material_v2 for level not in
{1, 2, 3}, preventing silent fallthrough to L1 sampling. Add unit test
for the new validation.
Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
@@ -525,8 +525,12 @@ def sample_material_v2(
|
|||||||
MaterialContext 实例。
|
MaterialContext 实例。
|
||||||
|
|
||||||
异常:
|
异常:
|
||||||
|
ValueError: level 不在 {1, 2, 3} 中。
|
||||||
RuntimeError: 耗尽 max_attempts 次尝试仍无法满足约束。
|
RuntimeError: 耗尽 max_attempts 次尝试仍无法满足约束。
|
||||||
"""
|
"""
|
||||||
|
if level not in (1, 2, 3):
|
||||||
|
msg = f"level 必须为 1、2 或 3,收到: {level}"
|
||||||
|
raise ValueError(msg)
|
||||||
|
|
||||||
for attempt in range(max_attempts):
|
for attempt in range(max_attempts):
|
||||||
# Phase 1: 按层级采样候选节点
|
# Phase 1: 按层级采样候选节点
|
||||||
|
|||||||
@@ -291,6 +291,21 @@ class TestSampleMaterialV2:
|
|||||||
has_related = any(pattern in s for s in result.subtitle_sentences)
|
has_related = any(pattern in s for s in result.subtitle_sentences)
|
||||||
assert has_related
|
assert has_related
|
||||||
|
|
||||||
|
def test_invalid_level_raises(self, real_tree: TreeIndex) -> None:
|
||||||
|
"""无效的 level 参数应抛出 ValueError。"""
|
||||||
|
from app.question_gen.sampler_v2 import sample_material_v2
|
||||||
|
|
||||||
|
rng = random.Random(42)
|
||||||
|
with pytest.raises(ValueError, match="level 必须为"):
|
||||||
|
sample_material_v2(
|
||||||
|
tree=real_tree,
|
||||||
|
task_type="Object Recognition",
|
||||||
|
used_node_ids=set(),
|
||||||
|
rng=rng,
|
||||||
|
level=4,
|
||||||
|
constraint=RETRIEVAL_FAMILY.sampling,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class TestValidateSamplingConstraints:
|
class TestValidateSamplingConstraints:
|
||||||
"""_validate_sampling_constraints 辅助函数测试。"""
|
"""_validate_sampling_constraints 辅助函数测试。"""
|
||||||
|
|||||||
Reference in New Issue
Block a user