Skip to main content

cortiq_server/
api.rs

1//! Cortiq extension API endpoints.
2
3use 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
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    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// ─── Masks ───────────────────────────────────────────────
40
41#[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    /// Held-out quality value; null = not measured (never a declaration).
52    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// ─── Task Switch ─────────────────────────────────────────
81
82#[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}