Skip to main content

cortiq_server/
api.rs

1//! Cortiq extension API endpoints.
2
3use crate::AppState;
4use axum::{
5    extract::State,
6    response::Json,
7    routing::{get, post},
8    Router,
9};
10use serde::{Deserialize, Serialize};
11use std::sync::Arc;
12
13/// Register Cortiq extension routes.
14pub 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
21// ─── Status ──────────────────────────────────────────────
22
23async 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// ─── Masks ───────────────────────────────────────────────
29
30#[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    /// Held-out quality value; null = not measured (never a declaration).
41    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// ─── Task Switch ─────────────────────────────────────────
70
71#[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}