Files
worldquant-alpha-system/backend/tests/test_submission.py
T
yuxuanhui a39894a9f1
Deploy production / deploy (push) Successful in 56s
fix: 修复 Description 生成格式并按需轮询后台状态
2026-09-10 15:32:59 +08:00

571 lines
23 KiB
Python

import copy
import json
from contextlib import asynccontextmanager
import httpx
import pytest
from pydantic import ValidationError
from pydantic_ai.messages import ModelResponse, ToolCallPart
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.worldquant import WqClient
from tests.conftest import alpha
FIELDS = {
"idea": "Short term price reversal is a hypothesis for this signal.",
"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.",
}
CHECKS = [{"name": "PROD_CORRELATION", "result": "FAIL", "value": 0.8, "limit": 0.7}]
class Platform:
def __init__(self, raw):
self.raw = copy.deepcopy(raw)
self.calls = []
self.patches = []
self.pending = 0
self.fail_patch = False
self.unknown_patch = False
self.fail_check = False
def __call__(self, request):
self.calls.append((request.method, request.url.path))
if request.url.path == "/authentication":
return httpx.Response(201, json={})
if request.method == "PATCH" and request.url.path == "/alphas/alpha1":
self.patches.append(json.loads(request.content))
if self.fail_patch:
return httpx.Response(400, json={"secret": "not-for-client"})
for section, value in self.patches[-1].items():
self.raw[section].update(value)
if self.unknown_patch:
self.unknown_patch = False
raise httpx.ReadTimeout("secret", request=request)
return httpx.Response(204)
if request.method == "GET" and request.url.path == "/alphas/alpha1":
return httpx.Response(200, json=self.raw)
if request.method == "GET" and request.url.path == "/alphas/alpha1/check":
if self.fail_check:
return httpx.Response(403, json={})
if self.pending:
self.pending -= 1
return httpx.Response(202, headers={"Retry-After": "0"}, json={})
return httpx.Response(200, json={"is": {"checks": CHECKS}})
raise AssertionError(f"Unexpected platform operation {request.method} {request.url.path}")
async def setup(app, raw=None):
raw = raw or alpha()
async with app.state.sessions.begin() as db:
await upsert_alpha(db, raw)
account = await db.get(Account, 1)
account.email = "test@example.com"
account.password_encrypted = cipher(app.state.settings).encrypt(b"fake-password").decode()
account.connection_status = "connected"
db.add(SelfCorrelation(alpha_id="alpha1", region="USA", result={"status": "low"}, stale=False))
(await db.get(Research, "alpha1")).note = "preserve local research"
platform = Platform(raw)
await app.state.runner.client.close()
app.state.runner.client = WqClient(app.state.settings, transport=httpx.MockTransport(platform))
return platform
async def enqueue(client, descriptions=None):
state = (await client.get("/api/v1/alphas/alpha1/submission")).json()
return await client.post(
"/api/v1/alphas/alpha1/submission-check",
json={
"snapshot": state["snapshot"],
"descriptions": descriptions or {"regular": FIELDS},
},
)
@pytest.mark.parametrize("kind", ["REGULAR", "SUPER"])
async def test_write_check_poll_and_preserve_snapshot(app, logged_in, kind):
raw = (
alpha()
if kind == "REGULAR"
else alpha(type="SUPER", selection={"code": "rank(close)"}, combo={"code": "alpha"})
)
platform = await setup(app, raw)
platform.pending = 2
descriptions = {key: FIELDS for key in (["regular"] if kind == "REGULAR" else ["selection", "combo"])}
response = await enqueue(logged_in, descriptions)
assert response.status_code == 202, response.text
job_id = response.json()["id"]
duplicate = await enqueue(logged_in, descriptions)
assert duplicate.json()["id"] == job_id
await app.state.runner.execute(job_id)
async with app.state.sessions() as db:
job = await db.get(Job, job_id)
assert job.status == "completed", job.error
assert job.processed == 1 and job.checkpoint["phase"] == "checked"
item = await db.get(Alpha, "alpha1")
assert item.check_type == "FAIL_1" and item.prod_correlation == 0.8
assert item.checks == CHECKS and item.is_metrics["checks"] == CHECKS
assert item.expression == raw["regular"]["code"] and item.sharpe == 1.5
assert item.raw["settings"] == raw["settings"]
assert (await db.get(Research, "alpha1")).note == "preserve local research"
assert not (await db.get(SelfCorrelation, "alpha1")).stale
assert platform.patches == [{key: {"description": Description(**FIELDS).text()} for key in descriptions}]
assert platform.calls.count(("GET", "/alphas/alpha1/check")) == 3
assert not any(path.endswith("/submit") for _, path in platform.calls)
# Durable completion does not repeat either operation.
await app.state.runner.execute(job_id)
assert len(platform.patches) == 1
state = (await logged_in.get("/api/v1/alphas/alpha1/submission")).json()
assert state["job"]["status"] == "completed"
@pytest.mark.parametrize(
"status,stale", [("high", False), ("partial", False), ("insufficient_data", False), ("low", True)]
)
async def test_local_correlation_blocks_upstream(app, logged_in, status, stale):
platform = await setup(app)
async with app.state.sessions.begin() as db:
row = await db.get(SelfCorrelation, "alpha1")
row.result, row.stale = {"status": status}, stale
response = await enqueue(logged_in)
assert response.status_code == 409
assert platform.calls == []
@pytest.mark.parametrize("change", ["description", "code", "settings", "status"])
async def test_remote_conflict_never_overwrites(app, logged_in, change):
platform = await setup(app)
response = await enqueue(logged_in)
if change in ("description", "code"):
platform.raw["regular"][change] = "changed by another user"
elif change == "settings":
platform.raw["settings"]["delay"] = 0
else:
platform.raw["status"] = "ACTIVE"
await app.state.runner.execute(response.json()["id"])
async with app.state.sessions() as db:
assert (await db.get(Job, response.json()["id"])).status == "failed"
assert platform.patches == []
assert ("GET", "/alphas/alpha1/check") not in platform.calls
@pytest.mark.parametrize("failure", ["fail_patch", "unknown_patch", "fail_check"])
async def test_retry_reconciles_partial_writes(app, logged_in, failure):
platform = await setup(app)
setattr(platform, failure, True)
response = await enqueue(logged_in)
job_id = response.json()["id"]
await app.state.runner.execute(job_id)
async with app.state.sessions() as db:
assert (await db.get(Job, job_id)).status == "failed"
item = await db.get(Alpha, "alpha1")
assert item.checks != CHECKS
if failure == "fail_check":
assert item.raw["regular"]["description"] == Description(**FIELDS).text()
if failure == "fail_patch":
assert ("GET", "/alphas/alpha1/check") not in platform.calls
setattr(platform, failure, False)
await app.state.runner.execute(job_id)
async with app.state.sessions() as db:
assert (await db.get(Job, job_id)).status == "completed"
assert len(platform.patches) == (2 if failure == "fail_patch" else 1)
async def test_complete_description_reused_verbatim(app, logged_in):
original = Description(**FIELDS).text().replace("\n", "\n\n")
platform = await setup(app, alpha(regular={"code": "rank(close)", "description": original}))
response = await enqueue(logged_in)
await app.state.runner.execute(response.json()["id"])
assert not platform.patches
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
):
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
async def model_factory(config, settings):
seen.append(
(
config.base_url,
config.model,
config.protocol,
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)}),
]
)
)
app.state.ai.model_factory = model_factory
response = await logged_in.put(
"/api/v1/ai/settings",
json={
"base_url": "https://model.test/v1",
"model": "bot-model",
"description_model": "description-model",
"protocol": "responses",
"api_key": "shared-secret",
},
)
assert response.status_code == 200 and "shared-secret" not in response.text
state = (await logged_in.get("/api/v1/alphas/alpha1/submission")).json()
generated = await logged_in.post(
"/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 seen == [("https://model.test/v1", "description-model", "responses", "shared-secret")]
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)
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
await logged_in.put(
"/api/v1/ai/settings",
json={
"base_url": "https://model.test/v1",
"model": "bot",
"description_model": "description",
"api_key": "shared-secret",
},
)
@asynccontextmanager
async def broken(config, settings):
yield FunctionModel(
function=lambda messages, info: ModelResponse(
parts=[
ToolCallPart(
info.output_tools[0].name, {"descriptions": {"regular": {**FIELDS, "idea": " "}}}
),
]
)
)
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):
platform = await setup(app)
response = await logged_in.post(
"/api/v1/alphas/alpha1/submission-check",
json={
"snapshot": "0" * 64,
"descriptions": {"regular": FIELDS},
},
)
assert response.status_code == 409
assert (await enqueue(logged_in, {"combo": FIELDS})).status_code == 422
assert platform.calls == []
async def test_description_model_does_not_invalidate_bot_test(app, logged_in):
from tests.test_ai import CONFIG, configure
await configure(app, logged_in)
response = await logged_in.put(
"/api/v1/ai/settings",
json={
**{k: v for k, v in CONFIG.items() if k != "api_key"},
"enabled": True,
"description_model": " separate-description-model ",
},
)
assert response.json()["ready"] and response.json()["enabled"]
assert response.json()["description_model"] == "separate-description-model"
# Older clients saving bot settings do not clear the independently configured model.
response = await logged_in.put(
"/api/v1/ai/settings",
json={
**{k: v for k, v in CONFIG.items() if k != "api_key"},
"enabled": True,
},
)
assert response.json()["description_model"] == "separate-description-model"
def test_description_model_migration_preserves_existing_config(tmp_path):
import importlib.util
from pathlib import Path
import sqlalchemy as sa
from alembic.migration import MigrationContext
from alembic.operations import Operations
path = Path(__file__).parents[1] / "migrations/versions/0012_description_model.py"
spec = importlib.util.spec_from_file_location("description_migration", path)
migration = importlib.util.module_from_spec(spec)
spec.loader.exec_module(migration)
engine = sa.create_engine(f"sqlite:///{tmp_path}/migration.db")
with engine.begin() as connection:
connection.exec_driver_sql("CREATE TABLE ai_settings (id INTEGER PRIMARY KEY, model VARCHAR(200))")
connection.exec_driver_sql("INSERT INTO ai_settings VALUES (1, 'keep-bot-model')")
with Operations.context(MigrationContext.configure(connection)):
migration.upgrade()
assert connection.exec_driver_sql("SELECT model, description_model FROM ai_settings").one() == (
"keep-bot-model",
"",
)
migration.downgrade()
assert connection.exec_driver_sql("SELECT model FROM ai_settings").scalar() == "keep-bot-model"
assert "description_model" not in {
c["name"] for c in sa.inspect(connection).get_columns("ai_settings")
}
engine.dispose()
@pytest.mark.parametrize("cached", [False, True])
async def test_no_local_references_does_not_block_platform_check(app, logged_in, cached):
platform = await setup(app)
async with app.state.sessions.begin() as db:
result = await db.get(SelfCorrelation, "alpha1")
if cached:
result.result = {"status": "insufficient_data", "candidate_count": 0}
else:
await db.delete(result)
state = (await logged_in.get("/api/v1/alphas/alpha1/submission")).json()
assert state["can_check"] is True
response = await enqueue(logged_in)
assert response.status_code == 202, response.text
await app.state.runner.execute(response.json()["id"])
async with app.state.sessions() as db:
job = await db.get(Job, response.json()["id"])
assert job.status == "completed", job.error
assert ("GET", "/alphas/alpha1/check") in platform.calls
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:
(await db.get(SelfCorrelation, "alpha1")).result = {
"status": "insufficient_data",
"candidate_count": 0,
}
response = await enqueue(logged_in)
assert response.status_code == 202
async with app.state.sessions.begin() as db:
await upsert_alpha(db, alpha(id="peer", status="ACTIVE"))
assert not (await logged_in.get("/api/v1/alphas/alpha1/submission")).json()["can_check"]
assert (await enqueue(logged_in)).status_code == 409
await app.state.runner.execute(response.json()["id"])
async with app.state.sessions() as db:
assert (await db.get(Job, response.json()["id"])).status == "failed"
assert not platform.patches
assert ("GET", "/alphas/alpha1/check") not in platform.calls
async def test_complete_text_input_is_written_as_three_paragraphs(app, logged_in):
platform = await setup(app)
text = Description(**FIELDS).text()
response = await enqueue(logged_in, {"regular": text})
assert response.status_code == 202
await app.state.runner.execute(response.json()["id"])
assert platform.patches == [{"regular": {"description": text}}]
assert len(text.split("\n\n")) == 3
state = (await logged_in.get("/api/v1/alphas/alpha1/submission")).json()
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
with pytest.raises(ValidationError):
GeneratedDescriptions(descriptions={"regular": bad})