feat: add AI descriptions and platform submission checks
This commit is contained in:
@@ -0,0 +1,338 @@
|
||||
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})
|
||||
|
||||
|
||||
async def test_ai_uses_independent_model_shared_connection_without_platform_write(app, logged_in):
|
||||
platform = await setup(app)
|
||||
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": {"regular": FIELDS}}),
|
||||
]
|
||||
)
|
||||
)
|
||||
|
||||
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"] == {"regular": FIELDS}
|
||||
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
|
||||
|
||||
|
||||
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()
|
||||
Reference in New Issue
Block a user