feat: add cached homepage information and manual AI interpretations
This commit is contained in:
@@ -0,0 +1,300 @@
|
||||
"""Observed platform schemas and homepage cache/AI boundaries."""
|
||||
import asyncio
|
||||
import importlib.util
|
||||
from contextlib import asynccontextmanager
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
import sqlalchemy as sa
|
||||
from alembic.migration import MigrationContext
|
||||
from alembic.operations import Operations
|
||||
from pydantic_ai.messages import ModelResponse, ToolCallPart
|
||||
from pydantic_ai.models.function import FunctionModel
|
||||
|
||||
from app import home_information as home
|
||||
from app.models import Account, AISettings, HomeInformation, HomeRankHistory
|
||||
from app.worldquant import WqError
|
||||
|
||||
PREFIX = "/api/v1/dashboard/information"
|
||||
QUERY = "?competition_id=ARC2026"
|
||||
|
||||
|
||||
def competition():
|
||||
return {"id": "ARC2026", "name": "All Region Competition", "description": "<b>Research</b>",
|
||||
"startDate": "2026-09-01T00:00:00-04:00", "endDate": "2026-10-11T23:59:59-04:00",
|
||||
"status": "ACCEPTED", "leaderboard": {"user": "TEST_USER", "rank": 7915, "alphas": 0}}
|
||||
|
||||
|
||||
def page(items, next=None):
|
||||
return {"results": items, "next": next, "count": len(items)}
|
||||
|
||||
|
||||
async def setup(app, monkeypatch):
|
||||
async with app.state.sessions.begin() as db:
|
||||
account = await db.get(Account, 1)
|
||||
account.connection_status, account.wq_user_id = "connected", "TEST_USER"
|
||||
calls, state = [], {"rank": 25, "fail": set(), "agreement": "Participants must use GLOBAL region. Delay must be 1."}
|
||||
|
||||
async def get(path, params=None, headers=None):
|
||||
calls.append((path, params))
|
||||
if path in state["fail"]:
|
||||
raise WqError("资源暂不可用")
|
||||
if path == "/users/self/messages":
|
||||
offset = params["offset"]
|
||||
return page([{"id": f"msg{offset}", "title": "Update", "description": "<p>News</p><script>secret()</script><a href='javascript:alert(1)'>bad</a>", "dateCreated": "2026-09-11T12:00:00-04:00"}],
|
||||
"https://api.worldquantbrain.com/users/self/messages?offset=1" if offset == 0 else None)
|
||||
if path == "/consultant/boards/leader":
|
||||
assert params == {"user": "TEST_USER"}
|
||||
return page([{"user": "TEST_USER", "dailyOsmosisRank": state["rank"], "valueFactor": 0.5}])
|
||||
if path == "/users/TEST_USER/competitions":
|
||||
return page([competition()])
|
||||
if path == "/events":
|
||||
return page([{"id": "event", "title": "Webinar", "start": "2026-09-20T09:00:00-04:00", "end": "2026-09-20T10:00:00-04:00", "timezone": "US/Eastern"}])
|
||||
if path == "/competitions/ARC2026":
|
||||
return competition()
|
||||
if path.endswith("/agreement"):
|
||||
return {"title": "Rules", "lastModified": "2026-09-07T04:16:26-04:00", "content": [{"type": "TEXT", "value": state["agreement"]}]}
|
||||
raise AssertionError(path)
|
||||
monkeypatch.setattr(app.state.runner.client, "get", get)
|
||||
return calls, state
|
||||
|
||||
|
||||
async def test_cache_pagination_failure_and_account_isolation(app, logged_in, monkeypatch):
|
||||
calls, state = await setup(app, monkeypatch)
|
||||
first = (await logged_in.get(PREFIX + "/messages")).json()
|
||||
assert first["content"]["next_offset"] == 1
|
||||
assert "secret()" not in str(first) and "javascript:" not in str(first)
|
||||
assert first["analysis"] is None and not first["can_generate"]
|
||||
assert (await logged_in.get(PREFIX + "/messages")).json() == first and len(calls) == 1
|
||||
second = (await logged_in.get(PREFIX + "/messages?offset=1")).json()
|
||||
assert second["content"]["items"][0]["id"] == "msg1"
|
||||
state["fail"].add("/users/self/messages")
|
||||
failed = (await logged_in.post(PREFIX + "/messages/refresh")).json()
|
||||
assert failed["stale"] and failed["error"]
|
||||
assert failed["content"] == first["content"] and failed["fetched_at"] == first["fetched_at"]
|
||||
assert (await logged_in.get(PREFIX + "/events")).json()["content"]["items"]
|
||||
async with app.state.sessions.begin() as db:
|
||||
(await db.get(Account, 1)).wq_user_id = "OTHER"
|
||||
other = (await logged_in.get(PREFIX + "/messages")).json()
|
||||
assert other["content"] is None and other["fetched_at"] is None
|
||||
|
||||
|
||||
async def test_auth_and_invalid_competition(app, client, logged_in, monkeypatch):
|
||||
calls, _ = await setup(app, monkeypatch)
|
||||
assert (await logged_in.get(PREFIX + "/competition?competition_id=../secrets")).status_code == 422
|
||||
assert (await logged_in.get(PREFIX + "/messages?offset=-1")).status_code == 422
|
||||
async with app.state.sessions.begin() as db:
|
||||
(await db.get(Account, 1)).connection_status = "disconnected"
|
||||
assert (await logged_in.get(PREFIX + "/events")).status_code == 409
|
||||
client.cookies.clear()
|
||||
assert (await client.post(PREFIX + "/events/refresh")).status_code == 401
|
||||
assert not calls
|
||||
|
||||
|
||||
async def test_rank_first_capture_zero_and_daily_comparability(app, logged_in, monkeypatch):
|
||||
_, state = await setup(app, monkeypatch)
|
||||
monkeypatch.setattr(home, "now", lambda: datetime.fromisoformat("2026-09-12T03:59:00+00:00"))
|
||||
first = (await logged_in.get(PREFIX + "/leaderboard")).json()["content"]
|
||||
assert first["rank"] == 25 and first["change"] is None
|
||||
state["rank"] = 20
|
||||
assert (await logged_in.post(PREFIX + "/leaderboard/refresh")).json()["content"]["change"] == 5
|
||||
monkeypatch.setattr(home, "now", lambda: datetime.fromisoformat("2026-09-12T04:00:00+00:00"))
|
||||
third = (await logged_in.post(PREFIX + "/leaderboard/refresh")).json()["content"]
|
||||
assert third["change"] is None and third["scope"] != first["scope"]
|
||||
state["rank"] = 0.0
|
||||
fourth = (await logged_in.post(PREFIX + "/leaderboard/refresh")).json()["content"]
|
||||
assert fourth["rank"] is None and fourth["change"] is None
|
||||
async with app.state.sessions() as db:
|
||||
assert len((await db.scalars(sa.select(HomeRankHistory))).all()) == 3
|
||||
|
||||
|
||||
def test_time_filter_sort_and_missing_timezone(monkeypatch):
|
||||
monkeypatch.setattr(home, "now", lambda: datetime.fromisoformat("2026-09-12T04:00:00+00:00"))
|
||||
content = {"items": [{"id": "old", "end": "2026-09-12T00:00:00-04:00"},
|
||||
{"id": "unknown", "end": "2026-09-20"},
|
||||
{"id": "future", "start": "2026-09-13T00:00:00-04:00", "end": "2026-09-13T01:00:00-04:00"},
|
||||
{"id": "ongoing", "start": "2026-09-11T00:00:00-04:00", "end": "2026-09-12T01:00:00-04:00"}]}
|
||||
assert [i["id"] for i in home.temporal(content, "events")["items"]] == ["ongoing", "future", "unknown"]
|
||||
assert [i["id"] for i in home.temporal(content, "competitions")["items"]] == ["ongoing", "future", "unknown", "old"]
|
||||
assert home.instant("2026-09-20") is None
|
||||
|
||||
|
||||
async def test_all_event_pages_and_untrusted_next_urls():
|
||||
class Client:
|
||||
async def get(self, path, params):
|
||||
offset = params["offset"]
|
||||
return page([{"id": str(offset)}], f"https://api.worldquantbrain.com/events?offset={offset+1}" if offset < 2 else None)
|
||||
assert len(await home.all_pages(Client(), "/events")) == 3
|
||||
for link in ["https://evil.test/events?offset=1", "https://api.worldquantbrain.com/other?offset=1", "https://api.worldquantbrain.com/events?offset=0"]:
|
||||
with pytest.raises(WqError):
|
||||
home.next_offset({"next": link}, "/events", 0)
|
||||
with pytest.raises(WqError):
|
||||
home.rows({"results": "bad"})
|
||||
|
||||
|
||||
async def model_setup(app, wrong_quote=False):
|
||||
calls = []
|
||||
def respond(messages, info):
|
||||
items = [{"title": title, "text": "未知", "evidence": []} for title in home.RULES]
|
||||
items[0].update(text="GLOBAL", evidence=[{"source_id": "agreement-0", "quote": "invented" if wrong_quote else "GLOBAL region"}])
|
||||
return ModelResponse(parts=[ToolCallPart(info.output_tools[0].name, {"items": items, "suggestions": ["请核对平台协议中的截止日期。"]})])
|
||||
@asynccontextmanager
|
||||
async def factory(config, settings):
|
||||
calls.append(config.model)
|
||||
yield FunctionModel(respond)
|
||||
app.state.ai.model_factory = factory
|
||||
async with app.state.sessions.begin() as db:
|
||||
config = await db.get(AISettings, 1)
|
||||
config.model = "research-model"
|
||||
config.description_model, config.base_url, config.api_key_encrypted = "basic-model", "https://model.test/v1", "test-key-encrypted"
|
||||
return calls
|
||||
|
||||
|
||||
async def test_manual_ai_cache_source_config_invalidation_and_no_fallback(app, logged_in, monkeypatch):
|
||||
_, state = await setup(app, monkeypatch)
|
||||
calls = await model_setup(app)
|
||||
url = PREFIX + "/competition"
|
||||
await logged_in.get(url + QUERY)
|
||||
await logged_in.get(url + QUERY)
|
||||
assert not calls
|
||||
generated = await logged_in.post(url + "/generate" + QUERY)
|
||||
assert generated.status_code == 200, generated.text
|
||||
saved = generated.json()["analysis"]
|
||||
assert saved["model"] == "basic-model" and not saved["outdated"] and calls == ["basic-model"]
|
||||
assert "test-key-encrypted" not in generated.text
|
||||
state["agreement"] += " Agreement updated."
|
||||
refreshed = (await logged_in.post(url + "/refresh" + QUERY)).json()
|
||||
assert refreshed["analysis"]["outdated"] and refreshed["analysis"]["source_version"] == saved["source_version"]
|
||||
await logged_in.post(url + "/generate" + QUERY)
|
||||
async with app.state.sessions.begin() as db:
|
||||
(await db.get(AISettings, 1)).model = "different-research-model"
|
||||
assert not (await logged_in.get(url + QUERY)).json()["analysis"]["outdated"]
|
||||
async with app.state.sessions.begin() as db:
|
||||
(await db.get(AISettings, 1)).description_model = "new-basic-model"
|
||||
assert (await logged_in.get(url + QUERY)).json()["analysis"]["outdated"]
|
||||
async with app.state.sessions.begin() as db:
|
||||
(await db.get(AISettings, 1)).description_model = ""
|
||||
assert (await logged_in.post(url + "/generate" + QUERY)).status_code == 409
|
||||
assert calls == ["basic-model", "basic-model"]
|
||||
|
||||
|
||||
async def test_failed_grounding_preserves_prior_analysis(app, logged_in, monkeypatch):
|
||||
await setup(app, monkeypatch)
|
||||
await model_setup(app)
|
||||
await logged_in.get(PREFIX + "/competition" + QUERY)
|
||||
saved = (await logged_in.post(PREFIX + "/competition/generate" + QUERY)).json()["analysis"]
|
||||
await model_setup(app, wrong_quote=True)
|
||||
assert (await logged_in.post(PREFIX + "/competition/generate" + QUERY)).status_code == 502
|
||||
assert (await logged_in.get(PREFIX + "/competition" + QUERY)).json()["analysis"] == saved
|
||||
|
||||
|
||||
async def test_concurrent_first_reads_are_coalesced(app, logged_in, monkeypatch):
|
||||
calls, _ = await setup(app, monkeypatch)
|
||||
results = await asyncio.gather(*[logged_in.get(PREFIX + "/messages") for _ in range(4)])
|
||||
assert all(r.status_code == 200 for r in results) and len(calls) == 1
|
||||
|
||||
|
||||
async def test_event_sources_include_competitions_and_current_day(app, logged_in, monkeypatch):
|
||||
await setup(app, monkeypatch)
|
||||
await logged_in.get(PREFIX + "/events")
|
||||
await logged_in.get(PREFIX + "/competitions")
|
||||
async with app.state.sessions() as db:
|
||||
record = await db.get(HomeInformation, ("TEST_USER", "events:"))
|
||||
request = SimpleNamespace(app=app)
|
||||
sources, initial = await home.analysis_sources(request, "TEST_USER", "events", record)
|
||||
assert [s["id"] for s in sources] == ["platform", "competitions", "today"]
|
||||
monkeypatch.setattr(home, "now", lambda: datetime.fromisoformat("2027-01-01T04:00:00+00:00"))
|
||||
assert initial != (await home.analysis_sources(request, "TEST_USER", "events", record))[1]
|
||||
|
||||
|
||||
def test_migration_keeps_model_fields(tmp_path):
|
||||
path = Path(__file__).parents[1] / "migrations/versions/0016_home_information.py"
|
||||
spec = importlib.util.spec_from_file_location("home_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 conn:
|
||||
conn.exec_driver_sql("CREATE TABLE ai_settings (model TEXT, description_model TEXT)")
|
||||
conn.exec_driver_sql("INSERT INTO ai_settings VALUES ('research', 'basic')")
|
||||
with Operations.context(MigrationContext.configure(conn)):
|
||||
migration.upgrade()
|
||||
assert {"home_information", "home_rank_history"}.issubset(sa.inspect(conn).get_table_names())
|
||||
migration.downgrade()
|
||||
assert conn.exec_driver_sql("SELECT * FROM ai_settings").one() == ("research", "basic")
|
||||
engine.dispose()
|
||||
|
||||
|
||||
def test_event_local_timezone_dst_and_missing_dates():
|
||||
assert home.instant("2026-09-12T00:00:00", "US/Eastern").isoformat() == "2026-09-12T04:00:00+00:00"
|
||||
assert home.instant("2026-11-01T01:30:00", "US/Eastern") is None
|
||||
assert home.instant("2026-03-08T02:30:00", "US/Eastern") is None
|
||||
assert home.instant("2026-09-12T00:00:00", "invalid-zone") is None
|
||||
assert home.instant(None) is None
|
||||
|
||||
|
||||
async def test_source_changes_during_generation_preserve_previous_analysis(app, logged_in, monkeypatch):
|
||||
_, state = await setup(app, monkeypatch)
|
||||
await model_setup(app)
|
||||
await logged_in.get(PREFIX + "/competition" + QUERY)
|
||||
initial = (await logged_in.post(PREFIX + "/competition/generate" + QUERY)).json()["analysis"]
|
||||
original_factory = app.state.ai.model_factory
|
||||
started, release = asyncio.Event(), asyncio.Event()
|
||||
@asynccontextmanager
|
||||
async def slow_factory(config, settings):
|
||||
started.set()
|
||||
await release.wait()
|
||||
async with original_factory(config, settings) as model:
|
||||
yield model
|
||||
app.state.ai.model_factory = slow_factory
|
||||
pending = asyncio.create_task(logged_in.post(PREFIX + "/competition/generate" + QUERY))
|
||||
await asyncio.wait_for(started.wait(), 3)
|
||||
assert (await logged_in.post(PREFIX + "/competition/generate" + QUERY)).status_code == 409
|
||||
state["agreement"] += " Changed terms."
|
||||
await logged_in.post(PREFIX + "/competition/refresh" + QUERY)
|
||||
release.set()
|
||||
response = await pending
|
||||
assert response.status_code == 409
|
||||
saved = (await logged_in.get(PREFIX + "/competition" + QUERY)).json()["analysis"]
|
||||
assert saved["source_version"] == initial["source_version"] and saved["outdated"]
|
||||
|
||||
|
||||
async def test_empty_and_partial_unavailability(app, logged_in, monkeypatch):
|
||||
await setup(app, monkeypatch)
|
||||
async def read(path, params=None, headers=None):
|
||||
if path == '/events':
|
||||
raise WqError('无权访问该平台资源', 'access_denied')
|
||||
return page([])
|
||||
monkeypatch.setattr(app.state.runner.client, 'get', read)
|
||||
for module in ['messages', 'competitions']:
|
||||
result = (await logged_in.get(PREFIX + '/' + module)).json()
|
||||
assert result['content']['items'] == [] and result['fetched_at'] and not result['error']
|
||||
rank = (await logged_in.get(PREFIX + '/leaderboard')).json()
|
||||
assert rank['content']['rank'] is None
|
||||
events = (await logged_in.get(PREFIX + '/events')).json()
|
||||
assert events['error'] and events['content'] is None and events['stale']
|
||||
assert (await logged_in.get('/api/v1/auth/me')).status_code == 200
|
||||
|
||||
|
||||
def test_competition_rank_requires_same_user_and_unknown_metrics():
|
||||
data = competition()
|
||||
normalized = home.competition(data, 'TEST_USER')
|
||||
assert normalized['rank'] == 7915 and normalized['alphas'] == 0
|
||||
assert normalized['progress'] is None and normalized['robustness_score'] is None
|
||||
assert home.competition(data, 'OTHER')['rank'] is None
|
||||
assert home.safe_url('https://[malformed') is None
|
||||
assert home.safe_url('javascript:alert(1)') is None
|
||||
assert home.safe_url('https://user:pass@example.com') is None
|
||||
|
||||
|
||||
@pytest.mark.parametrize('field,value', [('base_url','https://new.test/v1'), ('protocol','responses'), ('api_key_encrypted','changed-key')])
|
||||
async def test_shared_connection_change_invalidates_ai(app, logged_in, monkeypatch, field, value):
|
||||
await setup(app, monkeypatch)
|
||||
calls = await model_setup(app)
|
||||
await logged_in.get(PREFIX + '/competition' + QUERY)
|
||||
assert (await logged_in.post(PREFIX + '/competition/generate' + QUERY)).status_code == 200
|
||||
async with app.state.sessions.begin() as db:
|
||||
setattr(await db.get(AISettings, 1), field, value)
|
||||
assert (await logged_in.get(PREFIX + '/competition' + QUERY)).json()['analysis']['outdated']
|
||||
assert calls == ['basic-model']
|
||||
Reference in New Issue
Block a user