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": Description(**FIELDS).text()}} ), ] ) ) 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": 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 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})