You cannot select more than 25 topics
Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.
361 lines
13 KiB
Python
361 lines
13 KiB
Python
"""Experiment management API endpoints."""
|
|
from __future__ import annotations
|
|
|
|
import datetime
|
|
import json
|
|
import os
|
|
import re
|
|
from pathlib import Path
|
|
from typing import Any
|
|
|
|
from fastapi import APIRouter, HTTPException, Query
|
|
|
|
from apps.web.dependencies import get_configs_dir, get_journal_dir, get_project_root
|
|
from apps.web.experiment_baked_sleeves import read_effective_baked_sleeves
|
|
from libs.backtest.experiments import (
|
|
build_lineage_tree,
|
|
diff_experiments,
|
|
get_ancestor_chain,
|
|
normalize_experiment_status,
|
|
next_experiment_id,
|
|
search_experiments,
|
|
set_experiment_status,
|
|
create_experiment,
|
|
rebuild_experiment_index,
|
|
validate_all_experiments,
|
|
)
|
|
|
|
router = APIRouter(prefix="/experiments", tags=["experiments"])
|
|
|
|
_VALID_STATUSES = {"draft", "active", "promoted", "retired"}
|
|
|
|
|
|
def _read_json(path: Path) -> dict[str, Any]:
|
|
return json.loads(path.read_text())
|
|
|
|
|
|
@router.get("")
|
|
def list_experiments(
|
|
status: str | None = Query(None),
|
|
family: str | None = Query(None),
|
|
tag: str | None = Query(None),
|
|
pattern: str | None = Query(None),
|
|
parent: str | None = Query(None),
|
|
has_journal: bool | None = Query(None),
|
|
include_retired: bool = Query(False),
|
|
sort_by: str = Query("sqs", description="sqs | name | id | status | family"),
|
|
limit: int = Query(200, le=1000),
|
|
offset: int = 0,
|
|
) -> dict[str, Any]:
|
|
"""List/search experiments with optional filters."""
|
|
configs_dir = get_configs_dir()
|
|
journal_dir = get_journal_dir()
|
|
journal_path = journal_dir / "improvement_journal.jsonl"
|
|
|
|
results = search_experiments(
|
|
configs_dir=configs_dir,
|
|
journal_path=journal_path if journal_path.exists() else None,
|
|
tag=tag,
|
|
status=status,
|
|
include_retired=include_retired,
|
|
version_family=family,
|
|
parent=parent,
|
|
name_pattern=pattern,
|
|
has_journal_entry=has_journal,
|
|
)
|
|
|
|
# Enrich with latest SQS from registry (index often lacks SQS if built without journal)
|
|
registry_path = journal_dir / "experiment_registry.json"
|
|
if registry_path.exists():
|
|
try:
|
|
registry = json.loads(registry_path.read_text())
|
|
# Build map: experiment_name -> latest sqs_score (most recent entry wins)
|
|
sqs_map: dict[str, float | None] = {}
|
|
for entry in registry.get("entries", []):
|
|
name = entry.get("experiment_name")
|
|
sqs = entry.get("sqs_score")
|
|
if name and sqs is not None:
|
|
sqs_map[name] = sqs
|
|
for r in results:
|
|
if r.get("sqs_score") is None and r["name"] in sqs_map:
|
|
r["sqs_score"] = sqs_map[r["name"]]
|
|
r["has_journal_entry"] = True
|
|
except Exception:
|
|
pass
|
|
|
|
# Sort
|
|
if sort_by == "sqs":
|
|
results.sort(key=lambda e: (e.get("sqs_score") is None, -(e.get("sqs_score") or 0)))
|
|
elif sort_by == "id":
|
|
results.sort(key=lambda e: (e.get("id") is None, e.get("id") or 0))
|
|
elif sort_by == "name":
|
|
results.sort(key=lambda e: e.get("name", ""))
|
|
elif sort_by == "family":
|
|
results.sort(key=lambda e: (e.get("version_family") or "", e.get("name", "")))
|
|
elif sort_by == "status":
|
|
status_order = {"promoted": 0, "active": 1, "draft": 2, "retired": 3}
|
|
results.sort(key=lambda e: status_order.get(normalize_experiment_status(e.get("status", "active")), 4))
|
|
|
|
total = len(results)
|
|
return {"experiments": results[offset : offset + limit], "total": total}
|
|
|
|
|
|
@router.get("/families")
|
|
def list_families() -> dict[str, Any]:
|
|
"""List all distinct version families."""
|
|
results = search_experiments(configs_dir=get_configs_dir())
|
|
families = sorted({e.get("version_family") for e in results if e.get("version_family")})
|
|
return {"families": families}
|
|
|
|
|
|
@router.get("/tags")
|
|
def list_tags() -> dict[str, Any]:
|
|
"""List all distinct tags."""
|
|
results = search_experiments(configs_dir=get_configs_dir())
|
|
all_tags: set[str] = set()
|
|
for e in results:
|
|
all_tags.update(e.get("tags", []))
|
|
return {"tags": sorted(all_tags)}
|
|
|
|
|
|
@router.get("/{name}")
|
|
def get_experiment(name: str) -> dict[str, Any]:
|
|
"""Get full experiment config."""
|
|
configs_dir = get_configs_dir()
|
|
path = configs_dir / f"{name}.json"
|
|
if not path.exists():
|
|
raise HTTPException(status_code=404, detail=f"Experiment not found: {name}")
|
|
data = _read_json(path)
|
|
# Check if a journal entry exists for this experiment
|
|
journal_dir = get_journal_dir()
|
|
registry_path = journal_dir / "experiment_registry.json"
|
|
has_journal = False
|
|
if registry_path.exists():
|
|
try:
|
|
registry = json.loads(registry_path.read_text())
|
|
has_journal = any(
|
|
e.get("experiment_name") == name
|
|
for e in registry.get("entries", [])
|
|
)
|
|
except Exception:
|
|
pass
|
|
data["baked_sleeves"] = read_effective_baked_sleeves(path, get_project_root())
|
|
data["has_journal_entry"] = has_journal
|
|
return data
|
|
|
|
|
|
@router.get("/{name}/lineage")
|
|
def get_lineage(name: str) -> dict[str, Any]:
|
|
"""Get lineage tree rooted at name's root ancestor + ancestor chain."""
|
|
configs_dir = get_configs_dir()
|
|
path = configs_dir / f"{name}.json"
|
|
if not path.exists():
|
|
raise HTTPException(status_code=404, detail=f"Experiment not found: {name}")
|
|
|
|
# Walk up to root
|
|
chain = get_ancestor_chain(name, configs_dir)
|
|
root_name = chain[-1]["name"] if chain else name
|
|
|
|
try:
|
|
tree = build_lineage_tree(root_name, configs_dir)
|
|
except KeyError as exc:
|
|
raise HTTPException(status_code=404, detail=str(exc)) from exc
|
|
|
|
return {"tree": tree, "ancestor_chain": chain, "current": name}
|
|
|
|
|
|
@router.get("/{name}/diff/{other}")
|
|
def get_diff(name: str, other: str) -> dict[str, Any]:
|
|
"""Get diff between two experiments."""
|
|
configs_dir = get_configs_dir()
|
|
for n in (name, other):
|
|
if not (configs_dir / f"{n}.json").exists():
|
|
raise HTTPException(status_code=404, detail=f"Experiment not found: {n}")
|
|
try:
|
|
diffs = diff_experiments(name, other, configs_dir)
|
|
except FileNotFoundError as exc:
|
|
raise HTTPException(status_code=404, detail=str(exc)) from exc
|
|
|
|
# Convert tuple values to lists for JSON serialization
|
|
return {k: list(v) for k, v in diffs.items()}
|
|
|
|
|
|
@router.put("/{name}/status")
|
|
def update_status(name: str, body: dict[str, Any]) -> dict[str, Any]:
|
|
"""Promote or retire an experiment."""
|
|
configs_dir = get_configs_dir()
|
|
path = configs_dir / f"{name}.json"
|
|
if not path.exists():
|
|
raise HTTPException(status_code=404, detail=f"Experiment not found: {name}")
|
|
|
|
new_status = normalize_experiment_status(body.get("status"))
|
|
if new_status not in _VALID_STATUSES:
|
|
raise HTTPException(status_code=422, detail=f"Invalid status: {new_status}. Must be one of {_VALID_STATUSES}")
|
|
|
|
set_experiment_status(name, new_status, configs_dir=configs_dir)
|
|
|
|
# Handle alias separately (promote command pattern)
|
|
alias = body.get("alias")
|
|
if alias:
|
|
path = configs_dir / f"{name}.json"
|
|
data = json.loads(path.read_text())
|
|
aliases = data.get("aliases", [])
|
|
if alias not in aliases:
|
|
aliases.append(alias)
|
|
data["aliases"] = aliases
|
|
path.write_text(json.dumps(data, indent=2, ensure_ascii=False) + "\n")
|
|
|
|
return {"name": name, "status": new_status}
|
|
|
|
|
|
@router.post("/compose")
|
|
def compose_experiment(body: dict[str, Any]) -> dict[str, Any]:
|
|
"""Create a new experiment by composing a base strategy with sleeve presets."""
|
|
configs_dir = get_configs_dir()
|
|
base_name = body.get("base")
|
|
new_name = body.get("name")
|
|
changelog = body.get("changelog", "")
|
|
parking = body.get("parking") or None
|
|
idle_alpha = body.get("idle_alpha") or None
|
|
form4_sleeve = body.get("form4_sleeve") or None
|
|
ownership_sleeve = body.get("ownership_sleeve") or None
|
|
|
|
if not base_name:
|
|
raise HTTPException(status_code=422, detail="base is required")
|
|
if not new_name:
|
|
raise HTTPException(status_code=422, detail="name is required")
|
|
|
|
base_path = configs_dir / f"{base_name}.json"
|
|
if not base_path.exists():
|
|
raise HTTPException(status_code=404, detail=f"Base experiment not found: {base_name}")
|
|
|
|
new_path = configs_dir / f"{new_name}.json"
|
|
if new_path.exists():
|
|
raise HTTPException(status_code=400, detail=f"Experiment already exists: {new_name}")
|
|
|
|
# Clone base config
|
|
config = _read_json(base_path)
|
|
|
|
# Update metadata
|
|
config["experiment_name"] = new_name
|
|
config["id"] = next_experiment_id(configs_dir)
|
|
config["parent"] = base_name
|
|
config["changelog"] = changelog or f"Composed from {base_name} with sleeves"
|
|
config["status"] = "draft"
|
|
config["generation"] = (config.get("generation") or 0) + 1
|
|
config["created_at"] = datetime.datetime.now(datetime.timezone.utc).isoformat()
|
|
config["performance_summary"] = None
|
|
# Infer version_family from the new name
|
|
m = re.search(r"_(v\d+[a-z]*)[\._]", new_name) or re.search(r"_(v\d+[a-z]*)$", new_name)
|
|
config["version_family"] = m.group(1) if m else config.get("version_family")
|
|
|
|
# Merge sleeve presets into overrides
|
|
overrides = config.setdefault("overrides", {})
|
|
|
|
if parking:
|
|
risk_overrides = overrides.setdefault("risk", {})
|
|
risk_overrides["cash_parking_preset"] = parking
|
|
else:
|
|
if isinstance(overrides.get("risk"), dict):
|
|
overrides["risk"].pop("cash_parking_preset", None)
|
|
|
|
if idle_alpha:
|
|
overrides["idle_alpha_sleeve_preset"] = idle_alpha
|
|
else:
|
|
overrides.pop("idle_alpha_sleeve_preset", None)
|
|
|
|
if form4_sleeve:
|
|
overrides["form4_capture_sleeve_preset"] = form4_sleeve
|
|
else:
|
|
overrides.pop("form4_capture_sleeve_preset", None)
|
|
|
|
if ownership_sleeve:
|
|
overrides["ownership_capture_sleeve_preset"] = ownership_sleeve
|
|
else:
|
|
overrides.pop("ownership_capture_sleeve_preset", None)
|
|
|
|
new_path.write_text(json.dumps(config, indent=2, ensure_ascii=False) + "\n")
|
|
rebuild_experiment_index(configs_dir)
|
|
|
|
return {"name": new_name, "parent": base_name}
|
|
|
|
|
|
@router.post("")
|
|
def create_new_experiment(body: dict[str, Any]) -> dict[str, Any]:
|
|
"""Create a new experiment cloned from a parent."""
|
|
configs_dir = get_configs_dir()
|
|
parent = body.get("parent")
|
|
changelog = body.get("changelog", "")
|
|
new_name = body.get("name")
|
|
|
|
if not parent:
|
|
raise HTTPException(status_code=422, detail="parent is required")
|
|
if not new_name:
|
|
raise HTTPException(status_code=422, detail="name is required")
|
|
|
|
if not (configs_dir / f"{parent}.json").exists():
|
|
raise HTTPException(status_code=404, detail=f"Parent experiment not found: {parent}")
|
|
|
|
try:
|
|
result_path = create_experiment(
|
|
parent_name=parent,
|
|
new_name=new_name,
|
|
changelog=changelog,
|
|
configs_dir=configs_dir,
|
|
)
|
|
except Exception as exc:
|
|
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
|
|
|
return {"name": result_path.stem, "parent": parent}
|
|
|
|
|
|
@router.put("/{name}")
|
|
def update_experiment(name: str, body: dict[str, Any]) -> dict[str, Any]:
|
|
"""Update experiment config JSON directly."""
|
|
configs_dir = get_configs_dir()
|
|
path = configs_dir / f"{name}.json"
|
|
if not path.exists():
|
|
raise HTTPException(status_code=404, detail=f"Experiment not found: {name}")
|
|
|
|
# Safety: ensure experiment_name matches
|
|
body.pop("baked_sleeves", None)
|
|
body["experiment_name"] = name
|
|
path.write_text(json.dumps(body, indent=2, ensure_ascii=False) + "\n")
|
|
|
|
# Rebuild index to pick up changes
|
|
rebuild_experiment_index(configs_dir)
|
|
return {"name": name, "updated": True}
|
|
|
|
|
|
@router.delete("/{name}")
|
|
def delete_experiment(name: str) -> dict[str, Any]:
|
|
"""Delete an experiment config file."""
|
|
configs_dir = get_configs_dir()
|
|
path = configs_dir / f"{name}.json"
|
|
if not path.exists():
|
|
raise HTTPException(status_code=404, detail=f"Experiment not found: {name}")
|
|
path.unlink()
|
|
rebuild_experiment_index(configs_dir)
|
|
return {"name": name, "deleted": True}
|
|
|
|
|
|
@router.post("/validate")
|
|
def validate_experiments() -> dict[str, Any]:
|
|
"""Validate all experiment configs."""
|
|
configs_dir = get_configs_dir()
|
|
errors = validate_all_experiments(configs_dir)
|
|
return {"errors": errors, "valid": len(errors) == 0}
|
|
|
|
|
|
@router.post("/rebuild-index")
|
|
def rebuild_index() -> dict[str, Any]:
|
|
"""Force rebuild the experiment index."""
|
|
configs_dir = get_configs_dir()
|
|
journal_dir = get_journal_dir()
|
|
journal_path = journal_dir / "improvement_journal.jsonl"
|
|
rebuild_experiment_index(
|
|
configs_dir,
|
|
journal_path=journal_path if journal_path.exists() else None,
|
|
)
|
|
return {"rebuilt": True}
|