feat: add cached homepage information and manual AI interpretations

This commit is contained in:
yuxuanhui
2026-09-12 01:55:24 +08:00
parent 394438e753
commit d9fbaa7cf7
19 changed files with 1674 additions and 18 deletions
+300
View File
@@ -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']