matrix-spam-ml

git clone git://archive.git.mtrnord.blog/MTRNord/matrix-spam-ml.git
Log | Files | Refs | Submodules | README | LICENSE

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 }