Merge branch 'worktree-training-deploy'
This commit is contained in:
3
.gitignore
vendored
3
.gitignore
vendored
@@ -160,9 +160,6 @@ claude-flow
|
||||
.terraform.lock.hcl
|
||||
**/.terraform.lock.hcl
|
||||
|
||||
# Design docs and implementation plans (generated, local only)
|
||||
docs/plans/
|
||||
|
||||
# Data cache (downloaded market data, not checked in)
|
||||
data/cache/
|
||||
|
||||
|
||||
127
docs/plans/2026-02-27-training-deploy-design.md
Normal file
127
docs/plans/2026-02-27-training-deploy-design.md
Normal file
@@ -0,0 +1,127 @@
|
||||
# Training Pipeline & Model Deployment Design
|
||||
|
||||
**Date**: 2026-02-27
|
||||
**Status**: Approved
|
||||
**Tracks**: Two independent parallel tracks
|
||||
|
||||
## Overview
|
||||
|
||||
Wire the complete training-to-serving pipeline on Kapsule. Track 1 handles training job dispatch (fxt CLI → gRPC → PostgreSQL → K8s GPU Jobs). Track 2 handles model serving and promotion (S3 checkpoints → model loader → inference → operator-gated promotion).
|
||||
|
||||
Both tracks share S3 (`foxhunt-models/`) as the handoff point.
|
||||
|
||||
## Track 1: Training Pipeline
|
||||
|
||||
### Goal
|
||||
|
||||
`fxt train start tft ES.FUT --epochs 50` creates a real GPU training job on Kapsule.
|
||||
|
||||
### Flow
|
||||
|
||||
```
|
||||
fxt train start tft ES.FUT --epochs 50
|
||||
→ gRPC StartTraining (ml_training_service:50053)
|
||||
→ service.rs handler → JobSpawner::spawn_batch() → PostgreSQL
|
||||
→ In-process queue consumer (tokio::spawn, 5s poll)
|
||||
→ JobSpawner::get_next_pending_job()
|
||||
→ K8sDispatcher::dispatch() → batch/v1 Job
|
||||
├── training container (GPU, train_baseline_rl/supervised)
|
||||
└── uploader sidecar (S3 → foxhunt-models/)
|
||||
→ Job completion callback → update status in PostgreSQL
|
||||
```
|
||||
|
||||
### Existing Code
|
||||
|
||||
| Component | File | Status |
|
||||
|-----------|------|--------|
|
||||
| fxt train start | `bin/fxt/src/commands/train/start.rs` | Done (sends gRPC) |
|
||||
| fxt train status | `bin/fxt/src/commands/train/status.rs` | Done |
|
||||
| fxt train list | `bin/fxt/src/commands/train/list.rs` | Done |
|
||||
| K8sDispatcher | `services/ml_training_service/src/k8s_dispatcher.rs` | Done (builds Jobs) |
|
||||
| JobSpawner | `services/ml_training_service/src/job_spawner.rs` | Done (PostgreSQL CRUD) |
|
||||
| Job template | `infra/k8s/training/job-template.yaml` | Done (reference only) |
|
||||
|
||||
### Changes Needed
|
||||
|
||||
1. **Wire gRPC handler to JobSpawner**: `service.rs` StartTraining handler must call `JobSpawner::spawn_batch()` instead of returning a stub response.
|
||||
|
||||
2. **Add queue consumer loop**: In `main.rs`, `tokio::spawn` a loop that polls `get_next_pending_job()` every 5s and calls `K8sDispatcher::dispatch()`.
|
||||
|
||||
3. **Job completion callback**: The uploader sidecar already has `CALLBACK_ENDPOINT`. Add a gRPC or HTTP endpoint in ml_training_service that receives completion notifications and calls `JobSpawner::update_job_status()`.
|
||||
|
||||
4. **Deploy training-output-pvc**: Currently only `training-data-pvc` is deployed. Need `training-output-pvc` for job artifact staging.
|
||||
|
||||
5. **PostgreSQL schema**: Deploy `batch_jobs` and `child_jobs` tables. Add sqlx migration.
|
||||
|
||||
6. **Docker image**: Rebuild ml-training-service image with kube-rs K8sDispatcher compiled in.
|
||||
|
||||
### K8s Resources
|
||||
|
||||
- `ml-training-service` Deployment (exists, needs redeploy with new image)
|
||||
- `training-data-pvc` (exists, 10Gi scw-bssd)
|
||||
- `training-output-pvc` (new, needs creation)
|
||||
- `s3-credentials` Secret (exists)
|
||||
- `scw-registry` Secret (exists)
|
||||
- `gpu-training` node pool (Scaleway, autoscales 0→N)
|
||||
|
||||
## Track 2: Model Serving & Validation
|
||||
|
||||
### Goal
|
||||
|
||||
Trained model checkpoints deploy to a serving endpoint with operator-gated promotion and automatic rollback.
|
||||
|
||||
### Flow
|
||||
|
||||
```
|
||||
Training job completes → uploader sidecar → S3 (foxhunt-models/)
|
||||
→ Callback to ml_training_service
|
||||
→ PromotionManager::register_completion()
|
||||
├── No active model → auto-register
|
||||
└── Compare metrics → PendingPromotion / NoImprovement
|
||||
→ fxt model promote <model-id> (operator approval)
|
||||
→ DeploymentPipeline::deploy()
|
||||
├── Rolling update
|
||||
├── Health check (inference smoke test)
|
||||
└── Auto-rollback on failure
|
||||
```
|
||||
|
||||
### Existing Code
|
||||
|
||||
| Component | File | Status |
|
||||
|-----------|------|--------|
|
||||
| S3ModelLoader | `crates/model_loader/src/lib.rs` | Done (S3 + LRU cache + versioning) |
|
||||
| PromotionManager | `services/ml_training_service/src/promotion_manager.rs` | Done (metric comparison, pending/active tracking) |
|
||||
| DeploymentPipeline | `services/ml_training_service/src/deployment_pipeline.rs` | Done (rolling update, health check, rollback) |
|
||||
| ValidationPipeline | `services/ml_training_service/src/validation_pipeline.rs` | Exists |
|
||||
|
||||
### Changes Needed
|
||||
|
||||
1. **fxt model commands**: Add `fxt model promote`, `fxt model list`, `fxt model status` CLI subcommands.
|
||||
|
||||
2. **Wire promotion callback**: Uploader sidecar completion → PromotionManager::register_completion(). Extract metrics from training output.
|
||||
|
||||
3. **Model serving deployment**: K8s Deployment that loads active model from S3 via S3ModelLoader and serves predictions via gRPC.
|
||||
|
||||
4. **Validation test**: Send sample market data to inference endpoint, verify prediction response shape and latency.
|
||||
|
||||
5. **Proto additions**: Add `PromoteModel`, `ListModels`, `GetModelStatus` RPCs to ml_training proto.
|
||||
|
||||
## Shared Infrastructure
|
||||
|
||||
| Resource | Status | Notes |
|
||||
|----------|--------|-------|
|
||||
| PostgreSQL | Running in foxhunt ns | Needs schema migration |
|
||||
| S3 foxhunt-models | Exists | Bucket + credentials |
|
||||
| training-data-pvc | Deployed | 10Gi, read-only for jobs |
|
||||
| training-output-pvc | Needs creation | Staging for job artifacts |
|
||||
| scw-registry secret | Deployed | Image pull auth |
|
||||
| gpu-training pool | Configured | Autoscales 0→N |
|
||||
|
||||
## Decision Log
|
||||
|
||||
| Decision | Choice | Rationale |
|
||||
|----------|--------|-----------|
|
||||
| Job dispatch path | DB-backed (JobSpawner) | Full history, batch orchestration, rollback |
|
||||
| Queue consumer | In-process tokio task | Zero additional infra, simple |
|
||||
| Track split | Training vs Serving | Independent concerns, parallel development |
|
||||
| Promotion gating | Operator approval | Production safety for HFT system |
|
||||
738
docs/plans/2026-02-27-training-deploy-plan.md
Normal file
738
docs/plans/2026-02-27-training-deploy-plan.md
Normal file
@@ -0,0 +1,738 @@
|
||||
# Training Pipeline & Model Deployment Implementation Plan
|
||||
|
||||
> **For Claude:** REQUIRED SUB-SKILL: Use superpowers:executing-plans to implement this plan task-by-task.
|
||||
|
||||
**Goal:** Wire the complete fxt CLI → training service → K8s GPU jobs pipeline and model promotion flow.
|
||||
|
||||
**Architecture:** Two independent tracks. Track 1 adds JobSpawner (PostgreSQL-backed batch queue) to the start_training gRPC path and adds an in-process queue consumer that dispatches to K8sDispatcher. Track 2 wires the stub promotion/completion RPCs to the real PromotionManager. Both tracks share S3 as the checkpoint handoff.
|
||||
|
||||
**Tech Stack:** Rust, tonic (gRPC), kube-rs, sqlx (PostgreSQL), Scaleway Kapsule (K8s), S3
|
||||
|
||||
---
|
||||
|
||||
## Track 1: Training Pipeline
|
||||
|
||||
### Task 1: Add JobSpawner to MLTrainingServiceImpl
|
||||
|
||||
**Files:**
|
||||
- Modify: `services/ml_training_service/src/service.rs:53-82` (add job_spawner field)
|
||||
- Modify: `services/ml_training_service/src/main.rs:394-424` (construct JobSpawner, pass to service)
|
||||
|
||||
**Step 1: Write failing test**
|
||||
|
||||
Add to `services/ml_training_service/src/service.rs` tests:
|
||||
|
||||
```rust
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_service_impl_has_job_spawner() {
|
||||
// Verify the struct has a job_spawner field
|
||||
// This is a compile-time check — if the field doesn't exist, this won't compile
|
||||
let _: fn(&MLTrainingServiceImpl) -> bool = |s| s.job_spawner.is_some();
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
**Step 2: Run test to verify it fails**
|
||||
|
||||
Run: `SQLX_OFFLINE=true cargo test -p ml_training_service --lib test_service_impl_has_job_spawner`
|
||||
Expected: FAIL — `no field job_spawner on type MLTrainingServiceImpl`
|
||||
|
||||
**Step 3: Add job_spawner field to MLTrainingServiceImpl**
|
||||
|
||||
In `service.rs`, add to the struct (after line 60):
|
||||
|
||||
```rust
|
||||
/// Database-backed job spawner for batch orchestration -- None when DB unavailable.
|
||||
pub job_spawner: Option<Arc<crate::job_spawner::JobSpawner>>,
|
||||
```
|
||||
|
||||
Update `MLTrainingServiceImpl::new()` to accept and store it:
|
||||
|
||||
```rust
|
||||
pub fn new(
|
||||
orchestrator: Arc<TrainingOrchestrator>,
|
||||
tuning_manager: Arc<TuningManager>,
|
||||
_config: MLConfig,
|
||||
promotion_manager: Arc<crate::promotion_manager::PromotionManager>,
|
||||
k8s_dispatcher: Option<Arc<crate::k8s_dispatcher::K8sDispatcher>>,
|
||||
job_spawner: Option<Arc<crate::job_spawner::JobSpawner>>,
|
||||
) -> Self {
|
||||
// ... existing code ...
|
||||
Self {
|
||||
orchestrator,
|
||||
tuning_handlers,
|
||||
batch_tuning_manager,
|
||||
promotion_manager,
|
||||
k8s_dispatcher,
|
||||
job_spawner,
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
**Step 4: Update main.rs to construct JobSpawner**
|
||||
|
||||
In `main.rs` after the database initialization (line ~287), add:
|
||||
|
||||
```rust
|
||||
// Initialize job spawner for batch training orchestration
|
||||
let job_spawner = Arc::new(crate::job_spawner::JobSpawner::new(database.pool().clone()));
|
||||
info!("Job spawner initialized for batch training orchestration");
|
||||
```
|
||||
|
||||
Pass it to the service constructor:
|
||||
|
||||
```rust
|
||||
let training_service = MLTrainingServiceImpl::new(
|
||||
Arc::clone(&orchestrator),
|
||||
Arc::clone(&tuning_manager),
|
||||
ml_config.clone(),
|
||||
promotion_manager,
|
||||
k8s_dispatcher,
|
||||
Some(job_spawner.clone()),
|
||||
);
|
||||
```
|
||||
|
||||
**Step 5: Run test to verify it passes**
|
||||
|
||||
Run: `SQLX_OFFLINE=true cargo test -p ml_training_service --lib test_service_impl_has_job_spawner`
|
||||
Expected: PASS
|
||||
|
||||
**Step 6: Commit**
|
||||
|
||||
```bash
|
||||
git add services/ml_training_service/src/service.rs services/ml_training_service/src/main.rs
|
||||
git commit -m "feat(training): add JobSpawner to MLTrainingServiceImpl"
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### Task 2: Wire start_training to JobSpawner
|
||||
|
||||
**Files:**
|
||||
- Modify: `services/ml_training_service/src/service.rs:250-346` (start_training handler)
|
||||
|
||||
**Step 1: Write failing test**
|
||||
|
||||
```rust
|
||||
#[tokio::test]
|
||||
async fn test_start_training_creates_batch_job() {
|
||||
// This test validates the flow: start_training -> JobSpawner::spawn_batch
|
||||
// It will be an integration test that requires a database
|
||||
// For now, test the model-to-binary mapping which is pure
|
||||
use crate::k8s_dispatcher::training_binary_for_model;
|
||||
assert_eq!(training_binary_for_model("tft"), "train_baseline_supervised");
|
||||
assert_eq!(training_binary_for_model("dqn"), "train_baseline_rl");
|
||||
}
|
||||
```
|
||||
|
||||
**Step 2: Modify start_training handler**
|
||||
|
||||
In the `start_training` method, after the K8s dispatcher attempt (line ~322), add the JobSpawner path. The logic should be:
|
||||
|
||||
1. If `job_spawner` is Some AND `k8s_dispatcher` is Some → use batch flow (JobSpawner creates DB records, queue consumer dispatches)
|
||||
2. If only `k8s_dispatcher` → dispatch directly (current behavior)
|
||||
3. Fallback → in-process orchestrator
|
||||
|
||||
```rust
|
||||
// After the existing K8s dispatcher block, before the orchestrator fallback:
|
||||
if let Some(ref spawner) = self.job_spawner {
|
||||
let asset = crate::job_spawner::Asset {
|
||||
symbol: symbol.clone(),
|
||||
data_file: std::path::PathBuf::from(format!("/data/futures-baseline/{}", symbol)),
|
||||
};
|
||||
let model = common::model_types::ModelType::from_str(&req.model_type)
|
||||
.map_err(|e| Status::invalid_argument(format!("Invalid model type: {}", e)))?;
|
||||
|
||||
match spawner.spawn_batch(vec![asset], vec![model]).await {
|
||||
Ok(batch) => {
|
||||
info!(batch_id = %batch.batch_id, "Batch job created in database");
|
||||
return Ok(Response::new(StartTrainingResponse {
|
||||
job_id: batch.batch_id.to_string(),
|
||||
status: ProtoTrainingStatus::Pending as i32,
|
||||
message: format!("Batch training job queued: {}", batch.batch_id),
|
||||
}));
|
||||
}
|
||||
Err(e) => {
|
||||
warn!("JobSpawner failed, falling back to orchestrator: {}", e);
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
**Step 3: Run tests**
|
||||
|
||||
Run: `SQLX_OFFLINE=true cargo check -p ml_training_service`
|
||||
Expected: Compiles (may need to add `from_str` import or adjust ModelType conversion)
|
||||
|
||||
**Step 4: Commit**
|
||||
|
||||
```bash
|
||||
git add services/ml_training_service/src/service.rs
|
||||
git commit -m "feat(training): wire start_training gRPC handler to JobSpawner"
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### Task 3: Add queue consumer loop
|
||||
|
||||
**Files:**
|
||||
- Create: `services/ml_training_service/src/queue_consumer.rs`
|
||||
- Modify: `services/ml_training_service/src/lib.rs` (add module)
|
||||
- Modify: `services/ml_training_service/src/main.rs` (spawn consumer)
|
||||
|
||||
**Step 1: Write failing test**
|
||||
|
||||
In `queue_consumer.rs`:
|
||||
|
||||
```rust
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_queue_consumer_config_defaults() {
|
||||
let config = QueueConsumerConfig::default();
|
||||
assert_eq!(config.poll_interval_secs, 5);
|
||||
assert_eq!(config.max_concurrent_jobs, 1);
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
**Step 2: Implement queue consumer**
|
||||
|
||||
```rust
|
||||
//! Queue Consumer — polls JobSpawner for pending jobs and dispatches to K8s.
|
||||
//!
|
||||
//! Runs as a tokio::spawn'd task inside the ml_training_service process.
|
||||
|
||||
use std::sync::Arc;
|
||||
use std::time::Duration;
|
||||
|
||||
use tracing::{error, info, warn};
|
||||
|
||||
use crate::job_spawner::JobSpawner;
|
||||
use crate::k8s_dispatcher::{K8sDispatcher, TrainingJobParams, training_binary_for_model};
|
||||
|
||||
/// Configuration for the queue consumer.
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct QueueConsumerConfig {
|
||||
/// Seconds between polls for pending jobs.
|
||||
pub poll_interval_secs: u64,
|
||||
/// Maximum concurrent K8s jobs to dispatch (currently sequential).
|
||||
pub max_concurrent_jobs: u32,
|
||||
}
|
||||
|
||||
impl Default for QueueConsumerConfig {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
poll_interval_secs: 5,
|
||||
max_concurrent_jobs: 1,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Starts the queue consumer loop. Call via `tokio::spawn(run_queue_consumer(...))`.
|
||||
pub async fn run_queue_consumer(
|
||||
spawner: Arc<JobSpawner>,
|
||||
dispatcher: Arc<K8sDispatcher>,
|
||||
config: QueueConsumerConfig,
|
||||
) {
|
||||
info!(
|
||||
poll_interval = config.poll_interval_secs,
|
||||
"Queue consumer started"
|
||||
);
|
||||
|
||||
let interval = Duration::from_secs(config.poll_interval_secs);
|
||||
|
||||
loop {
|
||||
match spawner.get_next_pending_job().await {
|
||||
Ok(Some(job)) => {
|
||||
let job_id = job.id;
|
||||
info!(job_id = %job_id, model = %job.model_type, "Dispatching pending job");
|
||||
|
||||
// Mark as running
|
||||
if let Err(e) = spawner.update_job_status(job_id, "Running").await {
|
||||
error!(job_id = %job_id, "Failed to mark job as Running: {}", e);
|
||||
continue;
|
||||
}
|
||||
|
||||
// Extract symbol from config_json
|
||||
let symbol = job.config_json
|
||||
.get("asset")
|
||||
.and_then(|v| v.as_str())
|
||||
.unwrap_or("ES.FUT")
|
||||
.to_string();
|
||||
|
||||
let binary = training_binary_for_model(&job.model_type);
|
||||
|
||||
let params = TrainingJobParams {
|
||||
job_id,
|
||||
model_type: job.model_type.clone(),
|
||||
symbol,
|
||||
epochs: 50, // TODO: extract from config_json
|
||||
binary: binary.to_string(),
|
||||
};
|
||||
|
||||
match dispatcher.dispatch(¶ms).await {
|
||||
Ok(k8s_name) => {
|
||||
info!(
|
||||
job_id = %job_id,
|
||||
k8s_name = %k8s_name,
|
||||
"K8s job dispatched successfully"
|
||||
);
|
||||
}
|
||||
Err(e) => {
|
||||
error!(job_id = %job_id, "K8s dispatch failed: {}", e);
|
||||
if let Err(e2) = spawner.update_job_status(job_id, "Failed").await {
|
||||
error!(job_id = %job_id, "Failed to mark job as Failed: {}", e2);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
Ok(None) => {
|
||||
// No pending jobs — sleep and retry
|
||||
}
|
||||
Err(e) => {
|
||||
warn!("Error polling for pending jobs: {}", e);
|
||||
}
|
||||
}
|
||||
|
||||
tokio::time::sleep(interval).await;
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
**Step 3: Add module to lib.rs**
|
||||
|
||||
Add `pub mod queue_consumer;` to `services/ml_training_service/src/lib.rs`.
|
||||
|
||||
**Step 4: Spawn consumer in main.rs**
|
||||
|
||||
In `main.rs`, after the service is constructed (line ~424), add:
|
||||
|
||||
```rust
|
||||
// Spawn queue consumer if both JobSpawner and K8sDispatcher are available
|
||||
if let Some(ref dispatcher) = k8s_dispatcher {
|
||||
let consumer_spawner = job_spawner.clone();
|
||||
let consumer_dispatcher = Arc::clone(dispatcher);
|
||||
let consumer_config = ml_training_service::queue_consumer::QueueConsumerConfig::default();
|
||||
|
||||
tokio::spawn(ml_training_service::queue_consumer::run_queue_consumer(
|
||||
consumer_spawner,
|
||||
consumer_dispatcher,
|
||||
consumer_config,
|
||||
));
|
||||
info!("Queue consumer started — polling for pending training jobs");
|
||||
}
|
||||
```
|
||||
|
||||
**Step 5: Run tests**
|
||||
|
||||
Run: `SQLX_OFFLINE=true cargo test -p ml_training_service --lib test_queue_consumer`
|
||||
Expected: PASS
|
||||
|
||||
**Step 6: Commit**
|
||||
|
||||
```bash
|
||||
git add services/ml_training_service/src/queue_consumer.rs services/ml_training_service/src/lib.rs services/ml_training_service/src/main.rs
|
||||
git commit -m "feat(training): add in-process queue consumer for K8s job dispatch"
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### Task 4: Expand child_jobs model_type constraint
|
||||
|
||||
**Files:**
|
||||
- Create: `migrations/049_expand_child_jobs_model_types.sql`
|
||||
|
||||
The existing `chk_child_job_model_type` constraint only allows 5 model types (DQN, PPO, MAMBA-2, TFT, TLOB). We support 10. Fix it.
|
||||
|
||||
**Step 1: Write migration**
|
||||
|
||||
```sql
|
||||
-- Migration 049: Expand child_jobs model_type constraint to all 10 supported models
|
||||
ALTER TABLE child_jobs DROP CONSTRAINT IF EXISTS chk_child_job_model_type;
|
||||
ALTER TABLE child_jobs ADD CONSTRAINT chk_child_job_model_type CHECK (
|
||||
model_type IN ('DQN', 'PPO', 'MAMBA-2', 'TFT', 'TLOB', 'TGGN', 'LIQUID', 'KAN', 'XLSTM', 'DIFFUSION')
|
||||
);
|
||||
```
|
||||
|
||||
**Step 2: Commit**
|
||||
|
||||
```bash
|
||||
git add migrations/049_expand_child_jobs_model_types.sql
|
||||
git commit -m "fix(db): expand child_jobs model_type constraint to all 10 models"
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### Task 5: Deploy training-output-pvc
|
||||
|
||||
**Files:**
|
||||
- Verify: `infra/k8s/training/training-output-pvc.yaml` (already exists)
|
||||
|
||||
**Step 1: Check if PVC exists in cluster**
|
||||
|
||||
```bash
|
||||
kubectl get pvc -n foxhunt training-output-pvc
|
||||
```
|
||||
|
||||
If missing, apply it:
|
||||
|
||||
```bash
|
||||
kubectl apply -f infra/k8s/training/training-output-pvc.yaml
|
||||
```
|
||||
|
||||
**Step 2: Verify both PVCs bound**
|
||||
|
||||
```bash
|
||||
kubectl get pvc -n foxhunt | grep training
|
||||
```
|
||||
|
||||
Expected: Both `training-data-pvc` and `training-output-pvc` in `Bound` state.
|
||||
|
||||
**Step 3: Commit (no code change needed if YAML exists)**
|
||||
|
||||
---
|
||||
|
||||
## Track 2: Model Serving & Validation
|
||||
|
||||
### Task 6: Wire report_job_completion to PromotionManager
|
||||
|
||||
**Files:**
|
||||
- Modify: `services/ml_training_service/src/service.rs:1097-1118` (report_job_completion handler)
|
||||
|
||||
**Step 1: Write failing test**
|
||||
|
||||
```rust
|
||||
#[tokio::test]
|
||||
async fn test_promotion_manager_register() {
|
||||
let pm = crate::promotion_manager::PromotionManager::new();
|
||||
let status = pm.register_completion(
|
||||
"test-job-1",
|
||||
"TFT",
|
||||
"ES.FUT",
|
||||
"s3://foxhunt-models/test/checkpoint.safetensors",
|
||||
&[("sharpe".to_string(), 1.5)].into_iter().collect(),
|
||||
).await;
|
||||
// First model for this (type, symbol) pair should be Registered
|
||||
assert_eq!(status, crate::promotion_manager::PromotionStatus::Registered);
|
||||
}
|
||||
```
|
||||
|
||||
**Step 2: Wire report_job_completion to PromotionManager**
|
||||
|
||||
Replace the stub implementation:
|
||||
|
||||
```rust
|
||||
async fn report_job_completion(
|
||||
&self,
|
||||
request: Request<JobCompletionReport>,
|
||||
) -> Result<Response<JobCompletionAck>, Status> {
|
||||
let report = request.into_inner();
|
||||
info!(
|
||||
job_id = %report.job_id,
|
||||
model_type = %report.model_type,
|
||||
success = report.success,
|
||||
"Job completion report received"
|
||||
);
|
||||
|
||||
if !report.success {
|
||||
// Update job status to Failed if we have a spawner
|
||||
if let Some(ref spawner) = self.job_spawner {
|
||||
if let Ok(uuid) = report.job_id.parse::<uuid::Uuid>() {
|
||||
let _ = spawner.update_job_status(uuid, "Failed").await;
|
||||
}
|
||||
}
|
||||
return Ok(Response::new(JobCompletionAck {
|
||||
acknowledged: true,
|
||||
promotion_status: "failed".to_string(),
|
||||
}));
|
||||
}
|
||||
|
||||
// Update job status to Completed
|
||||
if let Some(ref spawner) = self.job_spawner {
|
||||
if let Ok(uuid) = report.job_id.parse::<uuid::Uuid>() {
|
||||
let _ = spawner.update_job_status(uuid, "Completed").await;
|
||||
}
|
||||
}
|
||||
|
||||
// Register with promotion manager
|
||||
let metrics: std::collections::HashMap<String, f64> = report.metrics.clone();
|
||||
let status = self.promotion_manager.register_completion(
|
||||
&report.job_id,
|
||||
&report.model_type,
|
||||
&report.symbol,
|
||||
&report.s3_path,
|
||||
&metrics,
|
||||
).await;
|
||||
|
||||
let promotion_str = match status {
|
||||
crate::promotion_manager::PromotionStatus::PendingPromotion => "pending_promotion",
|
||||
crate::promotion_manager::PromotionStatus::NoImprovement => "no_improvement",
|
||||
crate::promotion_manager::PromotionStatus::Registered => "registered",
|
||||
crate::promotion_manager::PromotionStatus::Error(ref e) => {
|
||||
warn!("Promotion registration error: {}", e);
|
||||
"error"
|
||||
}
|
||||
};
|
||||
|
||||
Ok(Response::new(JobCompletionAck {
|
||||
acknowledged: true,
|
||||
promotion_status: promotion_str.to_string(),
|
||||
}))
|
||||
}
|
||||
```
|
||||
|
||||
**Step 3: Run tests**
|
||||
|
||||
Run: `SQLX_OFFLINE=true cargo check -p ml_training_service`
|
||||
|
||||
**Step 4: Commit**
|
||||
|
||||
```bash
|
||||
git add services/ml_training_service/src/service.rs
|
||||
git commit -m "feat(serving): wire report_job_completion to PromotionManager"
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### Task 7: Wire list_pending_promotions
|
||||
|
||||
**Files:**
|
||||
- Modify: `services/ml_training_service/src/service.rs:1120-1128`
|
||||
- Modify: `services/ml_training_service/src/promotion_manager.rs` (add list_pending method if needed)
|
||||
|
||||
**Step 1: Check if PromotionManager has list_pending**
|
||||
|
||||
Read `promotion_manager.rs` to verify `list_pending()` exists. If not, add:
|
||||
|
||||
```rust
|
||||
/// List all models pending promotion.
|
||||
pub async fn list_pending(&self) -> Vec<PendingModel> {
|
||||
self.pending.read().await.values().cloned().collect()
|
||||
}
|
||||
```
|
||||
|
||||
**Step 2: Wire the gRPC handler**
|
||||
|
||||
```rust
|
||||
async fn list_pending_promotions(
|
||||
&self,
|
||||
_request: Request<ListPendingPromotionsRequest>,
|
||||
) -> Result<Response<ListPendingPromotionsResponse>, Status> {
|
||||
let pending = self.promotion_manager.list_pending().await;
|
||||
|
||||
let promotions = pending.iter().map(|p| proto::PendingPromotion {
|
||||
model_id: p.model_id.clone(),
|
||||
model_type: p.model_type.clone(),
|
||||
symbol: p.symbol.clone(),
|
||||
s3_path: p.s3_path.clone(),
|
||||
job_id: p.job_id.clone(),
|
||||
new_metrics: p.new_metrics.clone(),
|
||||
current_metrics: p.current_metrics.clone(),
|
||||
}).collect();
|
||||
|
||||
Ok(Response::new(ListPendingPromotionsResponse { promotions }))
|
||||
}
|
||||
```
|
||||
|
||||
**Step 3: Check proto definition has PendingPromotion message**
|
||||
|
||||
If the proto doesn't have the `PendingPromotion` message, this will need adjustment to match whatever fields the proto actually defines. Check the generated code.
|
||||
|
||||
**Step 4: Commit**
|
||||
|
||||
```bash
|
||||
git add services/ml_training_service/src/service.rs services/ml_training_service/src/promotion_manager.rs
|
||||
git commit -m "feat(serving): wire list_pending_promotions to PromotionManager"
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### Task 8: Wire approve_promotion and reject_promotion
|
||||
|
||||
**Files:**
|
||||
- Modify: `services/ml_training_service/src/service.rs:1130-1155`
|
||||
- Modify: `services/ml_training_service/src/promotion_manager.rs` (add approve/reject methods)
|
||||
|
||||
**Step 1: Add approve/reject to PromotionManager**
|
||||
|
||||
```rust
|
||||
/// Approve a pending promotion — moves model to active.
|
||||
pub async fn approve(&self, model_id: &str) -> Result<ActiveModel, String> {
|
||||
let mut pending = self.pending.write().await;
|
||||
let model = pending.remove(model_id)
|
||||
.ok_or_else(|| format!("No pending model: {}", model_id))?;
|
||||
|
||||
let active = ActiveModel {
|
||||
model_id: model.model_id.clone(),
|
||||
s3_path: model.s3_path.clone(),
|
||||
metrics: model.new_metrics.clone(),
|
||||
promoted_at: chrono::Utc::now(),
|
||||
};
|
||||
|
||||
let key = (model.model_type.clone(), model.symbol.clone());
|
||||
self.active_models.write().await.insert(key, active.clone());
|
||||
|
||||
info!(model_id = %model_id, "Model promoted to active");
|
||||
Ok(active)
|
||||
}
|
||||
|
||||
/// Reject a pending promotion — removes from pending list.
|
||||
pub async fn reject(&self, model_id: &str, reason: &str) -> Result<(), String> {
|
||||
let mut pending = self.pending.write().await;
|
||||
if pending.remove(model_id).is_none() {
|
||||
return Err(format!("No pending model: {}", model_id));
|
||||
}
|
||||
info!(model_id = %model_id, reason = %reason, "Model promotion rejected");
|
||||
Ok(())
|
||||
}
|
||||
```
|
||||
|
||||
**Step 2: Wire gRPC handlers**
|
||||
|
||||
```rust
|
||||
async fn approve_promotion(
|
||||
&self,
|
||||
request: Request<ApprovePromotionRequest>,
|
||||
) -> Result<Response<ApprovePromotionResponse>, Status> {
|
||||
let req = request.into_inner();
|
||||
info!(model_id = %req.model_id, "Promotion approval requested");
|
||||
|
||||
match self.promotion_manager.approve(&req.model_id).await {
|
||||
Ok(active) => Ok(Response::new(ApprovePromotionResponse {
|
||||
success: true,
|
||||
message: format!("Model {} promoted. S3: {}", active.model_id, active.s3_path),
|
||||
})),
|
||||
Err(e) => Ok(Response::new(ApprovePromotionResponse {
|
||||
success: false,
|
||||
message: e,
|
||||
})),
|
||||
}
|
||||
}
|
||||
|
||||
async fn reject_promotion(
|
||||
&self,
|
||||
request: Request<RejectPromotionRequest>,
|
||||
) -> Result<Response<RejectPromotionResponse>, Status> {
|
||||
let req = request.into_inner();
|
||||
info!(model_id = %req.model_id, reason = %req.reason, "Promotion rejection requested");
|
||||
|
||||
match self.promotion_manager.reject(&req.model_id, &req.reason).await {
|
||||
Ok(()) => Ok(Response::new(RejectPromotionResponse {
|
||||
success: true,
|
||||
message: format!("Model {} promotion rejected", req.model_id),
|
||||
})),
|
||||
Err(e) => Ok(Response::new(RejectPromotionResponse {
|
||||
success: false,
|
||||
message: e,
|
||||
})),
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
**Step 3: Run tests**
|
||||
|
||||
Run: `SQLX_OFFLINE=true cargo check -p ml_training_service`
|
||||
|
||||
**Step 4: Commit**
|
||||
|
||||
```bash
|
||||
git add services/ml_training_service/src/service.rs services/ml_training_service/src/promotion_manager.rs
|
||||
git commit -m "feat(serving): wire approve/reject promotion to PromotionManager"
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### Task 9: Verify fxt model commands work end-to-end
|
||||
|
||||
**Files:**
|
||||
- Read: `bin/fxt/src/commands/model/list.rs`
|
||||
- Read: `bin/fxt/src/commands/model/approve.rs`
|
||||
- Read: `bin/fxt/src/commands/model/reject.rs`
|
||||
|
||||
**Step 1: Verify fxt model commands compile**
|
||||
|
||||
Run: `SQLX_OFFLINE=true cargo check -p fxt`
|
||||
|
||||
**Step 2: Test locally (if training service is running)**
|
||||
|
||||
```bash
|
||||
fxt model list
|
||||
fxt model approve <model-id>
|
||||
fxt model reject <model-id> --reason "test"
|
||||
```
|
||||
|
||||
**Step 3: Commit (if any fixes needed)**
|
||||
|
||||
---
|
||||
|
||||
### Task 10: Build and push Docker images
|
||||
|
||||
**Files:**
|
||||
- Read: `infra/docker/Dockerfile.training` (training binaries image)
|
||||
- Read: K8s deployment for ml-training-service
|
||||
|
||||
**Step 1: Build training image**
|
||||
|
||||
```bash
|
||||
docker build -f infra/docker/Dockerfile.training -t rg.fr-par.scw.cloud/foxhunt-ci/training:latest .
|
||||
docker push rg.fr-par.scw.cloud/foxhunt-ci/training:latest
|
||||
```
|
||||
|
||||
**Step 2: Build and push ml-training-service image**
|
||||
|
||||
```bash
|
||||
# This would use the service Dockerfile
|
||||
docker build -f infra/docker/Dockerfile.ml-training-service -t rg.fr-par.scw.cloud/foxhunt-ci/ml-training-service:latest .
|
||||
docker push rg.fr-par.scw.cloud/foxhunt-ci/ml-training-service:latest
|
||||
```
|
||||
|
||||
**Step 3: Redeploy to K8s**
|
||||
|
||||
```bash
|
||||
kubectl rollout restart deployment/ml-training-service -n foxhunt
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### Task 11: End-to-end validation
|
||||
|
||||
**Step 1: Verify training service is running**
|
||||
|
||||
```bash
|
||||
kubectl get pods -n foxhunt -l app=ml-training-service
|
||||
kubectl logs -n foxhunt deployment/ml-training-service --tail=50
|
||||
```
|
||||
|
||||
Look for: "Queue consumer started" and "K8s dispatcher initialized"
|
||||
|
||||
**Step 2: Trigger a training job**
|
||||
|
||||
```bash
|
||||
fxt train start tft ES.FUT --epochs 10
|
||||
```
|
||||
|
||||
Expected: Job ID returned, status PENDING
|
||||
|
||||
**Step 3: Check job was dispatched**
|
||||
|
||||
```bash
|
||||
fxt train status <job-id>
|
||||
kubectl get jobs -n foxhunt -l foxhunt/job-type=training
|
||||
```
|
||||
|
||||
Expected: K8s Job created on gpu-training pool
|
||||
|
||||
**Step 4: Check model promotion after completion**
|
||||
|
||||
```bash
|
||||
fxt model list
|
||||
```
|
||||
|
||||
Expected: New model appears as pending (or registered if first for this type/symbol)
|
||||
@@ -8,8 +8,8 @@ metadata:
|
||||
app.kubernetes.io/part-of: foxhunt
|
||||
spec:
|
||||
accessModes:
|
||||
- ReadWriteMany
|
||||
- ReadWriteOnce
|
||||
resources:
|
||||
requests:
|
||||
storage: 50Gi
|
||||
storageClassName: scw-bssd-nfs
|
||||
storageClassName: scw-bssd
|
||||
|
||||
7
migrations/049_expand_child_jobs_model_types.sql
Normal file
7
migrations/049_expand_child_jobs_model_types.sql
Normal file
@@ -0,0 +1,7 @@
|
||||
-- Migration 049: Expand child_jobs model_type constraint to all 10 supported models
|
||||
-- The original constraint (migration 046) only allowed DQN, PPO, MAMBA-2, TFT, TLOB.
|
||||
-- We now support 10 model types via UnifiedTrainable adapters.
|
||||
ALTER TABLE child_jobs DROP CONSTRAINT IF EXISTS chk_child_job_model_type;
|
||||
ALTER TABLE child_jobs ADD CONSTRAINT chk_child_job_model_type CHECK (
|
||||
model_type IN ('DQN', 'PPO', 'MAMBA-2', 'TFT', 'TLOB', 'TGGN', 'LIQUID', 'KAN', 'XLSTM', 'DIFFUSION')
|
||||
);
|
||||
@@ -139,6 +139,13 @@ impl DatabaseManager {
|
||||
Ok(manager)
|
||||
}
|
||||
|
||||
/// Get a clone of the underlying `PgPool` for use by other components (e.g. `JobSpawner`).
|
||||
///
|
||||
/// `PgPool` is an `Arc`-wrapped handle, so cloning is cheap.
|
||||
pub fn pg_pool(&self) -> sqlx::PgPool {
|
||||
self.db_pool.pool().clone()
|
||||
}
|
||||
|
||||
/// Run database migrations
|
||||
pub async fn run_migrations(&self) -> Result<()> {
|
||||
info!("Running database migrations");
|
||||
|
||||
@@ -503,6 +503,33 @@ impl JobSpawner {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Fetch a single child job by its ID
|
||||
///
|
||||
/// # Arguments
|
||||
///
|
||||
/// * `job_id` - Child job identifier
|
||||
///
|
||||
/// # Returns
|
||||
///
|
||||
/// * `Ok(Some(ChildJob))` - The child job
|
||||
/// * `Ok(None)` - Job not found
|
||||
/// * `Err` - Database error
|
||||
pub async fn get_job_by_id(&self, job_id: Uuid) -> Result<Option<ChildJob>> {
|
||||
let job = sqlx::query_as::<_, ChildJob>(
|
||||
r#"
|
||||
SELECT id, batch_id, model_type, status, created_at, config_json
|
||||
FROM child_jobs
|
||||
WHERE id = $1
|
||||
"#,
|
||||
)
|
||||
.bind(job_id)
|
||||
.fetch_optional(&self.db_pool)
|
||||
.await
|
||||
.context("Failed to fetch job by id")?;
|
||||
|
||||
Ok(job)
|
||||
}
|
||||
|
||||
/// Get all child jobs for a batch
|
||||
///
|
||||
/// # Arguments
|
||||
|
||||
@@ -38,6 +38,7 @@ pub mod monitoring;
|
||||
pub mod optuna_persistence;
|
||||
pub mod orchestrator;
|
||||
pub mod promotion_manager;
|
||||
pub mod queue_consumer;
|
||||
pub mod schema_types;
|
||||
pub mod service;
|
||||
pub mod simple_metrics;
|
||||
|
||||
@@ -414,6 +414,26 @@ async fn serve(args: ServeArgs) -> Result<()> {
|
||||
}
|
||||
};
|
||||
|
||||
// Initialize JobSpawner for persisting training jobs to PostgreSQL
|
||||
let job_spawner = Arc::new(ml_training_service::job_spawner::JobSpawner::new(
|
||||
database.pg_pool(),
|
||||
));
|
||||
info!("Job spawner initialized -- training jobs will be persisted to PostgreSQL");
|
||||
|
||||
// Spawn queue consumer to poll pending jobs and dispatch to K8s
|
||||
if let Some(ref dispatcher) = k8s_dispatcher {
|
||||
let consumer_spawner = Arc::clone(&job_spawner);
|
||||
let consumer_dispatcher = Arc::clone(dispatcher);
|
||||
let consumer_config =
|
||||
ml_training_service::queue_consumer::QueueConsumerConfig::default();
|
||||
tokio::spawn(ml_training_service::queue_consumer::run_queue_consumer(
|
||||
consumer_spawner,
|
||||
consumer_dispatcher,
|
||||
consumer_config,
|
||||
));
|
||||
info!("Queue consumer started -- polling for pending training jobs");
|
||||
}
|
||||
|
||||
// Create gRPC service
|
||||
let training_service = MLTrainingServiceImpl::new(
|
||||
Arc::clone(&orchestrator),
|
||||
@@ -421,6 +441,7 @@ async fn serve(args: ServeArgs) -> Result<()> {
|
||||
ml_config.clone(),
|
||||
promotion_manager,
|
||||
k8s_dispatcher,
|
||||
Some(job_spawner),
|
||||
);
|
||||
|
||||
// Build server with reflection
|
||||
|
||||
@@ -183,6 +183,12 @@ impl PromotionManager {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Expose the active_models map for test assertions (test-only).
|
||||
#[cfg(test)]
|
||||
pub fn active_models_for_test(&self) -> Arc<RwLock<HashMap<(String, String), ActiveModel>>> {
|
||||
Arc::clone(&self.active_models)
|
||||
}
|
||||
|
||||
/// Compare new metrics against current active metrics.
|
||||
///
|
||||
/// Returns true if the new model is better:
|
||||
|
||||
181
services/ml_training_service/src/queue_consumer.rs
Normal file
181
services/ml_training_service/src/queue_consumer.rs
Normal file
@@ -0,0 +1,181 @@
|
||||
//! Queue Consumer — background loop that polls pending training jobs and dispatches them to K8s.
|
||||
//!
|
||||
//! Runs as a `tokio::spawn`-ed task inside the `ml_training_service` process. It checks
|
||||
//! [`JobSpawner::get_next_pending_job`] every `poll_interval_secs` seconds and, when a job
|
||||
//! is found, marks it as `Running`, builds [`TrainingJobParams`], and calls
|
||||
//! [`K8sDispatcher::dispatch`]. On dispatch failure the job is marked `Failed`.
|
||||
//!
|
||||
//! The consumer is intentionally simple — one job at a time — because GPU training jobs
|
||||
//! are heavyweight and we rely on K8s scheduling for parallelism.
|
||||
|
||||
use std::sync::Arc;
|
||||
use std::time::Duration;
|
||||
|
||||
use tracing::{debug, error, info, warn};
|
||||
|
||||
use crate::job_spawner::JobSpawner;
|
||||
use crate::k8s_dispatcher::{training_binary_for_model, K8sDispatcher, TrainingJobParams};
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Configuration
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// Configuration for the queue consumer loop.
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct QueueConsumerConfig {
|
||||
/// How often (in seconds) to poll the database for pending jobs.
|
||||
pub poll_interval_secs: u64,
|
||||
/// Maximum number of concurrent K8s jobs the consumer will keep in flight.
|
||||
/// Currently only `1` is supported — the field exists for forward compatibility.
|
||||
pub max_concurrent_jobs: u32,
|
||||
/// Default number of training epochs when the job config does not specify one.
|
||||
pub default_epochs: u32,
|
||||
}
|
||||
|
||||
impl Default for QueueConsumerConfig {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
poll_interval_secs: 5,
|
||||
max_concurrent_jobs: 1,
|
||||
default_epochs: 50,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Consumer loop
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// Runs the queue consumer forever.
|
||||
///
|
||||
/// This function is designed to be passed to [`tokio::spawn`]. It never returns under
|
||||
/// normal operation — it only exits if the `JobSpawner` or `K8sDispatcher` is dropped
|
||||
/// (which causes DB / K8s calls to fail permanently).
|
||||
pub async fn run_queue_consumer(
|
||||
spawner: Arc<JobSpawner>,
|
||||
dispatcher: Arc<K8sDispatcher>,
|
||||
config: QueueConsumerConfig,
|
||||
) {
|
||||
let poll_interval = Duration::from_secs(config.poll_interval_secs);
|
||||
|
||||
info!(
|
||||
poll_interval_secs = config.poll_interval_secs,
|
||||
max_concurrent_jobs = config.max_concurrent_jobs,
|
||||
"queue consumer started"
|
||||
);
|
||||
|
||||
loop {
|
||||
match spawner.get_next_pending_job().await {
|
||||
Ok(Some(job)) => {
|
||||
let job_id = job.id;
|
||||
let model_type = job.model_type.clone();
|
||||
|
||||
info!(
|
||||
job_id = %job_id,
|
||||
model_type = %model_type,
|
||||
batch_id = %job.batch_id,
|
||||
"picked up pending job"
|
||||
);
|
||||
|
||||
// Mark as Running *before* dispatching so no other consumer grabs it.
|
||||
if let Err(e) = spawner.update_job_status(job_id, "Running").await {
|
||||
warn!(
|
||||
job_id = %job_id,
|
||||
error = %e,
|
||||
"failed to mark job as Running — skipping"
|
||||
);
|
||||
tokio::time::sleep(poll_interval).await;
|
||||
continue;
|
||||
}
|
||||
|
||||
// Extract parameters from config_json.
|
||||
let symbol = job
|
||||
.config_json
|
||||
.get("asset")
|
||||
.and_then(|v| v.as_str())
|
||||
.unwrap_or("ES.FUT")
|
||||
.to_string();
|
||||
|
||||
let epochs = job
|
||||
.config_json
|
||||
.get("epochs")
|
||||
.and_then(|v| v.as_u64())
|
||||
.map(|e| e as u32)
|
||||
.unwrap_or(config.default_epochs);
|
||||
|
||||
let binary = training_binary_for_model(&model_type).to_string();
|
||||
|
||||
let params = TrainingJobParams {
|
||||
job_id,
|
||||
model_type: model_type.clone(),
|
||||
symbol,
|
||||
epochs,
|
||||
binary,
|
||||
};
|
||||
|
||||
match dispatcher.dispatch(¶ms).await {
|
||||
Ok(k8s_name) => {
|
||||
info!(
|
||||
job_id = %job_id,
|
||||
k8s_name = %k8s_name,
|
||||
"dispatched training job to K8s"
|
||||
);
|
||||
}
|
||||
Err(e) => {
|
||||
error!(
|
||||
job_id = %job_id,
|
||||
error = %e,
|
||||
"K8s dispatch failed — marking job as Failed"
|
||||
);
|
||||
if let Err(update_err) =
|
||||
spawner.update_job_status(job_id, "Failed").await
|
||||
{
|
||||
error!(
|
||||
job_id = %job_id,
|
||||
error = %update_err,
|
||||
"failed to mark job as Failed after dispatch error"
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
Ok(None) => {
|
||||
debug!("no pending jobs — sleeping");
|
||||
}
|
||||
Err(e) => {
|
||||
warn!(error = %e, "error polling for pending jobs — will retry");
|
||||
}
|
||||
}
|
||||
|
||||
tokio::time::sleep(poll_interval).await;
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Tests
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_queue_consumer_config_defaults() {
|
||||
let config = QueueConsumerConfig::default();
|
||||
assert_eq!(config.poll_interval_secs, 5);
|
||||
assert_eq!(config.max_concurrent_jobs, 1);
|
||||
assert_eq!(config.default_epochs, 50);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_queue_consumer_config_custom() {
|
||||
let config = QueueConsumerConfig {
|
||||
poll_interval_secs: 10,
|
||||
max_concurrent_jobs: 3,
|
||||
default_epochs: 100,
|
||||
};
|
||||
assert_eq!(config.poll_interval_secs, 10);
|
||||
assert_eq!(config.max_concurrent_jobs, 3);
|
||||
assert_eq!(config.default_epochs, 100);
|
||||
}
|
||||
}
|
||||
@@ -35,7 +35,7 @@ use proto::{
|
||||
TrainingStatusUpdate as ProtoStatusUpdate,
|
||||
// Job completion & model promotion (on-demand training pipeline)
|
||||
JobCompletionReport, JobCompletionAck,
|
||||
ListPendingPromotionsRequest, ListPendingPromotionsResponse,
|
||||
ListPendingPromotionsRequest, ListPendingPromotionsResponse, PendingPromotion,
|
||||
ApprovePromotionRequest, ApprovePromotionResponse,
|
||||
RejectPromotionRequest, RejectPromotionResponse,
|
||||
};
|
||||
@@ -58,6 +58,9 @@ pub struct MLTrainingServiceImpl {
|
||||
pub promotion_manager: Arc<crate::promotion_manager::PromotionManager>,
|
||||
/// K8s job dispatcher -- None when running outside a K8s cluster.
|
||||
pub k8s_dispatcher: Option<Arc<crate::k8s_dispatcher::K8sDispatcher>>,
|
||||
/// Job spawner for persisting training jobs to PostgreSQL before K8s dispatch.
|
||||
/// None when no database pool is available.
|
||||
pub job_spawner: Option<Arc<crate::job_spawner::JobSpawner>>,
|
||||
}
|
||||
|
||||
impl MLTrainingServiceImpl {
|
||||
@@ -68,6 +71,7 @@ impl MLTrainingServiceImpl {
|
||||
_config: MLConfig,
|
||||
promotion_manager: Arc<crate::promotion_manager::PromotionManager>,
|
||||
k8s_dispatcher: Option<Arc<crate::k8s_dispatcher::K8sDispatcher>>,
|
||||
job_spawner: Option<Arc<crate::job_spawner::JobSpawner>>,
|
||||
) -> Self {
|
||||
let tuning_handlers = Arc::new(TuningHandlers::new(Arc::clone(&tuning_manager)));
|
||||
let working_dir = std::env::var("TUNING_WORKING_DIR").unwrap_or_else(|_| ".".to_string());
|
||||
@@ -78,6 +82,7 @@ impl MLTrainingServiceImpl {
|
||||
batch_tuning_manager,
|
||||
promotion_manager,
|
||||
k8s_dispatcher,
|
||||
job_spawner,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -261,18 +266,48 @@ impl MlTrainingService for MLTrainingServiceImpl {
|
||||
// Generate a job ID
|
||||
let job_id = Uuid::new_v4();
|
||||
|
||||
// Try K8s dispatch first, fall back to in-process orchestrator
|
||||
if let Some(ref dispatcher) = self.k8s_dispatcher {
|
||||
// Extract symbol from data_source file_path, or use default
|
||||
let symbol = req
|
||||
.data_source
|
||||
.as_ref()
|
||||
.and_then(|ds| match &ds.source {
|
||||
Some(proto::data_source::Source::FilePath(p)) => Some(p.clone()),
|
||||
_ => None,
|
||||
})
|
||||
.unwrap_or_else(|| "ES.FUT".to_string());
|
||||
// Extract symbol from data_source file_path, or use default
|
||||
let symbol = req
|
||||
.data_source
|
||||
.as_ref()
|
||||
.and_then(|ds| match &ds.source {
|
||||
Some(proto::data_source::Source::FilePath(p)) => Some(p.clone()),
|
||||
_ => None,
|
||||
})
|
||||
.unwrap_or_else(|| "ES.FUT".to_string());
|
||||
|
||||
// Priority 1: Batch flow — persist job to DB, queue consumer dispatches to K8s
|
||||
if let Some(ref spawner) = self.job_spawner {
|
||||
let model = common::model_types::ModelType::from_str(&req.model_type).ok_or_else(
|
||||
|| Status::invalid_argument(format!("Unknown model type: {}", req.model_type)),
|
||||
)?;
|
||||
|
||||
let asset = crate::job_spawner::Asset {
|
||||
symbol: symbol.clone(),
|
||||
data_file: std::path::PathBuf::from(format!(
|
||||
"/data/futures-baseline/{}",
|
||||
symbol
|
||||
)),
|
||||
};
|
||||
|
||||
match spawner.spawn_batch(vec![asset], vec![model]).await {
|
||||
Ok(batch) => {
|
||||
info!(batch_id = %batch.batch_id, "Batch job created in database");
|
||||
return Ok(Response::new(StartTrainingResponse {
|
||||
job_id: batch.batch_id.to_string(),
|
||||
status: ProtoTrainingStatus::Pending as i32,
|
||||
message: format!("Batch training job queued: {}", batch.batch_id),
|
||||
}));
|
||||
}
|
||||
Err(e) => {
|
||||
warn!("JobSpawner failed, falling back to K8s/orchestrator: {}", e);
|
||||
// Fall through to K8s dispatcher or in-process orchestrator
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Priority 2: Direct K8s dispatch
|
||||
if let Some(ref dispatcher) = self.k8s_dispatcher {
|
||||
let epochs = req
|
||||
.hyperparameters
|
||||
.as_ref()
|
||||
@@ -1106,14 +1141,92 @@ impl MlTrainingService for MLTrainingServiceImpl {
|
||||
"received job completion report"
|
||||
);
|
||||
|
||||
// TODO(task-6): wire to promotion pipeline
|
||||
// Parse the job UUID from the report
|
||||
let job_uuid = uuid::Uuid::parse_str(&report.job_id).map_err(|e| {
|
||||
Status::invalid_argument(format!("invalid job_id UUID: {}", e))
|
||||
})?;
|
||||
|
||||
// ── Failed job path ──────────────────────────────────────────────
|
||||
if !report.success {
|
||||
// Best-effort DB status update; log warning if it fails
|
||||
if let Some(ref spawner) = self.job_spawner {
|
||||
if let Err(e) = spawner.update_job_status(job_uuid, "Failed").await {
|
||||
warn!(job_id = %report.job_id, error = %e, "failed to update job status to Failed");
|
||||
}
|
||||
}
|
||||
return Ok(Response::new(JobCompletionAck {
|
||||
accepted: true,
|
||||
promotion_status: "failed".to_string(),
|
||||
}));
|
||||
}
|
||||
|
||||
// ── Successful job path ──────────────────────────────────────────
|
||||
// 1. Update DB status to Completed
|
||||
if let Some(ref spawner) = self.job_spawner {
|
||||
if let Err(e) = spawner.update_job_status(job_uuid, "Completed").await {
|
||||
warn!(job_id = %report.job_id, error = %e, "failed to update job status to Completed");
|
||||
}
|
||||
}
|
||||
|
||||
// 2. Look up model_type and symbol from the child_jobs table
|
||||
let (model_type, symbol) = match self.job_spawner {
|
||||
Some(ref spawner) => {
|
||||
match spawner.get_job_by_id(job_uuid).await {
|
||||
Ok(Some(child_job)) => {
|
||||
// Symbol is stored in config_json.asset
|
||||
let sym = child_job
|
||||
.config_json
|
||||
.get("asset")
|
||||
.and_then(|v| v.as_str())
|
||||
.unwrap_or("UNKNOWN")
|
||||
.to_string();
|
||||
(child_job.model_type, sym)
|
||||
}
|
||||
Ok(None) => {
|
||||
warn!(job_id = %report.job_id, "child job not found in DB, using defaults");
|
||||
("UNKNOWN".to_string(), "UNKNOWN".to_string())
|
||||
}
|
||||
Err(e) => {
|
||||
warn!(job_id = %report.job_id, error = %e, "failed to look up child job");
|
||||
("UNKNOWN".to_string(), "UNKNOWN".to_string())
|
||||
}
|
||||
}
|
||||
}
|
||||
None => {
|
||||
// No DB — cannot resolve model_type/symbol
|
||||
debug!(job_id = %report.job_id, "no JobSpawner, cannot resolve model metadata");
|
||||
("UNKNOWN".to_string(), "UNKNOWN".to_string())
|
||||
}
|
||||
};
|
||||
|
||||
// 3. Register with PromotionManager
|
||||
let metrics: std::collections::HashMap<String, f64> = report.metrics.into_iter().collect();
|
||||
let status = self
|
||||
.promotion_manager
|
||||
.register_completion(&report.job_id, &model_type, &symbol, &report.s3_path, metrics)
|
||||
.await;
|
||||
|
||||
let promotion_status = match status {
|
||||
crate::promotion_manager::PromotionStatus::PendingPromotion => "pending_promotion",
|
||||
crate::promotion_manager::PromotionStatus::NoImprovement => "no_improvement",
|
||||
crate::promotion_manager::PromotionStatus::Registered => "registered",
|
||||
crate::promotion_manager::PromotionStatus::Error(ref msg) => {
|
||||
warn!(job_id = %report.job_id, error = %msg, "promotion registration error");
|
||||
"error"
|
||||
}
|
||||
};
|
||||
|
||||
info!(
|
||||
job_id = %report.job_id,
|
||||
model_type = %model_type,
|
||||
symbol = %symbol,
|
||||
promotion_status = %promotion_status,
|
||||
"job completion processed"
|
||||
);
|
||||
|
||||
Ok(Response::new(JobCompletionAck {
|
||||
accepted: true,
|
||||
promotion_status: if report.success {
|
||||
"pending_promotion".to_string()
|
||||
} else {
|
||||
"failed".to_string()
|
||||
},
|
||||
promotion_status: promotion_status.to_string(),
|
||||
}))
|
||||
}
|
||||
|
||||
@@ -1121,10 +1234,24 @@ impl MlTrainingService for MLTrainingServiceImpl {
|
||||
&self,
|
||||
_request: Request<ListPendingPromotionsRequest>,
|
||||
) -> Result<Response<ListPendingPromotionsResponse>, Status> {
|
||||
// TODO(task-6): query promotion store
|
||||
Ok(Response::new(ListPendingPromotionsResponse {
|
||||
promotions: vec![],
|
||||
}))
|
||||
let pending = self.promotion_manager.list_pending().await;
|
||||
info!(count = pending.len(), "listing pending promotions");
|
||||
|
||||
let promotions = pending
|
||||
.into_iter()
|
||||
.map(|m| PendingPromotion {
|
||||
model_id: m.model_id,
|
||||
model_type: m.model_type,
|
||||
symbol: m.symbol,
|
||||
s3_path: m.s3_path,
|
||||
new_metrics: m.new_metrics.into_iter().collect(),
|
||||
current_metrics: m.current_metrics.into_iter().collect(),
|
||||
trained_at: m.trained_at.timestamp(),
|
||||
job_id: m.job_id,
|
||||
})
|
||||
.collect();
|
||||
|
||||
Ok(Response::new(ListPendingPromotionsResponse { promotions }))
|
||||
}
|
||||
|
||||
async fn approve_promotion(
|
||||
@@ -1132,12 +1259,27 @@ impl MlTrainingService for MLTrainingServiceImpl {
|
||||
request: Request<ApprovePromotionRequest>,
|
||||
) -> Result<Response<ApprovePromotionResponse>, Status> {
|
||||
let req = request.into_inner();
|
||||
info!(model_id = %req.model_id, "promotion approved");
|
||||
// TODO(task-6): execute promotion (swap live model)
|
||||
Ok(Response::new(ApprovePromotionResponse {
|
||||
success: true,
|
||||
message: format!("Promotion of model {} approved", req.model_id),
|
||||
}))
|
||||
if req.model_id.is_empty() {
|
||||
return Err(Status::invalid_argument("model_id is required"));
|
||||
}
|
||||
info!(model_id = %req.model_id, "approve_promotion request received");
|
||||
|
||||
match self.promotion_manager.approve(&req.model_id).await {
|
||||
Ok(()) => {
|
||||
info!(model_id = %req.model_id, "model promotion approved and activated");
|
||||
Ok(Response::new(ApprovePromotionResponse {
|
||||
success: true,
|
||||
message: format!("Model {} promoted to active", req.model_id),
|
||||
}))
|
||||
}
|
||||
Err(e) => {
|
||||
warn!(model_id = %req.model_id, error = %e, "approve_promotion failed");
|
||||
Ok(Response::new(ApprovePromotionResponse {
|
||||
success: false,
|
||||
message: format!("Failed to approve promotion: {}", e),
|
||||
}))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn reject_promotion(
|
||||
@@ -1145,12 +1287,32 @@ impl MlTrainingService for MLTrainingServiceImpl {
|
||||
request: Request<RejectPromotionRequest>,
|
||||
) -> Result<Response<RejectPromotionResponse>, Status> {
|
||||
let req = request.into_inner();
|
||||
info!(model_id = %req.model_id, reason = %req.reason, "promotion rejected");
|
||||
// TODO(task-6): mark promotion as rejected
|
||||
Ok(Response::new(RejectPromotionResponse {
|
||||
success: true,
|
||||
message: format!("Promotion of model {} rejected: {}", req.model_id, req.reason),
|
||||
}))
|
||||
if req.model_id.is_empty() {
|
||||
return Err(Status::invalid_argument("model_id is required"));
|
||||
}
|
||||
let reason = if req.reason.is_empty() {
|
||||
"no reason provided".to_string()
|
||||
} else {
|
||||
req.reason
|
||||
};
|
||||
info!(model_id = %req.model_id, reason = %reason, "reject_promotion request received");
|
||||
|
||||
match self.promotion_manager.reject(&req.model_id, &reason).await {
|
||||
Ok(()) => {
|
||||
info!(model_id = %req.model_id, reason = %reason, "model promotion rejected");
|
||||
Ok(Response::new(RejectPromotionResponse {
|
||||
success: true,
|
||||
message: format!("Promotion of model {} rejected: {}", req.model_id, reason),
|
||||
}))
|
||||
}
|
||||
Err(e) => {
|
||||
warn!(model_id = %req.model_id, error = %e, "reject_promotion failed");
|
||||
Ok(Response::new(RejectPromotionResponse {
|
||||
success: false,
|
||||
message: format!("Failed to reject promotion: {}", e),
|
||||
}))
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1582,4 +1744,259 @@ mod tests {
|
||||
assert!(config.training_params.validation_split >= 0.0);
|
||||
assert!(config.training_params.validation_split <= 1.0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_service_impl_has_job_spawner() {
|
||||
let _: fn(&MLTrainingServiceImpl) -> bool = |s| s.job_spawner.is_some();
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_start_training_model_binary_mapping() {
|
||||
use crate::k8s_dispatcher::training_binary_for_model;
|
||||
assert_eq!(training_binary_for_model("tft"), "train_baseline_supervised");
|
||||
assert_eq!(training_binary_for_model("dqn"), "train_baseline_rl");
|
||||
assert_eq!(training_binary_for_model("ppo"), "train_baseline_rl");
|
||||
assert_eq!(
|
||||
training_binary_for_model("mamba2"),
|
||||
"train_baseline_supervised"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_promotion_status_string_mapping() {
|
||||
// Verify the string representations match what gRPC clients expect
|
||||
use crate::promotion_manager::PromotionStatus;
|
||||
|
||||
let cases: Vec<(PromotionStatus, &str)> = vec![
|
||||
(PromotionStatus::PendingPromotion, "pending_promotion"),
|
||||
(PromotionStatus::NoImprovement, "no_improvement"),
|
||||
(PromotionStatus::Registered, "registered"),
|
||||
(PromotionStatus::Error("test".to_string()), "error"),
|
||||
];
|
||||
|
||||
for (status, expected) in cases {
|
||||
let result = match status {
|
||||
PromotionStatus::PendingPromotion => "pending_promotion",
|
||||
PromotionStatus::NoImprovement => "no_improvement",
|
||||
PromotionStatus::Registered => "registered",
|
||||
PromotionStatus::Error(_) => "error",
|
||||
};
|
||||
assert_eq!(result, expected);
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_report_job_completion_registers_with_promotion_manager() {
|
||||
// Integration test: verify that a successful completion report
|
||||
// flows through to the PromotionManager and returns the correct status
|
||||
use crate::promotion_manager::PromotionManager;
|
||||
|
||||
let pm = PromotionManager::new();
|
||||
let metrics = HashMap::from([
|
||||
("best_val_loss".to_string(), 0.05),
|
||||
("sharpe_ratio".to_string(), 1.5),
|
||||
]);
|
||||
|
||||
// First registration for a type+symbol returns Registered
|
||||
let status = pm
|
||||
.register_completion("job-1", "DQN", "ES.FUT", "s3://b/m.bin", metrics)
|
||||
.await;
|
||||
assert_eq!(
|
||||
status,
|
||||
crate::promotion_manager::PromotionStatus::Registered
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_list_pending_promotions_maps_to_proto() {
|
||||
// Verify that PendingModel fields map correctly to proto PendingPromotion
|
||||
use crate::promotion_manager::PendingModel;
|
||||
use chrono::Utc;
|
||||
|
||||
let now = Utc::now();
|
||||
let model = PendingModel {
|
||||
model_id: "m-123".to_string(),
|
||||
model_type: "DQN".to_string(),
|
||||
symbol: "ES.FUT".to_string(),
|
||||
s3_path: "s3://bucket/better.bin".to_string(),
|
||||
new_metrics: HashMap::from([
|
||||
("best_val_loss".to_string(), 0.03),
|
||||
("sharpe_ratio".to_string(), 2.5),
|
||||
]),
|
||||
current_metrics: HashMap::from([
|
||||
("best_val_loss".to_string(), 0.10),
|
||||
("sharpe_ratio".to_string(), 1.0),
|
||||
]),
|
||||
trained_at: now,
|
||||
job_id: "job-99".to_string(),
|
||||
};
|
||||
|
||||
// Map to proto the same way the handler does
|
||||
let proto = super::proto::PendingPromotion {
|
||||
model_id: model.model_id.clone(),
|
||||
model_type: model.model_type.clone(),
|
||||
symbol: model.symbol.clone(),
|
||||
s3_path: model.s3_path.clone(),
|
||||
new_metrics: model.new_metrics.clone().into_iter().collect(),
|
||||
current_metrics: model.current_metrics.clone().into_iter().collect(),
|
||||
trained_at: model.trained_at.timestamp(),
|
||||
job_id: model.job_id.clone(),
|
||||
};
|
||||
|
||||
assert_eq!(proto.model_id, "m-123");
|
||||
assert_eq!(proto.model_type, "DQN");
|
||||
assert_eq!(proto.symbol, "ES.FUT");
|
||||
assert_eq!(proto.s3_path, "s3://bucket/better.bin");
|
||||
assert_eq!(proto.job_id, "job-99");
|
||||
assert_eq!(proto.trained_at, now.timestamp());
|
||||
assert_eq!(proto.new_metrics.len(), 2);
|
||||
assert_eq!(proto.current_metrics.len(), 2);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_approve_promotion_activates_pending_model() {
|
||||
use crate::promotion_manager::{ActiveModel, PromotionManager};
|
||||
use chrono::Utc;
|
||||
|
||||
let pm = Arc::new(PromotionManager::new());
|
||||
|
||||
// Seed an active model so register_completion returns PendingPromotion
|
||||
{
|
||||
let active = pm.active_models_for_test();
|
||||
let mut w = active.write().await;
|
||||
w.insert(
|
||||
("DQN".to_string(), "ES.FUT".to_string()),
|
||||
ActiveModel {
|
||||
model_id: "old-model".to_string(),
|
||||
s3_path: "s3://bucket/old.bin".to_string(),
|
||||
metrics: HashMap::from([
|
||||
("best_val_loss".to_string(), 0.10),
|
||||
("sharpe_ratio".to_string(), 1.0),
|
||||
]),
|
||||
promoted_at: Utc::now(),
|
||||
},
|
||||
);
|
||||
}
|
||||
|
||||
// Register a better model
|
||||
let new_metrics = HashMap::from([
|
||||
("best_val_loss".to_string(), 0.03),
|
||||
("sharpe_ratio".to_string(), 2.5),
|
||||
]);
|
||||
let status = pm
|
||||
.register_completion("job-a", "DQN", "ES.FUT", "s3://bucket/better.bin", new_metrics)
|
||||
.await;
|
||||
assert_eq!(
|
||||
status,
|
||||
crate::promotion_manager::PromotionStatus::PendingPromotion
|
||||
);
|
||||
|
||||
let pending = pm.list_pending().await;
|
||||
assert_eq!(pending.len(), 1);
|
||||
let model_id = pending[0].model_id.clone();
|
||||
|
||||
// Approve it
|
||||
pm.approve(&model_id).await.unwrap();
|
||||
|
||||
// Pending list should be empty
|
||||
assert!(pm.list_pending().await.is_empty());
|
||||
|
||||
// Active model should be updated
|
||||
let active = pm.active_models_for_test();
|
||||
let r = active.read().await;
|
||||
let am = r
|
||||
.get(&("DQN".to_string(), "ES.FUT".to_string()))
|
||||
.unwrap();
|
||||
assert_eq!(am.model_id, model_id);
|
||||
assert_eq!(am.s3_path, "s3://bucket/better.bin");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_reject_promotion_removes_pending_model() {
|
||||
use crate::promotion_manager::{ActiveModel, PromotionManager};
|
||||
use chrono::Utc;
|
||||
|
||||
let pm = Arc::new(PromotionManager::new());
|
||||
|
||||
// Seed an active model
|
||||
{
|
||||
let active = pm.active_models_for_test();
|
||||
let mut w = active.write().await;
|
||||
w.insert(
|
||||
("PPO".to_string(), "NQ.FUT".to_string()),
|
||||
ActiveModel {
|
||||
model_id: "old-ppo".to_string(),
|
||||
s3_path: "s3://bucket/old-ppo.bin".to_string(),
|
||||
metrics: HashMap::from([
|
||||
("best_val_loss".to_string(), 0.10),
|
||||
("sharpe_ratio".to_string(), 1.0),
|
||||
]),
|
||||
promoted_at: Utc::now(),
|
||||
},
|
||||
);
|
||||
}
|
||||
|
||||
// Register a better model
|
||||
let new_metrics = HashMap::from([
|
||||
("best_val_loss".to_string(), 0.04),
|
||||
("sharpe_ratio".to_string(), 2.0),
|
||||
]);
|
||||
let status = pm
|
||||
.register_completion("job-r", "PPO", "NQ.FUT", "s3://bucket/cand.bin", new_metrics)
|
||||
.await;
|
||||
assert_eq!(
|
||||
status,
|
||||
crate::promotion_manager::PromotionStatus::PendingPromotion
|
||||
);
|
||||
|
||||
let pending = pm.list_pending().await;
|
||||
assert_eq!(pending.len(), 1);
|
||||
let model_id = pending[0].model_id.clone();
|
||||
|
||||
// Reject it
|
||||
pm.reject(&model_id, "needs more walk-forward windows")
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
// Pending list should be empty
|
||||
assert!(pm.list_pending().await.is_empty());
|
||||
|
||||
// Active model should NOT have changed
|
||||
let active = pm.active_models_for_test();
|
||||
let r = active.read().await;
|
||||
let am = r
|
||||
.get(&("PPO".to_string(), "NQ.FUT".to_string()))
|
||||
.unwrap();
|
||||
assert_eq!(am.model_id, "old-ppo");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_approve_nonexistent_model_returns_error() {
|
||||
use crate::promotion_manager::PromotionManager;
|
||||
|
||||
let pm = PromotionManager::new();
|
||||
let result = pm.approve("nonexistent-id").await;
|
||||
assert!(result.is_err());
|
||||
assert!(
|
||||
result
|
||||
.unwrap_err()
|
||||
.to_string()
|
||||
.contains("pending model not found")
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_reject_nonexistent_model_returns_error() {
|
||||
use crate::promotion_manager::PromotionManager;
|
||||
|
||||
let pm = PromotionManager::new();
|
||||
let result = pm.reject("nonexistent-id", "bad model").await;
|
||||
assert!(result.is_err());
|
||||
assert!(
|
||||
result
|
||||
.unwrap_err()
|
||||
.to_string()
|
||||
.contains("pending model not found")
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -113,6 +113,7 @@ async fn setup_test_service() -> (Arc<TuningManager>, Arc<MLTrainingServiceImpl>
|
||||
ml_config,
|
||||
promotion_mgr,
|
||||
None, // no K8s dispatcher in tests
|
||||
None, // no job spawner in tests
|
||||
));
|
||||
|
||||
(tuning_manager, service, temp_dir)
|
||||
|
||||
@@ -75,6 +75,7 @@ async fn setup_ml_training_service() -> Result<MLTrainingServiceImpl> {
|
||||
config,
|
||||
promotion_manager,
|
||||
None, // no K8s dispatcher in tests
|
||||
None, // no job spawner in tests
|
||||
))
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user