feat(data-process): 完善后台生成与失败重试
This commit is contained in:
@@ -88,6 +88,344 @@ def test_generate_model_records_uses_prompt_auth_and_stable_split() -> None:
|
||||
assert progress_updates == [(1, 1)]
|
||||
|
||||
|
||||
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)
|
||||
@@ -414,6 +752,112 @@ def test_generate_model_records_retries_short_batch_then_marks_it_invalid() -> N
|
||||
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,
|
||||
|
||||
Reference in New Issue
Block a user