fix: 修复 Description 生成格式并按需轮询后台状态
Deploy production / deploy (push) Successful in 56s

This commit is contained in:
yuxuanhui
2026-09-10 15:32:59 +08:00
parent 1a7d92a708
commit a39894a9f1
3 changed files with 279 additions and 38 deletions
+164 -6
View File
@@ -190,8 +190,21 @@ def test_description_nonempty(bad):
Description(**{**FIELDS, "idea": bad})
async def test_ai_uses_independent_model_shared_connection_without_platform_write(app, logged_in):
platform = await setup(app)
@pytest.mark.parametrize("kind", ["REGULAR", "SUPER"])
@pytest.mark.parametrize("format", ["fields", "complete_text", "single_newlines"])
async def test_ai_uses_independent_model_shared_connection_without_platform_write(
app, logged_in, kind, format
):
raw = (
alpha()
if kind == "REGULAR"
else alpha(type="SUPER", selection={"code": "rank(close)"}, combo={"code": "alpha"})
)
platform = await setup(app, raw)
sections = ["regular"] if kind == "REGULAR" else ["selection", "combo"]
value = FIELDS if format == "fields" else Description(**FIELDS).text()
if format == "single_newlines":
value = value.replace("\n\n", "\n")
seen = []
@asynccontextmanager
@@ -207,9 +220,7 @@ async def test_ai_uses_independent_model_shared_connection_without_platform_writ
yield FunctionModel(
function=lambda messages, info: ModelResponse(
parts=[
ToolCallPart(
info.output_tools[0].name, {"descriptions": {"regular": Description(**FIELDS).text()}}
),
ToolCallPart(info.output_tools[0].name, {"descriptions": dict.fromkeys(sections, value)}),
]
)
)
@@ -231,7 +242,7 @@ async def test_ai_uses_independent_model_shared_connection_without_platform_writ
"/api/v1/alphas/alpha1/description/generate", json={"snapshot": state["snapshot"]}
)
assert generated.status_code == 200, generated.text
assert generated.json()["descriptions"] == {"regular": Description(**FIELDS).text()}
assert generated.json()["descriptions"] == dict.fromkeys(sections, Description(**FIELDS).text())
assert seen == [("https://model.test/v1", "description-model", "responses", "shared-secret")]
assert not platform.calls
async with app.state.sessions() as db:
@@ -269,6 +280,153 @@ async def test_generation_requires_config_and_rejects_invalid_output(app, logged
app.state.ai.model_factory = broken
response = await logged_in.post(path, json={"snapshot": state["snapshot"]})
assert response.status_code == 502 and "shared-secret" not in response.text
assert "格式" in response.json()["detail"]
@pytest.mark.parametrize("failure", ["length", "sections", "plain_text", "paragraph_breaks"])
@pytest.mark.parametrize("corrected", [True, False])
async def test_generation_corrects_output_once_and_keeps_platform_untouched(
app, logged_in, failure, corrected, caplog
):
from pydantic_ai.messages import TextPart
platform = await setup(app)
await logged_in.put(
"/api/v1/ai/settings",
json={
"base_url": "https://model.test/v1",
"model": "bot",
"description_model": "description",
"api_key": "shared-secret",
},
)
calls = []
def answer(messages, info):
calls.append(messages)
if corrected and len(calls) == 2:
value = {"regular": FIELDS}
elif failure == "plain_text":
return ModelResponse(parts=[TextPart("shared-secret private model output")])
elif failure == "sections":
value = {"combo": FIELDS}
elif failure == "paragraph_breaks":
value = {"regular": {**FIELDS, "idea": FIELDS["idea"] + "\n\nExtra paragraph."}}
else:
value = {"regular": {**FIELDS, "idea": "shared-secret" * 50}}
return ModelResponse(parts=[ToolCallPart(info.output_tools[0].name, {"descriptions": value})])
@asynccontextmanager
async def model_factory(config, settings):
yield FunctionModel(function=answer)
app.state.ai.model_factory = model_factory
state = (await logged_in.get("/api/v1/alphas/alpha1/submission")).json()
response = await logged_in.post(
"/api/v1/alphas/alpha1/description/generate", json={"snapshot": state["snapshot"]}
)
assert len(calls) == 2
if corrected:
assert response.status_code == 200, response.text
assert response.json()["descriptions"] == {"regular": Description(**FIELDS).text()}
else:
assert response.status_code == 502
assert "格式" in response.json()["detail"]
assert "UnexpectedModelBehavior" in caplog.text
assert "shared-secret" not in response.text + caplog.text
assert not platform.calls
@pytest.mark.parametrize("protocol", ["chat_completions", "responses"])
async def test_description_structured_output_through_real_provider(app, logged_in, protocol):
from app.ai.provider import model_connection
platform = await setup(app)
await logged_in.put(
"/api/v1/ai/settings",
json={
"base_url": "https://model.test/v1",
"model": "bot",
"description_model": "description-wire-model",
"protocol": protocol,
"api_key": "synthetic-key",
},
)
paths = []
def gateway(request):
paths.append(request.url.path)
body = json.loads(request.content)
assert request.headers["authorization"] == "Bearer synthetic-key"
assert body["model"] == "description-wire-model" and body["stream"] is False
tool = body["tools"][0]
tool = tool["function"] if protocol == "chat_completions" else tool
# Exercise the real SDK's request schema, not just a FunctionModel's output adapter.
schema = json.dumps(tool["parameters"])
assert all(field in schema for field in FIELDS)
arguments = json.dumps({"descriptions": {"regular": FIELDS}})
if protocol == "chat_completions":
return httpx.Response(
200,
json={
"id": "chat-description",
"object": "chat.completion",
"created": 1789000000,
"model": body["model"],
"choices": [
{
"index": 0,
"finish_reason": "tool_calls",
"message": {
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "description-call",
"type": "function",
"function": {"name": tool["name"], "arguments": arguments},
}
],
},
}
],
},
)
return httpx.Response(
200,
json={
"id": "resp-description",
"object": "response",
"created_at": 1789000000,
"model": body["model"],
"status": "completed",
"output": [
{
"id": "fc-description",
"call_id": "description-call",
"type": "function_call",
"name": tool["name"],
"arguments": arguments,
"status": "completed",
}
],
},
)
@asynccontextmanager
async def model_factory(config, settings):
async with model_connection(config, settings, httpx.MockTransport(gateway)) as model:
yield model
app.state.ai.model_factory = model_factory
state = (await logged_in.get("/api/v1/alphas/alpha1/submission")).json()
response = await logged_in.post(
"/api/v1/alphas/alpha1/description/generate", json={"snapshot": state["snapshot"]}
)
assert response.status_code == 200, response.text
assert response.json()["descriptions"] == {"regular": Description(**FIELDS).text()}
assert paths == ["/v1/" + ("chat/completions" if protocol == "chat_completions" else "responses")]
assert not platform.calls
async def test_invalid_snapshot_and_section_rejected(app, logged_in):