1use crate::AppState;
4use axum::{
5 Router,
6 extract::State,
7 response::Json,
8 routing::{get, post},
9};
10use serde::{Deserialize, Serialize};
11use std::sync::Arc;
12
13pub fn routes() -> Router<Arc<AppState>> {
15 Router::new()
16 .route("/v1/cortiq/status", get(get_status))
17 .route("/v1/cortiq/masks", get(list_masks))
18 .route("/v1/cortiq/switch", post(switch_task))
19}
20
21async fn get_status(State(state): State<Arc<AppState>>) -> Json<serde_json::Value> {
24 let status = state.runtime.status().await;
25 Json(serde_json::to_value(status).unwrap_or_default())
26}
27
28#[derive(Serialize)]
31struct MaskListResponse {
32 masks: Vec<MaskInfo>,
33}
34
35#[derive(Serialize)]
36struct MaskInfo {
37 task_id: u32,
38 name: String,
39 sparsity: f32,
40 quality_score: Option<f32>,
42 quality_metric: Option<String>,
43 active_layers: usize,
44 active_neurons_avg: f64,
45 has_hot_pack: bool,
46}
47
48async fn list_masks(State(state): State<Arc<AppState>>) -> Json<MaskListResponse> {
49 let masks: Vec<MaskInfo> = state
50 .runtime
51 .masks()
52 .masks
53 .iter()
54 .map(|m| MaskInfo {
55 task_id: m.task_id,
56 name: m.name.clone(),
57 sparsity: m.sparsity,
58 quality_score: m.quality.as_ref().map(|q| q.value),
59 quality_metric: m.quality.as_ref().map(|q| q.metric.clone()),
60 active_layers: m.active_layer_count(),
61 active_neurons_avg: m.avg_active_neurons(),
62 has_hot_pack: m.has_hot_pack,
63 })
64 .collect();
65
66 Json(MaskListResponse { masks })
67}
68
69#[derive(Deserialize)]
72struct SwitchRequest {
73 task: String,
74}
75
76async fn switch_task(
77 State(state): State<Arc<AppState>>,
78 Json(req): Json<SwitchRequest>,
79) -> Result<Json<serde_json::Value>, axum::http::StatusCode> {
80 match state.runtime.switch_task(&req.task).await {
81 Ok(result) => Ok(Json(serde_json::to_value(result).unwrap_or_default())),
82 Err(e) => {
83 tracing::error!("Task switch failed: {}", e);
84 Err(axum::http::StatusCode::NOT_FOUND)
85 }
86 }
87}