@@ -11,7 +11,7 @@ use axum::{
1111use futures:: StreamExt ;
1212use serde_json:: json;
1313use tower_http:: cors:: { Any , CorsLayer } ;
14- use tracing:: { info , error } ;
14+ use tracing:: { error , info } ;
1515
1616use crate :: engine:: InferenceEngine ;
1717use crate :: error:: MohawkError ;
@@ -26,32 +26,27 @@ pub struct AppState {
2626/// Create the API router with all endpoints
2727pub fn create_router ( engine : InferenceEngine ) -> Router {
2828 let state = AppState { engine } ;
29-
29+
3030 // Configure CORS for GUI access
3131 let cors = CorsLayer :: new ( )
3232 . allow_origin ( Any )
3333 . allow_methods ( Any )
3434 . allow_headers ( Any ) ;
35-
35+
3636 Router :: new ( )
3737 // Health & metrics
3838 . route ( "/health" , get ( health_check) )
3939 . route ( "/metrics" , get ( get_metrics) )
40-
4140 // Model management (OpenAI compatible)
4241 . route ( "/v1/models" , get ( list_models) )
4342 . route ( "/v1/models/:model_id" , get ( get_model) )
44-
4543 // Model loading/unloading (Mohawk extension)
4644 . route ( "/api/models/load" , post ( load_model) )
4745 . route ( "/api/models/unload" , post ( unload_model) )
48-
4946 // Chat completions (OpenAI compatible)
5047 . route ( "/v1/chat/completions" , post ( chat_completions) )
51-
5248 // Legacy completions
5349 . route ( "/v1/completions" , post ( completions) )
54-
5550 . layer ( cors)
5651 . with_state ( state)
5752}
@@ -65,7 +60,7 @@ async fn health_check(State(state): State<AppState>) -> Json<HealthResponse> {
6560/// Metrics endpoint (Prometheus format)
6661async fn get_metrics ( State ( state) : State < AppState > ) -> impl IntoResponse {
6762 let stats = state. engine . get_stats ( ) . await ;
68-
63+
6964 let metrics = format ! (
7065 r#"# HELP mohawk_uptime_seconds Server uptime in seconds
7166# TYPE mohawk_uptime_seconds counter
@@ -79,18 +74,16 @@ mohawk_requests_total {}
7974# TYPE mohawk_models_loaded gauge
8075mohawk_models_loaded {}
8176"# ,
82- stats. uptime_secs,
83- stats. requests_total,
84- stats. models_loaded
77+ stats. uptime_secs, stats. requests_total, stats. models_loaded
8578 ) ;
86-
79+
8780 ( StatusCode :: OK , metrics)
8881}
8982
9083/// List available models (OpenAI compatible)
9184async fn list_models ( State ( state) : State < AppState > ) -> Json < ModelListResponse > {
9285 let models = state. engine . list_models ( ) . await ;
93-
86+
9487 Json ( ModelListResponse {
9588 object : "list" . to_string ( ) ,
9689 data : models,
@@ -103,11 +96,12 @@ async fn get_model(
10396 axum:: extract:: Path ( model_id) : axum:: extract:: Path < String > ,
10497) -> Result < Json < ModelInfo > , MohawkError > {
10598 let models = state. engine . list_models ( ) . await ;
106-
107- let model = models. into_iter ( )
99+
100+ let model = models
101+ . into_iter ( )
108102 . find ( |m| m. id == model_id)
109- . ok_or_else ( || MohawkError :: ModelNotFound ( model_id) ) ?;
110-
103+ . ok_or ( MohawkError :: ModelNotFound ( model_id) ) ?;
104+
111105 Ok ( Json ( model) )
112106}
113107
@@ -120,15 +114,18 @@ async fn load_model(
120114 . as_str ( )
121115 . ok_or_else ( || MohawkError :: InvalidRequest ( "model_id required" . to_string ( ) ) ) ?
122116 . to_string ( ) ;
123-
117+
124118 state. engine . load_model ( & model_id) . await ?;
125-
119+
126120 info ! ( "Model loaded: {}" , model_id) ;
127- Ok ( ( StatusCode :: OK , Json ( json ! ( {
128- "success" : true ,
129- "model_id" : model_id,
130- "status" : "loaded"
131- } ) ) ) )
121+ Ok ( (
122+ StatusCode :: OK ,
123+ Json ( json ! ( {
124+ "success" : true ,
125+ "model_id" : model_id,
126+ "status" : "loaded"
127+ } ) ) ,
128+ ) )
132129}
133130
134131/// Unload a model from memory
@@ -140,15 +137,18 @@ async fn unload_model(
140137 . as_str ( )
141138 . ok_or_else ( || MohawkError :: InvalidRequest ( "model_id required" . to_string ( ) ) ) ?
142139 . to_string ( ) ;
143-
140+
144141 state. engine . unload_model ( & model_id) . await ?;
145-
142+
146143 info ! ( "Model unloaded: {}" , model_id) ;
147- Ok ( ( StatusCode :: OK , Json ( json ! ( {
148- "success" : true ,
149- "model_id" : model_id,
150- "status" : "unloaded"
151- } ) ) ) )
144+ Ok ( (
145+ StatusCode :: OK ,
146+ Json ( json ! ( {
147+ "success" : true ,
148+ "model_id" : model_id,
149+ "status" : "unloaded"
150+ } ) ) ,
151+ ) )
152152}
153153
154154/// Chat completions endpoint (OpenAI compatible)
@@ -159,27 +159,24 @@ async fn chat_completions(
159159 if request. stream {
160160 // Streaming response
161161 let stream = state. engine . generate_stream ( request) . await ?;
162-
163- let stream_body = axum:: body:: Body :: from_stream (
164- stream. map ( |result| {
165- match result {
166- Ok ( token) => {
167- let json = serde_json:: to_string ( & token) . unwrap ( ) ;
168- Ok :: < _ , MohawkError > ( format ! ( "data: {}\n \n " , json) )
169- }
170- Err ( e) => {
171- error ! ( "Stream error: {}" , e) ;
172- Ok ( "data: [DONE]\n \n " . to_string ( ) )
173- }
174- }
175- } )
176- ) ;
177-
162+
163+ let stream_body = axum:: body:: Body :: from_stream ( stream. map ( |result| match result {
164+ Ok ( token) => {
165+ let json = serde_json:: to_string ( & token) . unwrap ( ) ;
166+ Ok :: < _ , MohawkError > ( format ! ( "data: {}\n \n " , json) )
167+ }
168+ Err ( e) => {
169+ error ! ( "Stream error: {}" , e) ;
170+ Ok ( "data: [DONE]\n \n " . to_string ( ) )
171+ }
172+ } ) ) ;
173+
178174 Ok ( (
179175 StatusCode :: OK ,
180176 [ ( "Content-Type" , "text/event-stream" ) ] ,
181177 stream_body,
182- ) . into_response ( ) )
178+ )
179+ . into_response ( ) )
183180 } else {
184181 // Non-streaming response
185182 let response = state. engine . generate ( request) . await ?;
@@ -193,23 +190,26 @@ async fn completions(
193190 Json ( payload) : Json < serde_json:: Value > ,
194191) -> Result < impl IntoResponse , MohawkError > {
195192 // Convert legacy format to chat format
196- let prompt = payload[ "prompt" ]
197- . as_str ( )
198- . unwrap_or ( "" )
199- . to_string ( ) ;
200-
193+ let prompt = payload[ "prompt" ] . as_str ( ) . unwrap_or ( "" ) . to_string ( ) ;
194+
201195 let messages = vec ! [ Message {
202196 role: "user" . to_string( ) ,
203197 content: prompt,
204198 } ] ;
205-
199+
206200 let request = InferenceRequest {
207201 messages,
208202 model : payload[ "model" ] . as_str ( ) . map ( |s| s. to_string ( ) ) ,
209203 temperature : payload[ "temperature" ] . as_f64 ( ) . map ( |f| f as f32 ) ,
210204 top_p : payload[ "top_p" ] . as_f64 ( ) . map ( |f| f as f32 ) ,
211- top_k : payload. get ( "top_k" ) . and_then ( |v| v. as_i64 ( ) ) . map ( |i| i as i32 ) ,
212- max_tokens : payload. get ( "max_tokens" ) . and_then ( |v| v. as_i64 ( ) ) . map ( |i| i as i32 ) ,
205+ top_k : payload
206+ . get ( "top_k" )
207+ . and_then ( |v| v. as_i64 ( ) )
208+ . map ( |i| i as i32 ) ,
209+ max_tokens : payload
210+ . get ( "max_tokens" )
211+ . and_then ( |v| v. as_i64 ( ) )
212+ . map ( |i| i as i32 ) ,
213213 stream : payload[ "stream" ] . as_bool ( ) . unwrap_or ( false ) ,
214214 stop : payload. get ( "stop" ) . and_then ( |v| v. as_array ( ) ) . map ( |arr| {
215215 arr. iter ( )
@@ -218,7 +218,7 @@ async fn completions(
218218 } ) ,
219219 system_prompt : None ,
220220 } ;
221-
221+
222222 let response = state. engine . generate ( request) . await ?;
223223 Ok ( Json ( response) )
224224}
0 commit comments