main.rs (7162B)
1 use askama_axum::Template; 2 use axum::routing::get; 3 use axum::{http::StatusCode, response::IntoResponse, routing::post, Json, Router}; 4 use axum_auth::AuthBearer; 5 use axum_macros::debug_handler; 6 use color_eyre::eyre::{bail, Result}; 7 use linkify::LinkFinder; 8 use once_cell::sync::OnceCell; 9 use serde::{Deserialize, Serialize}; 10 use std::io::Write; 11 use std::net::Ipv6Addr; 12 use std::{fs::OpenOptions, net::SocketAddr}; 13 use tensorflow::{Graph, SavedModelBundle, SessionOptions, SessionRunArgs, Tensor}; 14 use tracing::{error, info}; 15 use voca_rs::strip; 16 17 static GRAPH: OnceCell<Graph> = OnceCell::new(); 18 static MODEL: OnceCell<SavedModelBundle> = OnceCell::new(); 19 20 #[tokio::main] 21 async fn main() -> Result<()> { 22 color_eyre::install()?; 23 // initialize tracing 24 tracing_subscriber::fmt::init(); 25 info!("Starting up"); 26 27 let model_path = match std::env::var("MODEL_PATH") { 28 Ok(val) => val, 29 Err(_) => bail!("Missing MODEL_PATH env var"), 30 }; 31 32 let mut graph = Graph::new(); 33 let bundle = SavedModelBundle::load(&SessionOptions::new(), ["serve"], &mut graph, model_path)?; 34 GRAPH.set(graph).unwrap(); 35 MODEL.set(bundle).unwrap(); 36 37 // build our application with a route 38 let app = Router::new() 39 .route("/", get(index)) 40 .route("/health", get(health)) 41 // `GET /test` goes to `test` 42 .route("/test", post(test)) 43 // `POST /submit` goes to `submit` 44 .route("/submit", post(submit)) 45 .route("/submit_review", post(submit_for_review)); 46 47 let all_v6 = SocketAddr::new(Ipv6Addr::UNSPECIFIED.into(), 3000); 48 49 info!("listening on port 3000"); 50 axum::Server::bind(&all_v6) 51 .serve(app.into_make_service()) 52 .await?; 53 Ok(()) 54 } 55 56 async fn health() -> impl IntoResponse { 57 StatusCode::OK 58 } 59 60 #[derive(Template)] 61 #[template(path = "index.html")] 62 struct IndexTemplate {} 63 64 async fn index() -> IndexTemplate { 65 info!("index"); 66 IndexTemplate {} 67 } 68 69 async fn test(Json(payload): Json<TestData>) -> impl IntoResponse { 70 let bundle = MODEL.get().unwrap(); 71 let graph = GRAPH.get().unwrap(); 72 let session = &bundle.session; 73 let meta = bundle.meta_graph_def(); 74 let signature = meta 75 .get_signature(tensorflow::DEFAULT_SERVING_SIGNATURE_DEF_KEY) 76 .unwrap(); 77 let input_info = signature.get_input("input_1").unwrap(); 78 let output_info = signature.get_output("output_1").unwrap(); 79 80 let input_op = graph 81 .operation_by_name_required(&input_info.name().name) 82 .unwrap(); 83 let output_op = graph 84 .operation_by_name_required(&output_info.name().name) 85 .unwrap(); 86 87 let tensor: Tensor<String> = Tensor::from(&[payload.input_data.clone()]); 88 let mut args = SessionRunArgs::new(); 89 args.add_feed(&input_op, 0, &tensor); 90 91 let out = args.request_fetch(&output_op, 0); 92 93 session 94 .run(&mut args) 95 .expect("Error occurred during calculations"); 96 let out_res: f32 = args.fetch(out).unwrap()[0]; 97 98 let response = Prediction { 99 input_data: payload.input_data, 100 score: out_res, 101 }; 102 103 (StatusCode::OK, Json(response)) 104 } 105 106 #[debug_handler] 107 async fn submit( 108 AuthBearer(token): AuthBearer, 109 Json(payload): Json<SubmitData>, 110 ) -> impl IntoResponse { 111 let access_token = match std::env::var("ACCESS_TOKEN") { 112 Ok(val) => val, 113 Err(_) => { 114 error!("Missing ACCESS_TOKEN env var"); 115 return StatusCode::INTERNAL_SERVER_ERROR; 116 } 117 }; 118 if token != access_token { 119 return StatusCode::UNAUTHORIZED; 120 } 121 122 // TODO implement 123 StatusCode::NOT_IMPLEMENTED 124 } 125 126 #[debug_handler] 127 async fn submit_for_review( 128 AuthBearer(token): AuthBearer, 129 Json(payload): Json<SubmitReview>, 130 ) -> impl IntoResponse { 131 let access_token = match std::env::var("ACCESS_TOKEN") { 132 Ok(val) => val, 133 Err(_) => { 134 error!("Missing ACCESS_TOKEN env var"); 135 return StatusCode::INTERNAL_SERVER_ERROR; 136 } 137 }; 138 if token != access_token { 139 return StatusCode::UNAUTHORIZED; 140 } 141 142 std::fs::create_dir_all("./data/").unwrap(); 143 let file = OpenOptions::new() 144 .write(true) 145 .append(true) 146 .create(true) 147 .open("./data/review.txt"); 148 149 // Sanitize 150 // We remove newlines, html tags and links 151 let sanitized = strip::strip_tags(&payload.input_data); 152 //let sanitized = sanitized.replace(['\r', '\n'], " "); 153 let mut sanitized = trim_whitespace(&sanitized); 154 //let mut finder = LinkFinder::new(); 155 //let cloned_sanitized = sanitized.clone(); 156 //finder.url_must_have_scheme(false); 157 //let links: Vec<_> = finder.links(&cloned_sanitized).collect(); 158 //for link in links { 159 // sanitized = sanitized.replace(link.as_str(), " "); 160 //} 161 match file { 162 Ok(mut file) => { 163 if let Err(e) = writeln!(file, "{}", sanitized) { 164 eprintln!("Couldn't write to file: {}", e); 165 return StatusCode::INTERNAL_SERVER_ERROR; 166 } 167 } 168 Err(e) => { 169 eprintln!("Couldn't open file: {}", e); 170 return StatusCode::INTERNAL_SERVER_ERROR; 171 } 172 } 173 174 StatusCode::OK 175 } 176 177 fn trim_whitespace(s: &str) -> String { 178 let mut new_str = s.trim().to_owned(); 179 let mut prev = ' '; // The initial value doesn't really matter 180 new_str.retain(|ch| { 181 let result = ch != ' ' || prev != ' '; 182 prev = ch; 183 result 184 }); 185 new_str 186 } 187 188 #[derive(Deserialize, Serialize)] 189 #[cfg_attr(test, derive(schemars::JsonSchema))] 190 struct TestData { 191 input_data: String, 192 } 193 194 #[derive(Deserialize, Serialize)] 195 #[cfg_attr(test, derive(schemars::JsonSchema))] 196 struct Prediction { 197 input_data: String, 198 score: f32, 199 } 200 201 #[derive(Deserialize, Serialize)] 202 #[cfg_attr(test, derive(schemars::JsonSchema))] 203 struct SubmitData { 204 input_data: String, 205 spam: bool, 206 } 207 208 #[derive(Deserialize, Serialize)] 209 #[cfg_attr(test, derive(schemars::JsonSchema))] 210 struct SubmitReview { 211 input_data: String, 212 } 213 214 #[cfg(test)] 215 mod test { 216 use crate::{Prediction, SubmitData, SubmitReview, TestData}; 217 218 #[test] 219 fn generate_schema() { 220 let test_data_schema = schemars::schema_for!(TestData); 221 let prediction_schema = schemars::schema_for!(Prediction); 222 let submit_data_schema = schemars::schema_for!(SubmitData); 223 let submit_review_schema = schemars::schema_for!(SubmitReview); 224 225 std::fs::create_dir_all("./schemas").unwrap(); 226 std::fs::write( 227 "./schemas/test_data.json", 228 serde_json::to_string_pretty(&test_data_schema).unwrap(), 229 ) 230 .unwrap(); 231 std::fs::write( 232 "./schemas/prediction.json", 233 serde_json::to_string_pretty(&prediction_schema).unwrap(), 234 ) 235 .unwrap(); 236 std::fs::write( 237 "./schemas/submit_data.json", 238 serde_json::to_string_pretty(&submit_data_schema).unwrap(), 239 ) 240 .unwrap(); 241 std::fs::write( 242 "./schemas/submit_review.json", 243 serde_json::to_string_pretty(&submit_review_schema).unwrap(), 244 ) 245 .unwrap(); 246 } 247 }