392 lines
21 KiB
Python
392 lines
21 KiB
Python
"""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"
|
|
monkeypatch.setattr(home, "now", lambda: datetime.fromisoformat("2026-09-12T12:00:00+00:00"))
|
|
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{i}", "title": "Update", "description": "<p>News</p><script>secret()</script><a href='javascript:alert(1)'>bad</a>", "dateCreated": "2026-09-11T12:00:00-04:00"} for i in (range(10) if offset == 0 else [10])],
|
|
"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"] == 10
|
|
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) == 2
|
|
second = (await logged_in.get(PREFIX + "/messages?offset=10")).json()
|
|
assert second["content"]["items"][0]["id"] == "msg10"
|
|
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) == 2
|
|
|
|
|
|
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']
|
|
|
|
|
|
@pytest.mark.parametrize('at,expected', [
|
|
('2026-03-31T12:00:00-04:00', '2026-02-28T12:00:00-05:00'),
|
|
('2024-03-31T12:00:00-04:00', '2024-02-29T12:00:00-05:00'),
|
|
('2026-01-31T12:00:00-05:00', '2025-12-31T12:00:00-05:00'),
|
|
])
|
|
def test_recent_month_uses_eastern_calendar_month(monkeypatch, at, expected):
|
|
monkeypatch.setattr(home, 'now', lambda: datetime.fromisoformat(at))
|
|
assert home.month_start().isoformat() == expected
|
|
|
|
|
|
async def test_recent_messages_scan_pages_filter_unknown_future_and_boundary(monkeypatch):
|
|
monkeypatch.setattr(home, 'now', lambda: datetime.fromisoformat('2026-03-31T12:00:00-04:00'))
|
|
calls = []
|
|
class Client:
|
|
async def get(self, path, params):
|
|
calls.append(params)
|
|
if params['offset'] == 0:
|
|
return page([{'id': 'old', 'dateCreated': '2026-02-28T11:59:59-05:00'},
|
|
{'id': 'unknown'}, {'id': 'future', 'dateCreated': '2026-04-01T00:00:00Z'}],
|
|
'https://api.worldquantbrain.com/users/self/messages?offset=3&limit=100')
|
|
return page([{'id': 'boundary', 'dateCreated': '2026-02-28T12:00:00-05:00'},
|
|
{'id': 'today', 'dateCreated': '2026-03-31T12:00:00-04:00'}])
|
|
data = await home.source(Client(), 'messages', 'TEST_USER', 0, '')
|
|
assert [item['id'] for item in data['items']] == ['today', 'boundary']
|
|
assert calls == [{'offset': 0, 'limit': 100}, {'offset': 3, 'limit': 100}]
|
|
|
|
|
|
async def test_messages_and_ai_never_persist_and_expire_without_auto_generation(app, logged_in, monkeypatch):
|
|
platform_calls, _ = await setup(app, monkeypatch)
|
|
await model_setup(app)
|
|
model_calls = []
|
|
def respond(messages, info):
|
|
return ModelResponse(parts=[ToolCallPart(info.output_tools[0].name, {'items': [], 'suggestions': []})])
|
|
@asynccontextmanager
|
|
async def factory(config, settings):
|
|
model_calls.append(config.model)
|
|
yield FunctionModel(respond)
|
|
app.state.ai.model_factory = factory
|
|
# Legacy rows must never be used; migration will remove them on deployment.
|
|
async with app.state.sessions.begin() as db:
|
|
db.add(HomeInformation(user_id='TEST_USER', resource='messages:0', content={'items': [{'title': 'legacy'}]}))
|
|
first = (await logged_in.get(PREFIX + '/messages')).json()
|
|
assert first['ephemeral'] and first['content']['total'] == 11
|
|
generated = await logged_in.post(PREFIX + '/messages/generate')
|
|
assert generated.status_code == 200, generated.text
|
|
assert generated.json()['analysis']['model'] == 'basic-model'
|
|
await logged_in.get(PREFIX + '/messages')
|
|
await logged_in.get(PREFIX + '/messages?offset=10')
|
|
assert len(platform_calls) == 2 and model_calls == ['basic-model']
|
|
async with app.state.sessions() as db:
|
|
saved = (await db.scalars(sa.select(HomeInformation))).all()
|
|
assert len(saved) == 1 and saved[0].content == {'items': [{'title': 'legacy'}]}
|
|
assert saved[0].analysis is None
|
|
monkeypatch.setattr(home, 'now', lambda: datetime.fromisoformat('2026-09-12T12:15:00+00:00'))
|
|
assert (await logged_in.post(PREFIX + '/messages/generate')).status_code == 409
|
|
reloaded = (await logged_in.get(PREFIX + '/messages')).json()
|
|
assert reloaded['analysis'] is None and len(platform_calls) == 4
|
|
assert model_calls == ['basic-model']
|
|
|
|
|
|
async def test_message_window_rechecks_memory_and_removes_expired_interpretation(app, logged_in, monkeypatch):
|
|
await setup(app, monkeypatch)
|
|
monkeypatch.setattr(home, 'now', lambda: datetime.fromisoformat('2026-09-12T12:00:00+00:00'))
|
|
async def upstream(path, params=None):
|
|
return page([{'id': 'boundary', 'dateCreated': '2026-08-12T08:00:00-04:00'}])
|
|
monkeypatch.setattr(app.state.runner.client, 'get', upstream)
|
|
assert (await logged_in.get(PREFIX + '/messages')).json()['content']['total'] == 1
|
|
app.state.home_message_cache['TEST_USER'].analyses[0] = {'old': 'interpretation'}
|
|
monkeypatch.setattr(home, 'now', lambda: datetime.fromisoformat('2026-09-12T12:00:01+00:00'))
|
|
result = (await logged_in.get(PREFIX + '/messages')).json()
|
|
assert result['content']['total'] == 0 and result['analysis'] is None
|
|
|
|
|
|
def test_message_cleanup_migration_only_removes_message_cache(tmp_path):
|
|
path = Path(__file__).parents[1] / 'migrations/versions/0017_remove_message_cache.py'
|
|
spec = importlib.util.spec_from_file_location('message_cleanup', path)
|
|
migration = importlib.util.module_from_spec(spec)
|
|
spec.loader.exec_module(migration)
|
|
engine = sa.create_engine(f'sqlite:///{tmp_path}/cleanup.db')
|
|
with engine.begin() as conn:
|
|
conn.exec_driver_sql('CREATE TABLE home_information (user_id TEXT, resource TEXT, content TEXT, analysis TEXT)')
|
|
for user, resource in [('A', 'messages:0'), ('B', 'messages:10'), ('A', 'events:'), ('A', 'competition:ARC2026')]:
|
|
conn.exec_driver_sql('INSERT INTO home_information VALUES (?, ?, ?, ?)', (user, resource, 'source', 'ai'))
|
|
with Operations.context(MigrationContext.configure(conn)):
|
|
migration.upgrade()
|
|
migration.upgrade()
|
|
assert conn.exec_driver_sql('SELECT resource FROM home_information ORDER BY resource').scalars().all() == ['competition:ARC2026', 'events:']
|
|
engine.dispose()
|