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 let mut value = serde_json::to_value(status).unwrap_or_default();
26 if let Some(source) = state.runtime.model().arch().deepseek_v41.as_ref() {
27 let vision = cortiq_engine::dsv41_vision::VisionConfig::from_source(source).ok();
28 value["capabilities"] = serde_json::json!({
29 "tools": true,
30 "vision": vision.as_ref().is_some_and(|config| config.vision_enabled()),
31 "image_token_id": vision.map(|config| config.image_token_id),
32 "reasoning_effort": true,
33 "dsml": true
34 });
35 }
36 Json(value)
37}
38
39#[derive(Serialize)]
42struct MaskListResponse {
43 masks: Vec<MaskInfo>,
44}
45
46#[derive(Serialize)]
47struct MaskInfo {
48 task_id: u32,
49 name: String,
50 sparsity: f32,
51 quality_score: Option<f32>,
53 quality_metric: Option<String>,
54 active_layers: usize,
55 active_neurons_avg: f64,
56 has_hot_pack: bool,
57}
58
59async fn list_masks(State(state): State<Arc<AppState>>) -> Json<MaskListResponse> {
60 let masks: Vec<MaskInfo> = state
61 .runtime
62 .masks()
63 .masks
64 .iter()
65 .map(|m| MaskInfo {
66 task_id: m.task_id,
67 name: m.name.clone(),
68 sparsity: m.sparsity,
69 quality_score: m.quality.as_ref().map(|q| q.value),
70 quality_metric: m.quality.as_ref().map(|q| q.metric.clone()),
71 active_layers: m.active_layer_count(),
72 active_neurons_avg: m.avg_active_neurons(),
73 has_hot_pack: m.has_hot_pack,
74 })
75 .collect();
76
77 Json(MaskListResponse { masks })
78}
79
80#[derive(Deserialize)]
83struct SwitchRequest {
84 task: String,
85}
86
87async fn switch_task(
88 State(state): State<Arc<AppState>>,
89 Json(req): Json<SwitchRequest>,
90) -> Result<Json<serde_json::Value>, axum::http::StatusCode> {
91 match state.runtime.switch_task(&req.task).await {
92 Ok(result) => Ok(Json(serde_json::to_value(result).unwrap_or_default())),
93 Err(e) => {
94 tracing::error!("Task switch failed: {}", e);
95 Err(axum::http::StatusCode::NOT_FOUND)
96 }
97 }
98}