This commit is contained in:
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user