feat: Implement Alpha list and metrics enhancements
Deploy production / deploy (push) Successful in 1m12s

- Added new metrics fields: sub_universe_sharpe, robust_universe_sharpe, two_year_sharpe, prod_correlation, pnl, check_type, and neutralization to the Alpha model.
- Updated snapshot_columns function to derive new metrics and check types from platform snapshots.
- Enhanced API to include failed checks and check types in responses.
- Created migration script to backfill existing Alpha records with new metrics and check types.
- Updated frontend components to display new metrics and allow editing of custom tags.
- Improved filtering and sorting capabilities for new metrics in the Alpha list.
- Added tests for new functionality including checks classification and metrics filtering.
This commit is contained in:
yuxuanhui
2026-09-09 19:02:37 +08:00
parent e57b1f7a2e
commit ba0ed9d03f
13 changed files with 678 additions and 22 deletions
+61 -2
View File
@@ -9,6 +9,55 @@ from sqlalchemy import or_, select, update
from .models import Alpha, Research, ResearchTag, SelfCorrelation, now
from .research.provenance import source_alpha_ids
METRIC_FIELDS = (
"sharpe", "fitness", "returns", "turnover", "margin", "drawdown",
"sub_universe_sharpe", "robust_universe_sharpe", "two_year_sharpe", "prod_correlation", "pnl",
)
def failed_checks(checks):
"""Return failed platform check names; local correlation never changes this list."""
return [
check.get("name") if isinstance(check.get("name"), str) else "未命名检查"
for check in checks if isinstance(check, dict) and check.get("result") == "FAIL"
] if isinstance(checks, list) else []
def snapshot_columns(settings, metrics, checks):
"""Derive list fields from a platform snapshot, preserving missing metrics as null.
Only explicit FAIL results count. Empty, malformed and unfinished checks are
pending; all known checks passing without PROD_CORRELATION is only a pre-check.
No submission eligibility or activity eligibility is inferred here.
"""
settings = settings if isinstance(settings, dict) else {}
metrics = metrics if isinstance(metrics, dict) else {}
checks = checks if isinstance(checks, list) else []
valid = [check for check in checks if isinstance(check, dict)]
failures = len(failed_checks(checks))
by_name = {check["name"]: check for check in valid if isinstance(check.get("name"), str)}
if failures:
check_type = "FAIL_1" if failures == 1 else "FAIL_2"
elif not checks or len(valid) != len(checks) or any(check.get("result") != "PASS" for check in valid):
check_type = "PENDING"
else:
check_type = "PASS" if "PROD_CORRELATION" in by_name else "PRE_CHECK"
neutralization = settings.get("neutralization")
return {
"check_type": check_type,
"neutralization": neutralization if isinstance(neutralization, str) else None,
"pnl": number(metrics.get("pnl")),
**{
field: number(by_name.get(name, {}).get("value"))
for field, name in (
("sub_universe_sharpe", "LOW_SUB_UNIVERSE_SHARPE"),
("robust_universe_sharpe", "LOW_ROBUST_UNIVERSE_SHARPE"),
("two_year_sharpe", "LOW_2Y_SHARPE"),
("prod_correlation", "PROD_CORRELATION"),
)
},
}
def submission_condition(submission):
"""Match the platform list contract; a missing status is never assumed submitted."""
@@ -106,6 +155,8 @@ async def upsert_alpha(db, raw: dict):
item.settings, item.is_metrics = sanitize(settings), sanitize(metrics)
item.os_metrics = sanitize(raw.get("os")) if isinstance(raw.get("os"), dict) else {}
item.checks = sanitize(metrics.get("checks") or raw.get("checks") or [])
for key, value in snapshot_columns(item.settings, item.is_metrics, item.checks).items():
setattr(item, key, value)
for key in ("sharpe", "fitness", "returns", "turnover", "margin", "drawdown"):
setattr(item, key, number(metrics.get(key)))
item.date_created, item.date_submitted = date(raw.get("dateCreated")), date(raw.get("dateSubmitted"))
@@ -134,7 +185,7 @@ def list_statement(filters):
)
)
)
for name in ("region", "universe", "alpha_type", "language", "status", "stage", "hidden"):
for name in ("region", "universe", "alpha_type", "language", "status", "stage", "hidden", "check_type", "neutralization"):
value = getattr(filters, name)
if value is not None:
query = query.where(getattr(Alpha, name) == value)
@@ -148,7 +199,7 @@ def list_statement(filters):
query = query.where(Alpha.date_created >= filters.created_from)
if filters.created_to:
query = query.where(Alpha.date_created <= filters.created_to)
for name in ("sharpe", "fitness", "returns", "turnover", "margin", "drawdown"):
for name in METRIC_FIELDS:
for suffix, compare in (("min", "ge"), ("max", "le")):
value = getattr(filters, f"{name}_{suffix}")
if value is not None:
@@ -183,8 +234,16 @@ def summary(item: Alpha, research: Research):
"date_created",
"date_submitted",
"synced_at",
"check_type",
"neutralization",
"sub_universe_sharpe",
"robust_universe_sharpe",
"two_year_sharpe",
"prod_correlation",
"pnl",
)
result = {k: getattr(item, k) for k in keys}
result["failed_checks"] = failed_checks(item.checks)
result["expression_preview"] = (item.expression or item.selection or "")[:240]
result["research"] = {
k: getattr(research, k) for k in ("note", "tags", "favorite", "state", "updated_at", "version")
+10 -3
View File
@@ -16,7 +16,7 @@ from sqlalchemy import delete, select, text
from .ai.routes import router as ai_router
from .ai.runtime import AIRuntime
from .alphas import list_statement, sorted_statement
from .alphas import failed_checks, list_statement, sorted_statement
from .backtests.routes import router as backtest_router
from .business import Business, notify_job
from .catalog.research_routes import router as research_catalog_router
@@ -375,11 +375,18 @@ def create_app(settings=None, wq_client=None, ai_model_factory=None):
"turnover",
"margin",
"drawdown",
"sub_universe_sharpe",
"robust_universe_sharpe",
"two_year_sharpe",
"prod_correlation",
"pnl",
"neutralization",
"check_type",
"date_created",
"date_submitted",
"synced_at",
]
writer.writerow(columns + ["research_state", "favorite", "tags", "note"])
writer.writerow(columns + ["failed_checks", "research_state", "favorite", "tags", "note"])
yield buffer.getvalue()
buffer.seek(0)
buffer.truncate(0)
@@ -388,7 +395,7 @@ def create_app(settings=None, wq_client=None, ai_model_factory=None):
async for a, r in rows:
writer.writerow(
[csv_cell(getattr(a, key)) for key in columns]
+ [r.state, r.favorite, csv_cell(";".join(r.tags)), csv_cell(r.note)]
+ [csv_cell(";".join(failed_checks(a.checks))), r.state, r.favorite, csv_cell(";".join(r.tags)), csv_cell(r.note)]
)
yield buffer.getvalue()
buffer.seek(0)
+7
View File
@@ -79,6 +79,13 @@ class Alpha(Base):
turnover: Mapped[float | None] = mapped_column(Float)
margin: Mapped[float | None] = mapped_column(Float)
drawdown: Mapped[float | None] = mapped_column(Float)
sub_universe_sharpe: Mapped[float | None] = mapped_column(Float)
robust_universe_sharpe: Mapped[float | None] = mapped_column(Float)
two_year_sharpe: Mapped[float | None] = mapped_column(Float)
prod_correlation: Mapped[float | None] = mapped_column(Float)
pnl: Mapped[float | None] = mapped_column(Float)
neutralization: Mapped[str | None] = mapped_column(Text)
check_type: Mapped[str] = mapped_column(String(20), default="PENDING", server_default="PENDING", index=True)
date_created: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), index=True)
date_submitted: Mapped[datetime | None] = mapped_column(DateTime(timezone=True))
synced_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=now)
+30 -1
View File
@@ -9,6 +9,7 @@ from pydantic import BaseModel, ConfigDict, Field, field_validator, model_valida
ResearchState = Literal["inbox", "candidate", "optimizing", "archived"]
Submission = Literal["UNSUBMITTED", "SUBMITTED"]
CheckType = Literal["PENDING", "PRE_CHECK", "PASS", "FAIL_1", "FAIL_2"]
SortField = Literal[
"id",
"name",
@@ -18,6 +19,11 @@ SortField = Literal[
"turnover",
"margin",
"drawdown",
"sub_universe_sharpe",
"robust_universe_sharpe",
"two_year_sharpe",
"prod_correlation",
"pnl",
"date_created",
"date_submitted",
"synced_at",
@@ -76,6 +82,8 @@ class AlphaFilters(Contract):
status: str | None = None
stage: str | None = None
hidden: bool | None = None
check_type: CheckType | None = None
neutralization: str | None = None
research_state: ResearchState | None = None
favorite: bool | None = None
tag: str | None = Field(default=None, max_length=60)
@@ -93,6 +101,16 @@ class AlphaFilters(Contract):
margin_max: float | None = None
drawdown_min: float | None = None
drawdown_max: float | None = None
sub_universe_sharpe_min: float | None = None
sub_universe_sharpe_max: float | None = None
robust_universe_sharpe_min: float | None = None
robust_universe_sharpe_max: float | None = None
two_year_sharpe_min: float | None = None
two_year_sharpe_max: float | None = None
prod_correlation_min: float | None = None
prod_correlation_max: float | None = None
pnl_min: float | None = None
pnl_max: float | None = None
sort: SortField = "date_created"
direction: Literal["asc", "desc"] = "desc"
limit: int = Field(default=25, ge=1, le=100)
@@ -105,7 +123,10 @@ class AlphaFilters(Contract):
@model_validator(mode="after")
def range_order(self):
for key in ("sharpe", "fitness", "returns", "turnover", "margin", "drawdown"):
for key in (
"sharpe", "fitness", "returns", "turnover", "margin", "drawdown",
"sub_universe_sharpe", "robust_universe_sharpe", "two_year_sharpe", "prod_correlation", "pnl",
):
lo, hi = getattr(self, f"{key}_min"), getattr(self, f"{key}_max")
if lo is not None and hi is not None and lo > hi:
raise ValueError(f"{key} 最小值不能大于最大值")
@@ -222,6 +243,14 @@ class AlphaSummary(BaseModel):
date_created: datetime | None
date_submitted: datetime | None
synced_at: datetime
check_type: CheckType = "PENDING"
failed_checks: list[str] = Field(default_factory=list)
neutralization: str | None = None
sub_universe_sharpe: float | None = None
robust_universe_sharpe: float | None = None
two_year_sharpe: float | None = None
prod_correlation: float | None = None
pnl: float | None = None
research: ResearchOutput
local_correlation: dict | None = None
source_kinds: list[str] = Field(default_factory=list)
@@ -0,0 +1,125 @@
"""Index Alpha checks and metrics; backfill existing snapshots without upstream calls."""
import math
import sqlalchemy as sa
from alembic import op
revision = "0011"
down_revision = "0010"
branch_labels = None
depends_on = None
# Frozen normalization for historical snapshots; do not import mutable app code.
def number(value):
if value is None or isinstance(value, bool):
return None
try:
result = float(value)
return result if math.isfinite(result) else None
except (ValueError, TypeError):
return None
def failed_checks(checks):
"""Return failed platform check names; local correlation never changes this list."""
return (
[
check.get("name") if isinstance(check.get("name"), str) else "未命名检查"
for check in checks
if isinstance(check, dict) and check.get("result") == "FAIL"
]
if isinstance(checks, list)
else []
)
def snapshot_columns(settings, metrics, checks):
"""Derive list fields from a platform snapshot, preserving missing metrics as null.
Only explicit FAIL results count. Empty, malformed and unfinished checks are
pending; all known checks passing without PROD_CORRELATION is only a pre-check.
No submission eligibility or activity eligibility is inferred here.
"""
settings = settings if isinstance(settings, dict) else {}
metrics = metrics if isinstance(metrics, dict) else {}
checks = checks if isinstance(checks, list) else []
valid = [check for check in checks if isinstance(check, dict)]
failures = len(failed_checks(checks))
by_name = {check["name"]: check for check in valid if isinstance(check.get("name"), str)}
if failures:
check_type = "FAIL_1" if failures == 1 else "FAIL_2"
elif not checks or len(valid) != len(checks) or any(check.get("result") != "PASS" for check in valid):
check_type = "PENDING"
else:
check_type = "PASS" if "PROD_CORRELATION" in by_name else "PRE_CHECK"
neutralization = settings.get("neutralization")
return {
"check_type": check_type,
"neutralization": neutralization if isinstance(neutralization, str) else None,
"pnl": number(metrics.get("pnl")),
**{
field: number(by_name.get(name, {}).get("value"))
for field, name in (
("sub_universe_sharpe", "LOW_SUB_UNIVERSE_SHARPE"),
("robust_universe_sharpe", "LOW_ROBUST_UNIVERSE_SHARPE"),
("two_year_sharpe", "LOW_2Y_SHARPE"),
("prod_correlation", "PROD_CORRELATION"),
)
},
}
METRICS = ("sub_universe_sharpe", "robust_universe_sharpe", "two_year_sharpe", "prod_correlation", "pnl")
def upgrade():
columns = [sa.Column(name, sa.Float(), nullable=True) for name in METRICS]
columns += [
sa.Column("neutralization", sa.Text(), nullable=True),
sa.Column("check_type", sa.String(20), nullable=False, server_default="PENDING"),
]
for column in columns:
op.add_column("alphas", column)
op.create_index("ix_alphas_check_type", "alphas", ["check_type"])
table = sa.table(
"alphas",
sa.column("id", sa.String()),
sa.column("settings", sa.JSON()),
sa.column("is_metrics", sa.JSON()),
sa.column("checks", sa.JSON()),
*(sa.column(column.name, column.type) for column in columns),
)
connection = op.get_bind()
last_id = None
while True:
query = (
sa.select(table.c.id, table.c.settings, table.c.is_metrics, table.c.checks)
.order_by(table.c.id)
.limit(500)
)
if last_id is not None:
query = query.where(table.c.id > last_id)
rows = connection.execute(query).mappings().all()
if not rows:
break
connection.execute(
table.update()
.where(table.c.id == sa.bindparam("snapshot_id"))
.values({column.name: sa.bindparam(column.name) for column in columns}),
[
{
"snapshot_id": row["id"],
**snapshot_columns(row["settings"], row["is_metrics"], row["checks"]),
}
for row in rows
],
)
last_id = rows[-1]["id"]
def downgrade():
op.drop_index("ix_alphas_check_type", table_name="alphas")
for name in ("check_type", "neutralization", *reversed(METRICS)):
op.drop_column("alphas", name)
+227
View File
@@ -0,0 +1,227 @@
"""Alpha list filtering, snapshot extraction, local tags and historical migration."""
import csv
import io
from pathlib import Path
import pytest
import sqlalchemy as sa
from alembic import command
from alembic.config import Config
from cryptography.fernet import Fernet
from app.alphas import snapshot_columns, upsert_alpha
from tests.conftest import alpha
PREFIX = "/api/v1/alphas"
def checks(failures):
return [
{"name": name, "result": "FAIL" if index < failures else "PASS", "value": value}
for index, (name, value) in enumerate(
[
("LOW_SUB_UNIVERSE_SHARPE", 0),
("LOW_ROBUST_UNIVERSE_SHARPE", 1.2),
("LOW_2Y_SHARPE", 2.3),
("PROD_CORRELATION", 0),
]
)
]
@pytest.mark.parametrize(
"data, expected",
[
([], "PENDING"),
(None, "PENDING"),
([None], "PENDING"),
([{}], "PENDING"),
([{"name": "LOW_SHARPE", "result": "WARNING"}], "PENDING"),
([{"name": "LOW_SHARPE", "result": "PASS"}], "PRE_CHECK"),
([{"name": "PROD_CORRELATION", "result": "PENDING"}], "PENDING"),
(checks(0), "PASS"),
(checks(1), "FAIL_1"),
(checks(2), "FAIL_2"),
(checks(3), "FAIL_2"),
],
)
def test_check_classification_never_promotes_unknown_results(data, expected):
assert snapshot_columns({}, {}, data)["check_type"] == expected
async def test_checks_filter_before_pagination_and_share_export_scope(app, logged_in):
async with app.state.sessions.begin() as db:
for count in range(4):
await upsert_alpha(db, alpha(f"failed{count}", **{"is": {"pnl": 0, "checks": checks(count)}}))
await upsert_alpha(db, alpha("unknown", **{"is": {}}))
for check_type, expected in [
("FAIL_1", ["failed1"]),
("FAIL_2", ["failed2", "failed3"]),
("PASS", ["failed0"]),
("PENDING", ["unknown"]),
]:
response = await logged_in.get(
PREFIX, params={"check_type": check_type, "sort": "id", "direction": "asc"}
)
assert response.status_code == 200
assert [row["id"] for row in response.json()["items"]] == expected
query = "check_type=FAIL_2&submission=UNSUBMITTED&region=USA&limit=1&offset=1&sort=id&direction=asc"
page = (await logged_in.get(f"{PREFIX}?{query}")).json()
assert page["total"] == 2 and [a["id"] for a in page["items"]] == ["failed3"]
row = page["items"][0]
assert row["failed_checks"] == ["LOW_SUB_UNIVERSE_SHARPE", "LOW_ROBUST_UNIVERSE_SHARPE", "LOW_2Y_SHARPE"]
assert row["prod_correlation"] == row["sub_universe_sharpe"] == row["pnl"] == 0
exported = await logged_in.get(f"{PREFIX}/export?{query}")
rows = list(csv.DictReader(io.StringIO(exported.text.lstrip("\ufeff"))))
assert [r["id"] for r in rows] == ["failed2", "failed3"]
assert rows[0]["check_type"] == "FAIL_2" and rows[0]["pnl"] == "0.0"
assert rows[0]["failed_checks"] == "LOW_SUB_UNIVERSE_SHARPE;LOW_ROBUST_UNIVERSE_SHARPE"
assert (await logged_in.get(f"{PREFIX}?check_type=FAIL_GT_2")).status_code == 422
async def test_extended_metrics_filter_sort_missing_values_and_snapshot_refresh(app, logged_in):
fields = ["sub_universe_sharpe", "robust_universe_sharpe", "two_year_sharpe", "prod_correlation", "pnl"]
async with app.state.sessions.begin() as db:
for name, value in [("zero", 0), ("positive", 2), ("negative", -1)]:
sample_checks = [{**c, "value": value} for c in checks(0)]
await upsert_alpha(
db,
alpha(
name,
settings={"neutralization": "INDUSTRY"},
**{"is": {"pnl": value, "checks": sample_checks}},
),
)
await upsert_alpha(db, alpha("missing", **{"is": {}}))
for field in fields:
result = (await logged_in.get(PREFIX, params={"sort": field, "direction": "asc"})).json()
assert [r["id"] for r in result["items"]] == ["negative", "zero", "positive", "missing"]
result = (
await logged_in.get(
PREFIX, params={f"{field}_min": 0, f"{field}_max": 0, "neutralization": "INDUSTRY"}
)
).json()
assert [r["id"] for r in result["items"]] == ["zero"]
assert (await logged_in.get(PREFIX, params={f"{field}_min": 2, f"{field}_max": 1})).status_code == 422
async with app.state.sessions.begin() as db:
await upsert_alpha(db, alpha("zero", settings={}, **{"is": {}}))
detail = (await logged_in.get(f"{PREFIX}/zero")).json()
assert all(detail[field] is None for field in fields)
assert detail["neutralization"] is None and detail["check_type"] == "PENDING"
malformed = snapshot_columns({}, {"pnl": "nan"}, [{**c, "value": True} for c in checks(0)])
assert all(malformed[field] is None for field in fields)
async def test_ppac_tags_partial_edit_bulk_filter_and_sync_preservation(app, logged_in):
async with app.state.sessions.begin() as db:
for name in ("ppac", "other"):
await upsert_alpha(db, alpha(name, **{"is": {"checks": checks(1)}}))
assert (
await logged_in.patch(
f"{PREFIX}/ppac/research",
json={
"version": 1,
"note": "等活动轮到再提交",
"state": "candidate",
"favorite": True,
},
)
).status_code == 200
assert (
await logged_in.patch(
f"{PREFIX}/ppac/research",
json={
"version": 2,
"tags": [" PPAC ", "PPAC", "待活动提交"],
},
)
).status_code == 200
assert (
await logged_in.patch(f"{PREFIX}/ppac/research", json={"version": 2, "tags": []})
).status_code == 409
async with app.state.sessions.begin() as db:
await upsert_alpha(db, alpha("ppac", **{"is": {"checks": checks(2)}}))
result = (await logged_in.get(f"{PREFIX}?tag=PPAC&check_type=FAIL_2")).json()
assert result["total"] == 1
research = result["items"][0]["research"]
assert (
research["note"] == "等活动轮到再提交" and research["favorite"] and research["state"] == "candidate"
)
assert research["tags"] == ["PPAC", "待活动提交"] and research["version"] == 3
assert "PPAC" in (await logged_in.get(f"{PREFIX}/facets")).json()["tags"]
assert (
await logged_in.patch(
f"{PREFIX}/research/bulk",
json={
"alpha_ids": ["ppac", "other"],
"versions": {"ppac": 3, "other": 1},
"add_tags": ["活动候选"],
"remove_tags": ["待活动提交"],
},
)
).status_code == 200
assert (await logged_in.get(f"{PREFIX}?tag=活动候选")).json()["total"] == 2
assert (await logged_in.get(f"{PREFIX}?tag=待活动提交")).json()["total"] == 0
def test_migration_backfills_multiple_batches_and_preserves_research(tmp_path, monkeypatch):
path = tmp_path / "migration.db"
monkeypatch.setenv("DATABASE_URL", f"sqlite+aiosqlite:///{path}")
monkeypatch.setenv("ADMIN_PASSWORD", "migration-test-only")
monkeypatch.setenv("ENCRYPTION_KEY", Fernet.generate_key().decode())
monkeypatch.setenv("WQ_EMAIL", "")
monkeypatch.setenv("WQ_PASSWORD", "")
root = Path(__file__).resolve().parents[1]
config = Config(str(root / "alembic.ini"))
config.set_main_option("script_location", str(root / "migrations"))
command.upgrade(config, "0010")
engine = sa.create_engine(f"sqlite:///{path}")
metadata = sa.MetaData()
alphas = sa.Table("alphas", metadata, autoload_with=engine)
research = sa.Table("research", metadata, autoload_with=engine)
from app.models import now
with engine.begin() as db:
db.execute(
alphas.insert(),
[
{
"id": f"old{i:04}",
"hidden": False,
"settings": {"neutralization": "INDUSTRY"},
"is_metrics": {"pnl": 0},
"os_metrics": {},
"checks": checks(i % 4),
"synced_at": now(),
"raw": {},
}
for i in range(503)
],
)
db.execute(
research.insert(),
{
"alpha_id": "old0000",
"note": "keep",
"tags": ["PPAC"],
"favorite": True,
"state": "candidate",
"version": 7,
"updated_at": now(),
},
)
for _ in range(2):
command.upgrade(config, "head")
command.check(config)
alphas = sa.Table("alphas", sa.MetaData(), autoload_with=engine)
with engine.connect() as db:
rows = db.execute(sa.select(alphas).order_by(alphas.c.id)).mappings().all()
assert len(rows) == 503
for i, row in enumerate(rows):
expected = snapshot_columns(row["settings"], row["is_metrics"], checks(i % 4))
assert {key: row[key] for key in expected} == expected
record = db.execute(sa.select(research)).mappings().one()
assert record["tags"] == ["PPAC"] and record["note"] == "keep" and record["version"] == 7
command.downgrade(config, "0010")
engine.dispose()