Training über das Dashboard anstoßen
Neuer Abschnitt "Training" im Dashboard mit zwei Bedienelementen und den zugehörigen Endpunkten /control/train/history, /control/train/live und /control/training. Historisches Nachtraining - Kerzenanzahl je Symbol wählbar (500 bis 50 000), Fortschritt und Ergebnis werden im Dashboard angezeigt. - Läuft mit derselben Logik wie ein Backtest, aber ohne zu handeln, und speichert das Modell anschließend. - Handelsdurchlauf und Nachtraining teilen sich einen Mutex, damit sie nicht gleichzeitig auf Modell und Portfolio zugreifen. Die rechenintensive Schleife läuft in einem Worker-Thread, damit der Status-Server antwortbereit bleibt. - Ein zweiter Start wird abgelehnt, solange einer eingereiht ist oder läuft. Kontinuierliches Lernen - Schalter für das Online-Lernen im laufenden Betrieb. Ausgeschaltet handelt der Bot weiter, verändert das Modell aber nicht mehr. Label-Trennung - Vorgemerkte Labels tragen jetzt ein Tag. Ein Nachtraining darf die offenen Labels des Live-Betriebs weder auflösen noch verwerfen; ohne die Trennung würden sie gegen historische Kurse ausgewertet und das Modell mit falschen Ergebnissen gefüttert. - score() liest die Gewichte unter dem Lock, damit ein parallel laufendes Training keinen halb aktualisierten Vektor sichtbar macht. Absicherung - server.control_token (TRADEMIND_CONTROL_TOKEN) schützt alle Steuerbefehle über den Header X-TradeMind-Token; lesende Endpunkte bleiben offen. Ohne Token warnt der Bot beim Start, wenn der Port nicht nur lokal erreichbar ist. - server.enable_control: false entfernt die Routen vollständig. 153 Tests (25 neue), ruff sauber. Im gebauten Container geprüft: 202/409 beim Anstoßen, 401 ohne und mit falschem Token, +1099 Beobachtungen in 4,3 s bei weiterlaufendem Handels-Loop ohne Fehler.
This commit is contained in:
@@ -0,0 +1,321 @@
|
||||
"""Steuerung über das Dashboard: historisches Nachtraining und Live-Lernschalter."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
from aiohttp import web
|
||||
|
||||
from trademind.config import Config, ServerConfig
|
||||
from trademind.features import N_FEATURES
|
||||
from trademind.server import TOKEN_HEADER, StatusServer
|
||||
|
||||
from .conftest import make_candles
|
||||
from .test_engine import GrowingFeed, StaticFeed, build_engine, cyclical_series
|
||||
|
||||
|
||||
def control_config(base_config: Config, **server: object) -> Config:
|
||||
return Config.model_validate(
|
||||
{**base_config.model_dump(), "mode": "paper", "server": {"enabled": True, **server}}
|
||||
)
|
||||
|
||||
|
||||
# ------------------------------------------------------- Historisches Training
|
||||
|
||||
|
||||
async def test_history_training_can_be_started_and_completes(base_config):
|
||||
engine = build_engine(control_config(base_config), feed=GrowingFeed(cyclical_series(n=900), start=900))
|
||||
await engine.prepare()
|
||||
|
||||
result = engine.start_history_training(bars=900)
|
||||
assert result["accepted"] is True
|
||||
assert engine.training.state == "queued"
|
||||
|
||||
await engine._training_task
|
||||
job = engine.training
|
||||
assert job.state == "done"
|
||||
assert job.samples_gained > 0
|
||||
assert job.symbols_done == ["BTC/USDT"]
|
||||
assert job.bars_seen > 0
|
||||
assert engine.strategy.learner.stats.samples_seen == job.as_dict()["samples_total"]
|
||||
|
||||
|
||||
async def test_bars_are_clamped_to_sane_limits(base_config):
|
||||
engine = build_engine(control_config(base_config), feed=GrowingFeed(cyclical_series(n=900), start=900))
|
||||
await engine.prepare()
|
||||
|
||||
engine.start_history_training(bars=10)
|
||||
assert engine.training.bars_requested == 500 # Untergrenze
|
||||
await engine._training_task
|
||||
|
||||
engine.start_history_training(bars=10_000_000)
|
||||
assert engine.training.bars_requested == 50_000 # Obergrenze
|
||||
await engine._training_task
|
||||
|
||||
|
||||
async def test_second_training_is_rejected_while_one_runs(base_config):
|
||||
engine = build_engine(control_config(base_config), feed=GrowingFeed(cyclical_series(n=900), start=900))
|
||||
await engine.prepare()
|
||||
|
||||
engine.start_history_training(bars=900)
|
||||
second = engine.start_history_training(bars=900)
|
||||
assert second["accepted"] is False
|
||||
assert "läuft bereits" in second["reason"]
|
||||
await engine._training_task
|
||||
|
||||
|
||||
async def test_training_does_not_trade(base_config):
|
||||
engine = build_engine(control_config(base_config), feed=GrowingFeed(cyclical_series(n=900), start=900))
|
||||
await engine.prepare()
|
||||
engine.start_history_training(bars=900)
|
||||
await engine._training_task
|
||||
|
||||
assert engine.portfolio.trades == []
|
||||
assert engine.portfolio.positions == {}
|
||||
assert await engine.broker.cash() == pytest.approx(base_config.paper.starting_balance)
|
||||
|
||||
|
||||
async def test_training_leaves_live_labels_untouched(base_config):
|
||||
"""Historisches Training darf offene Live-Labels weder auflösen noch verwerfen."""
|
||||
engine = build_engine(control_config(base_config), feed=GrowingFeed(cyclical_series(n=900), start=900))
|
||||
await engine.prepare()
|
||||
learner = engine.strategy.learner
|
||||
|
||||
learner.register_candidate("BTC/USDT", np.ones(N_FEATURES), 100.0, bar_index=5, tag="live")
|
||||
assert learner.pending_count == 1
|
||||
|
||||
engine.start_history_training(bars=900)
|
||||
await engine._training_task
|
||||
|
||||
remaining = [p for p in learner._pending if p.tag == "live"]
|
||||
assert len(remaining) == 1
|
||||
assert remaining[0].entry_price == 100.0
|
||||
assert not [p for p in learner._pending if p.tag == "history"]
|
||||
|
||||
|
||||
async def test_training_does_not_shift_the_live_timeline(base_config):
|
||||
engine = build_engine(control_config(base_config), feed=GrowingFeed(cyclical_series(n=900), start=900))
|
||||
await engine.prepare()
|
||||
engine.bar_counter["BTC/USDT"] = 42
|
||||
engine.last_bar_ts["BTC/USDT"] = 1_234_567
|
||||
|
||||
engine.start_history_training(bars=900)
|
||||
await engine._training_task
|
||||
|
||||
assert engine.bar_counter["BTC/USDT"] == 42
|
||||
assert engine.last_bar_ts["BTC/USDT"] == 1_234_567
|
||||
|
||||
|
||||
async def test_bootstrap_still_adopts_the_timeline(base_config):
|
||||
"""Beim Kaltstart soll der Live-Loop dagegen an die Historie anschließen."""
|
||||
engine = build_engine(control_config(base_config), feed=GrowingFeed(cyclical_series(n=900), start=900))
|
||||
await engine.prepare()
|
||||
await engine.bootstrap_learner()
|
||||
assert engine.bar_counter["BTC/USDT"] > 0
|
||||
|
||||
|
||||
async def test_training_and_tick_do_not_overlap(base_config):
|
||||
engine = build_engine(control_config(base_config), feed=GrowingFeed(cyclical_series(n=900), start=900))
|
||||
await engine.prepare()
|
||||
|
||||
engine.start_history_training(bars=900)
|
||||
tick = asyncio.create_task(engine._tick(300))
|
||||
await asyncio.gather(engine._training_task, tick)
|
||||
|
||||
assert engine.training.state == "done"
|
||||
assert engine.errors == 0
|
||||
|
||||
|
||||
async def test_training_failure_is_reported_not_raised(base_config):
|
||||
class BrokenFeed(StaticFeed):
|
||||
async def fetch(self, symbol, timeframe, limit):
|
||||
raise RuntimeError("Börse nicht erreichbar")
|
||||
|
||||
engine = build_engine(control_config(base_config), feed=BrokenFeed())
|
||||
await engine.prepare()
|
||||
engine.start_history_training(bars=900)
|
||||
await engine._training_task
|
||||
|
||||
job = engine.training
|
||||
assert job.state == "done" # der Lauf endet geordnet
|
||||
assert job.symbols_done == []
|
||||
assert any("nicht erreichbar" in s for s in job.skipped)
|
||||
|
||||
|
||||
async def test_short_history_is_skipped(base_config):
|
||||
engine = build_engine(control_config(base_config), feed=StaticFeed(make_candles(n=60)))
|
||||
await engine.prepare()
|
||||
engine.start_history_training(bars=900)
|
||||
await engine._training_task
|
||||
assert any("Kerzen" in s for s in engine.training.skipped)
|
||||
|
||||
|
||||
async def test_rules_strategy_cannot_be_trained(base_config):
|
||||
config = Config.model_validate(
|
||||
{**control_config(base_config).model_dump(), "strategy": {"name": "rules"}}
|
||||
)
|
||||
engine = build_engine(config, feed=StaticFeed())
|
||||
await engine.prepare()
|
||||
result = engine.start_history_training()
|
||||
assert result["accepted"] is False
|
||||
assert "lernfähiges Modell" in result["reason"]
|
||||
|
||||
|
||||
# ------------------------------------------------------- Live-Lernen umschalten
|
||||
|
||||
|
||||
async def test_online_learning_can_be_toggled(base_config):
|
||||
engine = build_engine(control_config(base_config), feed=StaticFeed())
|
||||
await engine.prepare()
|
||||
assert engine.online_learning_enabled is True
|
||||
|
||||
assert engine.set_online_learning(False)["online_learning"] is False
|
||||
assert engine.strategy.learner.frozen is True
|
||||
assert engine.online_learning_enabled is False
|
||||
|
||||
assert engine.set_online_learning(True)["online_learning"] is True
|
||||
assert engine.strategy.learner.frozen is False
|
||||
|
||||
|
||||
async def test_frozen_learner_stops_updating_weights(base_config):
|
||||
engine = build_engine(control_config(base_config), feed=StaticFeed())
|
||||
await engine.prepare()
|
||||
learner = engine.strategy.learner
|
||||
for _ in range(60):
|
||||
learner.observe(np.ones(N_FEATURES), 1.0)
|
||||
|
||||
engine.set_online_learning(False)
|
||||
weights = learner.model.w.copy()
|
||||
for _ in range(60):
|
||||
learner.observe(np.ones(N_FEATURES), 0.0)
|
||||
assert np.allclose(learner.model.w, weights)
|
||||
|
||||
|
||||
async def test_control_disabled_rejects_commands(base_config):
|
||||
engine = build_engine(control_config(base_config, enable_control=False), feed=StaticFeed())
|
||||
await engine.prepare()
|
||||
assert engine.start_history_training()["accepted"] is False
|
||||
assert engine.set_online_learning(False)["accepted"] is False
|
||||
|
||||
|
||||
# --------------------------------------------------------------- HTTP-Schicht
|
||||
|
||||
|
||||
class SlowFeed(GrowingFeed):
|
||||
"""Verzögert den Abruf, damit sich Anfragen zuverlässig überlappen."""
|
||||
|
||||
async def fetch(self, symbol: str, timeframe: str, limit: int):
|
||||
await asyncio.sleep(0.3)
|
||||
return await super().fetch(symbol, timeframe, limit)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
async def client(aiohttp_client, base_config):
|
||||
"""Server mit angeschlossener Engine; gibt (client, engine) zurück."""
|
||||
|
||||
async def _make(feed=None, **server_options: object):
|
||||
config = control_config(base_config, **server_options)
|
||||
engine = build_engine(config, feed=feed or GrowingFeed(cyclical_series(n=900), start=900))
|
||||
await engine.prepare()
|
||||
srv = StatusServer(config.server, engine.status, controller=engine)
|
||||
return await aiohttp_client(srv._build_app()), engine
|
||||
|
||||
return _make
|
||||
|
||||
|
||||
async def test_dashboard_and_status_expose_training(client):
|
||||
http, _ = await client()
|
||||
page = await (await http.get("/")).text()
|
||||
assert "Training starten" in page
|
||||
assert "Kontinuierliches Lernen" in page
|
||||
|
||||
status = await (await http.get("/status")).json()
|
||||
assert status["training"]["control_enabled"] is True
|
||||
assert status["training"]["online_learning"] is True
|
||||
|
||||
|
||||
async def test_post_history_training_returns_202(client):
|
||||
http, engine = await client()
|
||||
response = await http.post("/control/train/history", json={"bars": 900})
|
||||
assert response.status == 202
|
||||
assert (await response.json())["accepted"] is True
|
||||
await engine._training_task
|
||||
|
||||
|
||||
async def test_post_history_training_conflict_returns_409(client):
|
||||
http, engine = await client(feed=SlowFeed(cyclical_series(n=900), start=900))
|
||||
first = await http.post("/control/train/history", json={"bars": 900})
|
||||
assert first.status == 202
|
||||
|
||||
second = await http.post("/control/train/history", json={"bars": 900})
|
||||
assert second.status == 409
|
||||
assert "läuft bereits" in (await second.json())["reason"]
|
||||
|
||||
await engine._training_task
|
||||
assert engine.training.state == "done"
|
||||
|
||||
|
||||
async def test_post_live_toggle(client):
|
||||
http, engine = await client()
|
||||
response = await http.post("/control/train/live", json={"enabled": False})
|
||||
assert response.status == 200
|
||||
assert (await response.json())["online_learning"] is False
|
||||
assert engine.strategy.learner.frozen is True
|
||||
|
||||
|
||||
async def test_live_toggle_rejects_bad_payload(client):
|
||||
http, _ = await client()
|
||||
assert (await http.post("/control/train/live", json={"enabled": "ja"})).status == 400
|
||||
assert (await http.post("/control/train/live", data="kein json")).status == 400
|
||||
|
||||
|
||||
async def test_token_is_required_when_configured(client):
|
||||
http, engine = await client(control_token="geheim")
|
||||
|
||||
assert (await http.post("/control/train/live", json={"enabled": False})).status == 401
|
||||
assert (await http.get("/control/training")).status == 401
|
||||
assert engine.strategy.learner.frozen is False # nichts passiert
|
||||
|
||||
ok = await http.post(
|
||||
"/control/train/live", json={"enabled": False}, headers={TOKEN_HEADER: "geheim"}
|
||||
)
|
||||
assert ok.status == 200
|
||||
assert engine.strategy.learner.frozen is True
|
||||
|
||||
|
||||
async def test_wrong_token_is_rejected(client):
|
||||
http, _ = await client(control_token="geheim")
|
||||
response = await http.post(
|
||||
"/control/train/live", json={"enabled": False}, headers={TOKEN_HEADER: "falsch"}
|
||||
)
|
||||
assert response.status == 401
|
||||
|
||||
|
||||
async def test_control_routes_absent_when_disabled(client):
|
||||
http, _ = await client(enable_control=False)
|
||||
assert (await http.post("/control/train/history", json={})).status == 404
|
||||
assert (await http.get("/control/training")).status == 404
|
||||
assert (await http.get("/status")).status == 200 # lesende Endpunkte bleiben
|
||||
|
||||
|
||||
async def test_read_only_endpoints_need_no_token(client):
|
||||
http, _ = await client(control_token="geheim")
|
||||
for path in ("/health", "/status", "/metrics", "/positions", "/trades"):
|
||||
assert (await http.get(path)).status == 200, path
|
||||
|
||||
|
||||
async def test_training_state_endpoint(client):
|
||||
http, _ = await client()
|
||||
data = await (await http.get("/control/training")).json()
|
||||
assert data["state"] == "idle"
|
||||
assert data["learning_available"] is True
|
||||
assert data["min_bars"] == 500
|
||||
|
||||
|
||||
def test_server_app_builds_without_controller(base_config):
|
||||
srv = StatusServer(ServerConfig(), lambda: {}, controller=None)
|
||||
app = srv._build_app()
|
||||
assert isinstance(app, web.Application)
|
||||
assert srv.control_available is False
|
||||
Reference in New Issue
Block a user