Files
YG_FT/backend/tests/test_data_process_generation.py
caoxiaozhu 78fb894307 feat(data_process): 优化问题生成提示语,提升问题自然度
- 系统提示语注入问题风格规则:像真实用户自然提问、避免"请描述/
  请说明/根据文档"等模板化开头、多条问题交替句式,并附正反示例;
  服务端注入对存量任务即时生效。
- 未配置提示语时的后端兜底从一句话充实为与前端同信息量的完整默认。
- 前端标准/思维链/DPO 三套内置默认提示语新增"问题表述自然"要求,
  旧版归档为 _4 常量并注册进 isBuiltInGenerationPrompt 迁移映射,
  自定义提示语不受影响。
- 新增兜底提示语专项测试与风格规则断言,生成测试 32/32 通过。
2026-08-18 15:49:00 +08:00

1021 lines
37 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
from __future__ import annotations
import json
import httpx
import pytest
from app.modules.data_process.generation import (
ModelGenerationError,
chat_completions_url,
generate_model_records,
)
def test_chat_completions_url_accepts_host_base_and_complete_url() -> None:
assert chat_completions_url("www.caoxiaozhu.com") == (
"https://www.caoxiaozhu.com/v1/chat/completions"
)
assert chat_completions_url("https://model.example/v1") == (
"https://model.example/v1/chat/completions"
)
complete = "https://model.example/openai/v1/chat/completions"
assert chat_completions_url(complete) == complete
def test_generate_model_records_uses_prompt_auth_and_stable_split() -> None:
requests: list[httpx.Request] = []
progress_updates: list[tuple[int, int]] = []
def handler(request: httpx.Request) -> httpx.Response:
requests.append(request)
payload = json.loads(request.content)
assert payload["model"] == "qwen-plus"
assert payload["response_format"] == {"type": "json_object"}
assert "客户反馈页面加载慢" in payload["messages"][1]["content"]
assert "你正在生成标准监督微调问答数据" in payload["messages"][0]["content"]
assert "禁止输出分析、推理过程" in payload["messages"][0]["content"]
assert "真实用户自然提出的问题" in payload["messages"][0]["content"]
assert "模板化开头" in payload["messages"][0]["content"]
return httpx.Response(
200,
json={
"choices": [
{
"message": {
"content": json.dumps(
{
"items": [
{
"instruction": "请生成简洁客服回复",
"input": "客户反馈页面加载慢",
"output": "已收到反馈,我们正在排查。",
}
]
},
ensure_ascii=False,
)
}
}
]
},
)
client = httpx.Client(transport=httpx.MockTransport(handler))
records = generate_model_records(
[{"id": "preview-1", "edited_content": "客户反馈页面加载慢"}],
model={
"name": "Qwen",
"online_model_name": "qwen-plus",
"api_url": "model.example",
"api_key": "test-secret",
},
config={
"generation_prompt": "请处理:{{ content }}",
"json_mode": True,
"temperature": 0.2,
"max_tokens": 512,
},
task_id="task-1",
split={"train": 100, "validation": 0, "test": 0},
qa_pairs_per_item=1,
client=client,
on_progress=lambda processed, total: progress_updates.append((processed, total)),
)
assert len(records) == 1
assert records[0]["status"] == "valid"
assert records[0]["split"] == "train"
assert requests[0].headers["Authorization"] == "Bearer test-secret"
assert progress_updates == [(1, 1)]
def test_generate_model_records_falls_back_to_rich_default_prompt() -> None:
captured: dict[str, str] = {}
def handler(request: httpx.Request) -> httpx.Response:
payload = json.loads(request.content)
captured["system"] = payload["messages"][0]["content"]
captured["user"] = payload["messages"][1]["content"]
return httpx.Response(
200,
json={
"choices": [
{
"message": {
"content": json.dumps(
{
"items": [
{
"instruction": "平台如何控制不同角色的菜单可见性?",
"input": "",
"output": "按角色分配权限。",
}
]
},
ensure_ascii=False,
)
}
}
]
},
)
client = httpx.Client(transport=httpx.MockTransport(handler))
records = generate_model_records(
[{"id": "preview-1", "edited_content": "平台按角色分配菜单权限。"}],
model={
"name": "Qwen",
"online_model_name": "qwen-plus",
"api_url": "model.example",
},
config={},
task_id="task-fallback",
split={"train": 100, "validation": 0, "test": 0},
qa_pairs_per_item=1,
client=client,
)
assert len(records) == 1
assert records[0]["status"] == "valid"
assert "数据生成专家" in captured["system"]
assert "真实用户自然提出的问题" in captured["system"]
assert "平台按角色分配菜单权限。" in captured["user"]
def test_generate_model_records_builds_native_dpo_pair() -> None:
def handler(request: httpx.Request) -> httpx.Response:
payload = json.loads(request.content)
system_prompt = payload["messages"][0]["content"]
assert '"chosen"' in system_prompt
assert '"rejected"' in system_prompt
assert "直接偏好优化" in system_prompt
return httpx.Response(
200,
json={
"choices": [{
"message": {
"content": json.dumps({
"items": [{
"instruction": "系统如何处理扫描 PDF",
"input": "",
"chosen": "仅在没有文本层时调用 OCR并保留页码。",
"rejected": "所有 PDF 都重复执行 OCR。",
}],
}, ensure_ascii=False),
},
}],
},
)
records = generate_model_records(
[{"id": "preview-dpo", "edited_content": "扫描 PDF 缺少文本层时执行 OCR。"}],
model={"name": "model", "api_url": "https://model.example/v1"},
config={"output_type": "dpo", "generation_retries": 0},
task_id="task-dpo",
split={"train": 100, "validation": 0, "test": 0},
qa_pairs_per_item=1,
client=httpx.Client(transport=httpx.MockTransport(handler)),
)
assert records[0]["status"] == "valid"
assert records[0]["chosen"] == "仅在没有文本层时调用 OCR,并保留页码。"
assert records[0]["rejected"] == "所有 PDF 都重复执行 OCR。"
assert records[0]["output"] == records[0]["chosen"]
def test_generate_model_records_rejects_equal_dpo_pair() -> None:
def handler(_request: httpx.Request) -> httpx.Response:
return httpx.Response(
200,
json={"choices": [{"message": {"content": json.dumps({
"items": [{
"instruction": "问题",
"chosen": "相同回答",
"rejected": "相同回答",
}],
}, ensure_ascii=False)}}]},
)
records = generate_model_records(
[{"id": "preview-dpo-invalid", "edited_content": "来源"}],
model={"name": "model", "api_url": "https://model.example/v1"},
config={"output_type": "dpo", "generation_retries": 0},
task_id="task-dpo-invalid",
split={"train": 100, "validation": 0, "test": 0},
qa_pairs_per_item=1,
client=httpx.Client(transport=httpx.MockTransport(handler)),
)
assert records[0]["status"] == "invalid"
assert "chosen equals rejected" in records[0]["error"]
def test_minimax_m3_uses_split_reasoning_and_completion_token_budget() -> None:
requests: list[dict[str, object]] = []
def handler(request: httpx.Request) -> httpx.Response:
payload = json.loads(request.content)
requests.append(payload)
return httpx.Response(
200,
json={
"choices": [{
"finish_reason": "stop",
"message": {
"reasoning_content": "模型内部思考不应混入业务 JSON",
"content": json.dumps({
"items": [{
"instruction": "申请编号有什么作用?",
"reasoning": "来源说明它用于标识报销申请。",
"answer": "它用于唯一标识一笔报销申请。",
}],
}, ensure_ascii=False),
},
}],
"output_sensitive": False,
"base_resp": {"status_code": 0, "status_msg": ""},
},
)
records = generate_model_records(
[{"id": "preview-minimax", "edited_content": "申请编号用于标识报销申请。"}],
model={
"name": "MiniMax",
"online_model_name": "MiniMax-M3",
"api_url": "https://api.minimaxi.com/v1",
},
config={
"output_type": "reasoning",
"json_mode": True,
"max_tokens": 1024,
"generation_retries": 0,
},
task_id="task-minimax",
split={"train": 100, "validation": 0, "test": 0},
qa_pairs_per_item=1,
client=httpx.Client(transport=httpx.MockTransport(handler)),
)
assert records[0]["status"] == "valid"
assert len(requests) == 1
assert requests[0]["reasoning_split"] is True
assert requests[0]["max_completion_tokens"] >= 4096
assert "max_tokens" not in requests[0]
assert "response_format" not in requests[0]
def test_minimax_m3_keeps_larger_configured_completion_budget() -> None:
requests: list[dict[str, object]] = []
def handler(request: httpx.Request) -> httpx.Response:
requests.append(json.loads(request.content))
return httpx.Response(
200,
json={
"choices": [{
"message": {
"content": json.dumps({
"items": [{
"instruction": "问题",
"output": "这是满足测试要求的完整答案。",
}],
}, ensure_ascii=False),
},
}],
},
)
generate_model_records(
[{"id": "preview-minimax-budget", "edited_content": "来源正文"}],
model={
"online_model_name": "MiniMax-M3",
"api_url": "https://api.minimax.io/v1",
},
config={"max_tokens": 8192, "generation_retries": 0},
task_id="task-minimax-budget",
split={"train": 100, "validation": 0, "test": 0},
qa_pairs_per_item=1,
client=httpx.Client(transport=httpx.MockTransport(handler)),
)
assert requests[0]["max_completion_tokens"] == 8192
def test_minimax_m3_name_on_custom_proxy_keeps_generic_openai_parameters() -> None:
requests: list[dict[str, object]] = []
def handler(request: httpx.Request) -> httpx.Response:
requests.append(json.loads(request.content))
return httpx.Response(
200,
json={
"choices": [{
"message": {
"content": json.dumps({
"items": [{
"instruction": "问题",
"output": "这是代理服务返回的完整答案。",
}],
}, ensure_ascii=False),
},
}],
},
)
generate_model_records(
[{"id": "preview-minimax-proxy", "edited_content": "来源正文"}],
model={
"online_model_name": "MiniMax-M3",
"api_url": "https://model-proxy.example/v1",
},
config={"max_tokens": 1024, "json_mode": True, "generation_retries": 0},
task_id="task-minimax-proxy",
split={"train": 100, "validation": 0, "test": 0},
qa_pairs_per_item=1,
client=httpx.Client(transport=httpx.MockTransport(handler)),
)
assert requests[0]["max_tokens"] == 1024
assert requests[0]["response_format"] == {"type": "json_object"}
assert "reasoning_split" not in requests[0]
assert "max_completion_tokens" not in requests[0]
def test_generate_model_records_extracts_json_surrounded_by_model_explanation() -> None:
content = "模型结果如下:\n```json\n" + json.dumps(
{
"items": [{
"instruction": "字段有什么作用?",
"output": "该字段用于唯一标识记录。",
}],
},
ensure_ascii=False,
) + "\n```\n生成完毕。"
client = httpx.Client(
transport=httpx.MockTransport(
lambda _: httpx.Response(
200,
json={"choices": [{"message": {"content": content}}]},
)
)
)
records = generate_model_records(
[{"id": "preview-explanation", "edited_content": "字段用于唯一标识记录。"}],
model={"name": "model", "api_url": "https://model.example/v1"},
config={"generation_retries": 0},
task_id="task-explanation",
split={"train": 100, "validation": 0, "test": 0},
qa_pairs_per_item=1,
client=client,
)
assert records[0]["status"] == "valid"
assert records[0]["output"] == "该字段用于唯一标识记录。"
def test_generate_model_records_reports_token_truncation_instead_of_json_error() -> None:
client = httpx.Client(
transport=httpx.MockTransport(
lambda _: httpx.Response(
200,
json={
"choices": [{
"finish_reason": "length",
"message": {"content": ""},
}],
"output_sensitive": False,
},
)
)
)
records = generate_model_records(
[{"id": "preview-truncated", "edited_content": "来源正文"}],
model={"name": "model", "api_url": "https://model.example/v1"},
config={"generation_retries": 0},
task_id="task-truncated",
split={"train": 100, "validation": 0, "test": 0},
qa_pairs_per_item=1,
client=client,
)
assert records[0]["status"] == "invalid"
assert "Token" in records[0]["error"]
assert "截断" in records[0]["error"]
def test_token_truncation_is_not_retried_even_when_json_looks_complete() -> None:
request_count = 0
content = json.dumps({
"items": [{
"instruction": "问题",
"output": "表面完整但服务端已声明截断。",
}],
}, ensure_ascii=False)
def handler(_: httpx.Request) -> httpx.Response:
nonlocal request_count
request_count += 1
return httpx.Response(
200,
json={
"choices": [{
"finish_reason": "length",
"message": {"content": content},
}],
},
)
records = generate_model_records(
[{"id": "preview-length", "edited_content": "来源正文"}],
model={"name": "model", "api_url": "https://model.example/v1"},
config={"generation_retries": 5},
task_id="task-length",
split={"train": 100, "validation": 0, "test": 0},
qa_pairs_per_item=1,
client=httpx.Client(transport=httpx.MockTransport(handler)),
)
assert request_count == 1
assert records[0]["status"] == "invalid"
assert "finish_reason=length" in records[0]["error"]
def test_sensitive_model_response_is_not_retried_or_saved() -> None:
request_count = 0
def handler(_: httpx.Request) -> httpx.Response:
nonlocal request_count
request_count += 1
return httpx.Response(
200,
json={
"choices": [{
"finish_reason": "stop",
"message": {"content": "{}"},
}],
"output_sensitive": True,
"base_resp": {"status_code": 1027, "status_msg": "output sensitive"},
},
)
records = generate_model_records(
[{"id": "preview-sensitive", "edited_content": "来源正文"}],
model={"name": "model", "api_url": "https://model.example/v1"},
config={"generation_retries": 5},
task_id="task-sensitive",
split={"train": 100, "validation": 0, "test": 0},
qa_pairs_per_item=1,
client=httpx.Client(transport=httpx.MockTransport(handler)),
)
assert request_count == 1
assert records[0]["status"] == "invalid"
assert "安全拦截" in records[0]["error"]
assert "1027" in records[0]["error"]
def test_empty_model_content_can_retry_then_succeed() -> None:
request_count = 0
def handler(_: httpx.Request) -> httpx.Response:
nonlocal request_count
request_count += 1
if request_count == 1:
return httpx.Response(
200,
json={"choices": [{"finish_reason": "stop", "message": {"content": ""}}]},
)
return httpx.Response(
200,
json={
"choices": [{
"finish_reason": "stop",
"message": {
"content": json.dumps({
"items": [{
"instruction": "问题",
"output": "第二次请求返回了完整答案。",
}],
}, ensure_ascii=False),
},
}],
},
)
records = generate_model_records(
[{"id": "preview-empty-retry", "edited_content": "来源正文"}],
model={"name": "model", "api_url": "https://model.example/v1"},
config={"generation_retries": 1},
task_id="task-empty-retry",
split={"train": 100, "validation": 0, "test": 0},
qa_pairs_per_item=1,
client=httpx.Client(transport=httpx.MockTransport(handler)),
)
assert request_count == 2
assert records[0]["status"] == "valid"
def test_multiple_top_level_json_documents_are_rejected_as_ambiguous() -> None:
first = json.dumps({
"items": [{"instruction": "问题一", "output": "答案一"}],
}, ensure_ascii=False)
second = json.dumps({
"items": [{"instruction": "问题二", "output": "答案二"}],
}, ensure_ascii=False)
client = httpx.Client(
transport=httpx.MockTransport(
lambda _: httpx.Response(
200,
json={"choices": [{"message": {"content": f"{first}\n{second}"}}]},
)
)
)
records = generate_model_records(
[{"id": "preview-ambiguous", "edited_content": "来源正文"}],
model={"name": "model", "api_url": "https://model.example/v1"},
config={"generation_retries": 0},
task_id="task-ambiguous",
split={"train": 100, "validation": 0, "test": 0},
qa_pairs_per_item=1,
client=client,
)
assert records[0]["status"] == "invalid"
assert "多个 JSON" in records[0]["error"]
def test_generate_model_records_builds_reasoning_output_with_think_tags() -> None:
def handler(request: httpx.Request) -> httpx.Response:
payload = json.loads(request.content)
system_prompt = payload["messages"][0]["content"]
assert '"reasoning":"...","answer":"..."' in system_prompt
assert "你正在生成用于训练推理模型的思维链数据" in system_prompt
assert "推理详细程度为“普通”" in system_prompt
assert "系统会在保存时统一组装" in system_prompt
content = "<think>模型接口自己的分析</think>" + json.dumps(
{
"items": [
{
"instruction": "计算两项费用合计",
"input": "交通费 30 元,餐费 20 元",
"reasoning": "先识别两项费用,再计算 30 + 20。",
"answer": "合计 50 元。",
}
]
},
ensure_ascii=False,
)
return httpx.Response(
200,
json={"choices": [{"message": {"content": content}}]},
)
records = generate_model_records(
[{"id": "preview-reasoning", "edited_content": "交通费 30 元,餐费 20 元"}],
model={"name": "model", "api_url": "https://model.example/v1"},
config={"output_type": "reasoning"},
task_id="task-reasoning",
split={"train": 100, "validation": 0, "test": 0},
qa_pairs_per_item=1,
client=httpx.Client(transport=httpx.MockTransport(handler)),
)
assert records[0]["status"] == "valid"
assert records[0]["output"] == (
"<think>\n先识别两项费用,再计算 30 + 20。\n</think>\n合计 50 元。"
)
def test_generate_model_records_uses_detailed_reasoning_instruction() -> None:
def handler(request: httpx.Request) -> httpx.Response:
system_prompt = json.loads(request.content)["messages"][0]["content"]
assert "推理详细程度为“详细”" in system_prompt
assert "完整展开问题条件、来源依据、中间计算或推导" in system_prompt
return httpx.Response(
200,
json={
"choices": [
{
"message": {
"content": json.dumps(
{
"items": [
{
"instruction": "计算报销总额",
"reasoning": "条件为交通费 30 元和餐费 20 元。分别核对后相加30 + 20 = 50。",
"answer": "报销总额为 50 元。",
}
]
},
ensure_ascii=False,
)
}
}
]
},
)
records = generate_model_records(
[{"id": "preview-detailed", "edited_content": "交通费 30 元,餐费 20 元"}],
model={"name": "model", "api_url": "https://model.example/v1"},
config={"output_type": "reasoning", "reasoning_detail": "detailed"},
task_id="task-detailed",
split={"train": 100, "validation": 0, "test": 0},
qa_pairs_per_item=1,
client=httpx.Client(transport=httpx.MockTransport(handler)),
)
assert records[0]["status"] == "valid"
assert "分别核对后相加" in records[0]["output"]
def test_generate_model_records_marks_reasoning_without_reasoning_field_invalid() -> None:
response = {
"choices": [
{
"message": {
"content": json.dumps(
{
"items": [
{
"instruction": "问题",
"answer": "只有最终答案",
}
]
},
ensure_ascii=False,
)
}
}
]
}
client = httpx.Client(
transport=httpx.MockTransport(lambda _: httpx.Response(200, json=response))
)
records = generate_model_records(
[{"id": "preview-missing-reasoning", "edited_content": "来源正文"}],
model={"name": "model", "api_url": "https://model.example/v1"},
config={"output_type": "reasoning"},
task_id="task-missing-reasoning",
split={"train": 100, "validation": 0, "test": 0},
qa_pairs_per_item=1,
client=client,
)
assert records[0]["status"] == "invalid"
assert records[0]["output"] == "只有最终答案"
assert "reasoning" in records[0]["error"]
def test_standard_output_removes_model_think_block() -> None:
content = json.dumps(
{
"items": [
{
"instruction": "问题",
"output": "<think>不应保存的分析</think>最终答案",
}
]
},
ensure_ascii=False,
)
client = httpx.Client(
transport=httpx.MockTransport(
lambda _: httpx.Response(
200,
json={"choices": [{"message": {"content": content}}]},
)
)
)
records = generate_model_records(
[{"id": "preview-standard", "edited_content": "来源正文"}],
model={"name": "model", "api_url": "https://model.example/v1"},
config={"output_type": "standard"},
task_id="task-standard",
split={"train": 100, "validation": 0, "test": 0},
qa_pairs_per_item=1,
client=client,
)
assert records[0]["output"] == "最终答案"
def test_generate_model_records_keeps_partial_failure_for_manual_repair() -> None:
client = httpx.Client(
transport=httpx.MockTransport(
lambda _: httpx.Response(200, json={"choices": [{"message": {"content": "not-json"}}]})
)
)
records = generate_model_records(
[{"id": "preview-1", "edited_content": "来源正文"}],
model={"name": "model", "api_url": "https://model.example/v1"},
config={"generation_retries": 1},
task_id="task-1",
split={"train": 80, "validation": 10, "test": 10},
qa_pairs_per_item=1,
client=client,
)
assert len(records) == 1
assert records[0]["status"] == "invalid"
assert records[0]["error"]
def test_generate_model_records_batches_fifty_results_with_unique_ids() -> None:
requests: list[httpx.Request] = []
def handler(request: httpx.Request) -> httpx.Response:
requests.append(request)
batch_start = (len(requests) - 1) * 10 + 1
batch_end = batch_start + 9
payload = json.loads(request.content)
system_prompt = payload["messages"][0]["content"]
assert "items 必须包含 10 条" in system_prompt
assert f"{batch_start}-{batch_end}" in system_prompt
return httpx.Response(
200,
json={
"choices": [
{
"message": {
"content": json.dumps(
{
"items": [
{
"instruction": "同一问题",
"input": "来源正文",
"output": "同一答案",
}
for _ in range(batch_start, batch_end + 1)
]
},
ensure_ascii=False,
)
}
}
]
},
)
records = generate_model_records(
[{"id": "preview-50", "edited_content": "来源正文"}],
model={"name": "model", "api_url": "https://model.example/v1"},
config={},
task_id="task-50",
split={"train": 100, "validation": 0, "test": 0},
qa_pairs_per_item=50,
client=httpx.Client(transport=httpx.MockTransport(handler)),
)
assert len(requests) == 5
assert len(records) == 50
assert len({record["id"] for record in records}) == 50
assert {record["instruction"] for record in records} == {"同一问题"}
assert all(record["status"] == "valid" for record in records)
def test_generate_model_records_preserves_successful_batches_when_one_fails() -> None:
request_count = 0
def handler(_: httpx.Request) -> httpx.Response:
nonlocal request_count
request_count += 1
if request_count == 2:
return httpx.Response(500)
return httpx.Response(
200,
json={
"choices": [
{
"message": {
"content": json.dumps(
{
"items": [
{
"instruction": f"问题 {index}",
"output": f"答案 {index}",
}
for index in range(1, 11)
]
},
ensure_ascii=False,
)
}
}
]
},
)
records = generate_model_records(
[{"id": "preview-partial", "edited_content": "来源正文"}],
model={"name": "model", "api_url": "https://model.example/v1"},
config={"generation_retries": 0},
task_id="task-partial",
split={"train": 100, "validation": 0, "test": 0},
qa_pairs_per_item=20,
client=httpx.Client(transport=httpx.MockTransport(handler)),
)
assert len(records) == 11
assert sum(record["status"] == "valid" for record in records) == 10
failed = next(record for record in records if record["status"] == "invalid")
assert "第 11-20 条" in failed["instruction"]
assert len({record["id"] for record in records}) == len(records)
def test_generate_model_records_retries_short_batch_then_marks_it_invalid() -> None:
request_count = 0
def handler(_: httpx.Request) -> httpx.Response:
nonlocal request_count
request_count += 1
return httpx.Response(
200,
json={
"choices": [
{
"message": {
"content": json.dumps(
{
"items": [
{
"instruction": "只有一条",
"output": "不足本批要求数量",
}
]
},
ensure_ascii=False,
)
}
}
]
},
)
records = generate_model_records(
[{"id": "preview-short", "edited_content": "来源正文"}],
model={"name": "model", "api_url": "https://model.example/v1"},
config={"generation_retries": 1},
task_id="task-short",
split={"train": 100, "validation": 0, "test": 0},
qa_pairs_per_item=10,
client=httpx.Client(transport=httpx.MockTransport(handler)),
)
assert request_count == 2
assert len(records) == 1
assert records[0]["status"] == "invalid"
assert "expected 10, got 1" in records[0]["error"]
def test_generate_model_records_does_not_retry_non_retryable_http_errors() -> None:
request_count = 0
def handler(_: httpx.Request) -> httpx.Response:
nonlocal request_count
request_count += 1
return httpx.Response(401, json={"error": {"message": "unauthorized"}})
records = generate_model_records(
[{"id": "preview-auth", "edited_content": "来源内容"}],
model={
"api_url": "https://model.example/v1",
"online_model_name": "test-model",
"api_key": "invalid",
},
config={"generation_retries": 5},
task_id="task-auth",
split={"train": 100, "validation": 0, "test": 0},
qa_pairs_per_item=1,
client=httpx.Client(transport=httpx.MockTransport(handler)),
)
assert request_count == 1
assert records[0]["status"] == "invalid"
assert "401" in records[0]["error"]
@pytest.mark.parametrize("status_code", [408, 425, 429, 500])
def test_generate_model_records_retries_retryable_http_statuses(
status_code: int,
) -> None:
request_count = 0
def handler(_: httpx.Request) -> httpx.Response:
nonlocal request_count
request_count += 1
if request_count == 1:
return httpx.Response(status_code)
return httpx.Response(
200,
json={
"choices": [{
"message": {
"content": json.dumps({
"items": [{
"instruction": "来源内容是什么?",
"output": "这是用于验证可重试错误的来源内容。",
}],
}, ensure_ascii=False),
},
}],
},
)
records = generate_model_records(
[{"id": "preview-retryable", "edited_content": "来源内容"}],
model={"name": "model", "api_url": "https://model.example/v1"},
config={"generation_retries": 1},
task_id="task-retryable",
split={"train": 100, "validation": 0, "test": 0},
qa_pairs_per_item=1,
client=httpx.Client(transport=httpx.MockTransport(handler)),
)
assert request_count == 2
assert records[0]["status"] == "valid"
def test_generate_model_records_retries_transient_network_errors() -> None:
request_count = 0
def handler(request: httpx.Request) -> httpx.Response:
nonlocal request_count
request_count += 1
if request_count == 1:
raise httpx.ConnectError("temporary connection failure", request=request)
return httpx.Response(
200,
json={
"choices": [{
"message": {
"content": json.dumps({
"items": [{
"instruction": "网络恢复了吗?",
"output": "临时连接错误后,第二次模型请求已经成功。",
}],
}, ensure_ascii=False),
},
}],
},
)
records = generate_model_records(
[{"id": "preview-network", "edited_content": "网络重试来源"}],
model={"name": "model", "api_url": "https://model.example/v1"},
config={"generation_retries": 1},
task_id="task-network",
split={"train": 100, "validation": 0, "test": 0},
qa_pairs_per_item=1,
client=httpx.Client(transport=httpx.MockTransport(handler)),
)
assert request_count == 2
assert records[0]["status"] == "valid"
@pytest.mark.parametrize("qa_pairs_per_item", [0, 51])
def test_generate_model_records_rejects_out_of_range_count(
qa_pairs_per_item: int,
) -> None:
with pytest.raises(ModelGenerationError, match=r"\[1, 50\]"):
generate_model_records(
[],
model={"name": "model", "api_url": "https://model.example/v1"},
config={},
task_id="task-invalid",
split={"train": 100, "validation": 0, "test": 0},
qa_pairs_per_item=qa_pairs_per_item,
)
def test_generate_model_records_rejects_unknown_output_type() -> None:
with pytest.raises(ModelGenerationError, match="output_type"):
generate_model_records(
[],
model={"name": "model", "api_url": "https://model.example/v1"},
config={"output_type": "unknown"},
task_id="task-invalid-output",
split={"train": 100, "validation": 0, "test": 0},
qa_pairs_per_item=1,
)
def test_generate_model_records_rejects_unknown_reasoning_detail() -> None:
with pytest.raises(ModelGenerationError, match="reasoning_detail"):
generate_model_records(
[],
model={"name": "model", "api_url": "https://model.example/v1"},
config={"output_type": "reasoning", "reasoning_detail": "verbose"},
task_id="task-invalid-reasoning-detail",
split={"train": 100, "validation": 0, "test": 0},
qa_pairs_per_item=1,
)