fix: 按提示词生成 Description 并保留原文
Deploy production / deploy (push) Successful in 56s

取消生成和写回时的内容、label 与长度校验,避免正常模型文本被本地 500 字符上限拒绝。模型直接返回普通文本,用户可在 textarea 中编辑后原样写回。

验证:62 项后端相关测试、前端构建与浏览器生成编辑写回检查通过;使用用户提供的表达式和 Chat Completions 服务实测一次生成 686 字符成功。
This commit is contained in:
yuxuanhui
2026-09-10 16:07:18 +08:00
parent a39894a9f1
commit 49bf8de9c8
3 changed files with 127 additions and 252 deletions
+81 -127
View File
@@ -4,14 +4,13 @@ from contextlib import asynccontextmanager
import httpx
import pytest
from pydantic import ValidationError
from pydantic_ai.messages import ModelResponse, ToolCallPart
from pydantic_ai.messages import ModelResponse, TextPart, UserPromptPart
from pydantic_ai.models.function import FunctionModel
from app.alphas import upsert_alpha
from app.models import Account, AISettings, Alpha, Job, Research, SelfCorrelation
from app.security import cipher
from app.submission import Description
from app.submission import HEADINGS, Description
from app.worldquant import WqClient
from tests.conftest import alpha
@@ -20,6 +19,12 @@ FIELDS = {
"data_rationale": "Close prices represent the observed price history of each instrument.",
"operator_rationale": "The delta measures five day change, rank compares stocks and negation reverses the ordering.",
}
# Captured from the failing gateway response: 649 characters including template labels.
ACCRUAL_FIELDS = {
"idea": "Combines ranked accrual-related signals and subtracts a ranked score signal, hypothesizing that higher accrual measures and lower fscore values identify relatively attractive stocks.",
"data_rationale": "The selected fields are named cashflow accruals, WC accruals, and fscore; their exact definitions are not provided, so interpretation is qualified. A 120-period backfill supplies missing observations.",
"operator_rationale": "ts_backfill extends the latest available value, rank makes fields cross-sectionally comparable, addition combines the accrual signals, and subtraction gives the fscore component a negative contribution.",
}
CHECKS = [{"name": "PROD_CORRELATION", "result": "FAIL", "value": 0.8, "limit": 0.7}]
@@ -184,17 +189,18 @@ async def test_complete_description_reused_verbatim(app, logged_in):
assert ("GET", "/alphas/alpha1/check") in platform.calls
@pytest.mark.parametrize("bad", ["", " ", "\n"])
def test_description_nonempty(bad):
with pytest.raises(ValidationError):
Description(**{**FIELDS, "idea": bad})
@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
):
@pytest.mark.parametrize(
"text",
[
"Brief draft",
"\n".join(HEADINGS),
Description(**ACCRUAL_FIELDS).text(),
" An idea.\n\nMore detail. ",
],
ids=["short_plain_text", "empty_template", "recorded_over_500", "whitespace"],
)
async def test_generated_text_is_returned_verbatim_without_content_validation(app, logged_in, kind, text):
raw = (
alpha()
if kind == "REGULAR"
@@ -202,10 +208,15 @@ async def test_ai_uses_independent_model_shared_connection_without_platform_writ
)
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 = []
prompts = []
def answer(messages, info):
assert not info.output_tools and not info.function_tools
prompts.append(
json.loads(next(p.content for m in messages for p in m.parts if isinstance(p, UserPromptPart)))
)
return ModelResponse(parts=[TextPart(text + prompts[-1]["section"])])
@asynccontextmanager
async def model_factory(config, settings):
@@ -217,13 +228,7 @@ async def test_ai_uses_independent_model_shared_connection_without_platform_writ
cipher(settings).decrypt(config.api_key_encrypted.encode()).decode(),
)
)
yield FunctionModel(
function=lambda messages, info: ModelResponse(
parts=[
ToolCallPart(info.output_tools[0].name, {"descriptions": dict.fromkeys(sections, value)}),
]
)
)
yield FunctionModel(function=answer)
app.state.ai.model_factory = model_factory
response = await logged_in.put(
@@ -242,16 +247,18 @@ 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"] == dict.fromkeys(sections, Description(**FIELDS).text())
assert generated.json()["descriptions"] == {key: text + key for key in sections}
assert seen == [("https://model.test/v1", "description-model", "responses", "shared-secret")]
assert [prompt["section"] for prompt in prompts] == sections
assert [prompt["code"] for prompt in prompts] == [raw[key]["code"] for key in sections]
assert not platform.calls
async with app.state.sessions() as db:
assert (await db.get(AISettings, 1)).model == "bot-model"
assert "description" not in (await db.get(Alpha, "alpha1")).raw["regular"]
async def test_generation_requires_config_and_rejects_invalid_output(app, logged_in):
await setup(app)
async def test_generation_requires_config_and_safely_reports_provider_failure(app, logged_in, caplog):
platform = await setup(app)
state = (await logged_in.get("/api/v1/alphas/alpha1/submission")).json()
path = "/api/v1/alphas/alpha1/description/generate"
assert (await logged_in.post(path, json={"snapshot": state["snapshot"]})).status_code == 409
@@ -265,80 +272,26 @@ async def test_generation_requires_config_and_rejects_invalid_output(app, logged
},
)
def fail(messages, info):
raise httpx.ReadTimeout("shared-secret private provider response")
@asynccontextmanager
async def broken(config, settings):
yield FunctionModel(
function=lambda messages, info: ModelResponse(
parts=[
ToolCallPart(
info.output_tools[0].name, {"descriptions": {"regular": {**FIELDS, "idea": " "}}}
),
]
)
)
yield FunctionModel(function=fail)
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 response.status_code == 502
assert "超时" in response.json()["detail"]
assert "ReadTimeout" in caplog.text
assert "shared-secret" not in response.text + caplog.text
assert "private provider response" 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):
@pytest.mark.parametrize("returned_fields", [FIELDS, ACCRUAL_FIELDS], ids=["short", "recorded_over_500"])
async def test_description_text_output_through_real_provider(app, logged_in, protocol, returned_fields):
from app.ai.provider import model_connection
platform = await setup(app)
@@ -359,12 +312,8 @@ async def test_description_structured_output_through_real_provider(app, logged_i
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}})
assert not body.get("tools") and not body.get("response_format")
text = Description(**returned_fields).text()
if protocol == "chat_completions":
return httpx.Response(
200,
@@ -376,17 +325,10 @@ async def test_description_structured_output_through_real_provider(app, logged_i
"choices": [
{
"index": 0,
"finish_reason": "tool_calls",
"finish_reason": "stop",
"message": {
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "description-call",
"type": "function",
"function": {"name": tool["name"], "arguments": arguments},
}
],
"content": text,
},
}
],
@@ -402,11 +344,10 @@ async def test_description_structured_output_through_real_provider(app, logged_i
"status": "completed",
"output": [
{
"id": "fc-description",
"call_id": "description-call",
"type": "function_call",
"name": tool["name"],
"arguments": arguments,
"id": "msg-description",
"type": "message",
"role": "assistant",
"content": [{"type": "output_text", "text": text, "annotations": []}],
"status": "completed",
}
],
@@ -424,11 +365,22 @@ async def test_description_structured_output_through_real_provider(app, logged_i
"/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 response.json()["descriptions"] == {"regular": Description(**returned_fields).text()}
assert paths == ["/v1/" + ("chat/completions" if protocol == "chat_completions" else "responses")]
assert not platform.calls
@pytest.mark.parametrize("format", ["fields", "complete_text"])
async def test_manual_description_over_500_is_written_completely(app, logged_in, format):
platform = await setup(app)
fields = {**FIELDS, "idea": "x" * 600}
text = "\n\n".join(heading + value for heading, value in zip(HEADINGS, fields.values()))
response = await enqueue(logged_in, {"regular": fields if format == "fields" else text})
assert response.status_code == 202, response.text
await app.state.runner.execute(response.json()["id"])
assert platform.patches == [{"regular": {"description": text}}]
async def test_invalid_snapshot_and_section_rejected(app, logged_in):
platform = await setup(app)
response = await logged_in.post(
@@ -519,17 +471,6 @@ async def test_no_local_references_does_not_block_platform_check(app, logged_in,
assert not any(path.endswith("/submit") for _, path in platform.calls)
@pytest.mark.parametrize("length", [499, 500, 501])
def test_description_total_character_limit(length):
fields = {**FIELDS, "idea": "x"}
fields["idea"] += "x" * (length - len(Description(**fields).text()))
if length > 500:
with pytest.raises(ValidationError):
Description(**fields)
else:
assert len(Description(**fields).text()) == length
async def test_new_reference_blocks_empty_cached_result_at_enqueue_and_execution(app, logged_in):
platform = await setup(app)
async with app.state.sessions.begin() as db:
@@ -562,9 +503,22 @@ async def test_complete_text_input_is_written_as_three_paragraphs(app, logged_in
assert state["descriptions"] == {"regular": text}
@pytest.mark.parametrize("bad", ["x" * 501, "x" * 100, Description(**FIELDS).text().replace("\n\n", "\n")])
def test_generation_rejects_incomplete_or_oversized_text(bad):
from app.submission import GeneratedDescriptions
@pytest.mark.parametrize("heading,field", list(zip(HEADINGS, FIELDS)))
@pytest.mark.parametrize("failure", ["missing_label", "empty_content"])
async def test_incomplete_template_can_be_reviewed_and_written(app, logged_in, heading, field, failure):
platform = await setup(app)
text = Description(**FIELDS).text()
bad = text.replace(heading, "") if failure == "missing_label" else text.replace(FIELDS[field], " \n ")
response = await enqueue(logged_in, {"regular": bad})
assert response.status_code == 202, response.text
await app.state.runner.execute(response.json()["id"])
assert platform.patches == [{"regular": {"description": bad}}]
with pytest.raises(ValidationError):
GeneratedDescriptions(descriptions={"regular": bad})
@pytest.mark.parametrize("text", ["", "short", "x" * 600, " Idea: draft\nMore details. "])
async def test_reviewed_plain_text_is_written_verbatim(app, logged_in, text):
platform = await setup(app)
response = await enqueue(logged_in, {"regular": text})
assert response.status_code == 202, response.text
await app.state.runner.execute(response.json()["id"])
assert platform.patches == ([{"regular": {"description": text}}] if text else [])