Skip to main content

snarkify_sdk/
prover.rs

1use async_trait::async_trait;
2use chrono::{DateTime, Utc};
3use cloudevents::{AttributesReader, Data, Event, EventBuilder, EventBuilderV10};
4use serde::{de::DeserializeOwned, Deserialize, Serialize};
5use tracing::*;
6
7use crate::datetime_serde::default_dt_formatter;
8
9pub const READ_INPUT_FROM_URL_HEADER: &str = "ce-inputurl";
10
11enum ProofStatus {
12    Success = 2,
13    Failure = 3,
14}
15
16#[derive(Serialize, Deserialize)]
17struct ProofResponse {
18    task_id: String,
19    result: String,
20    status: u8,
21    #[serde(with = "default_dt_formatter")]
22    started: DateTime<Utc>,
23    #[serde(with = "default_dt_formatter")]
24    finished: DateTime<Utc>,
25}
26
27#[async_trait]
28pub trait ProofHandler {
29    type Input: DeserializeOwned + Send + 'static;
30    type Output: Serialize + Send + 'static;
31    type Error: Serialize + Send + 'static;
32
33    async fn prove(data: Self::Input) -> Result<Self::Output, Self::Error>;
34
35    #[instrument(skip(event), fields(task_id = event.id()))] // Handle span fields
36    async fn handle(event: Event) -> Result<Event, actix_web::Error> {
37        let started = Utc::now();
38        trace!("Start time {started}");
39
40        let input_url = event.extension(READ_INPUT_FROM_URL_HEADER);
41        let result_with_input = match input_url {
42            Some(url) => {
43                // Get input from URL if provided
44                let url_str = url.to_string();
45                match reqwest::get(&url_str).await {
46                    Ok(response) => match response.text().await {
47                        Ok(text) => serde_json::from_str::<Self::Input>(&text).map_err(|err| {
48                            format!("Failed to parse JSON from prover input URL: {url_str}; err: {err:?}")
49                        }),
50                        Err(err) => Err(format!(
51                            "Failed to read HTTP response body from prover input URL: {url_str}. Error: {err:?}",
52                        )),
53                    },
54                    Err(err) => Err(format!(
55                        "Failed to fetch prover input from URL: {url_str}. Network error: {err:?}",
56                    )),
57                }
58            }
59            None => {
60                // Use event data if no URL provided
61                event
62                    .data()
63                    .ok_or("Event payload is missing".to_string())
64                    .and_then(|data| {
65                        match data {
66                            Data::Binary(v) => serde_json::from_slice(v),
67                            Data::String(v) => serde_json::from_str(v),
68                            Data::Json(v) => serde_json::from_value(v.clone()),
69                        }
70                        .map_err(|err| format!("Failed to parse Json from event payload. Error: {err:?}"))
71                    })
72            }
73        };
74
75        let (result, status) = match result_with_input {
76            Ok(input) => {
77                // Spawn the async task
78                let info_span = info_span!("prove");
79                let result = tokio::spawn(async move { Self::prove(input).instrument(info_span).await }).await;
80                match result {
81                    Ok(prove_result) => match prove_result {
82                        Ok(proof) => (
83                            serde_json::to_string(&proof).map_err(|err| {
84                                error!("Error while serializing success output: {err:?}");
85                                err
86                            })?,
87                            ProofStatus::Success,
88                        ),
89                        Err(error) => (
90                            serde_json::to_string(&error).map_err(|err| {
91                                error!("Error while serializing prove error: {err:?}");
92                                err
93                            })?,
94                            ProofStatus::Failure,
95                        ),
96                    },
97                    Err(err) => match err.try_into_panic() {
98                        Ok(panic) => {
99                            let panic_message = panic
100                                .downcast_ref::<String>()
101                                .map(|s| s.as_str())
102                                .or_else(|| panic.downcast_ref::<&str>().copied())
103                                .unwrap_or("A panic occurred during proof.");
104                            (panic_message.to_string(), ProofStatus::Failure)
105                        }
106                        Err(join_err) => (
107                            format!("Unexpected error while waiting for proof: {join_err:?}"),
108                            ProofStatus::Failure,
109                        ),
110                    },
111                }
112            }
113            Err(handle_request_err) => {
114                error!(handle_request_err);
115                // Initial JSON deserialization error
116                (handle_request_err.to_string(), ProofStatus::Failure)
117            }
118        };
119
120        let response = serde_json::to_value(ProofResponse {
121            task_id: event.id().to_string(),
122            result,
123            status: status as u8,
124            started,
125            finished: Utc::now(),
126        })?;
127
128        // events are routed to the right prover through the event.ty() attribute,
129        // which is the prover id
130        let source = format!("prover-{}", event.ty());
131        EventBuilderV10::new()
132            .id(event.id())
133            .source(source)
134            .ty("tenant-service-result") // Be careful of changing this, this is used as filter for backend service
135            .data("application/json", response)
136            .build()
137            .map_err(actix_web::error::ErrorInternalServerError)
138    }
139}
140
141#[cfg(test)]
142mod tests {
143    use super::*;
144    use mockito::Server;
145    use serde_json::json;
146
147    struct MyHandler {}
148    struct MyHandlerWithPanic {}
149
150    #[derive(Deserialize)]
151    struct MyInput {
152        name: String,
153    }
154
155    #[derive(Serialize)]
156    struct MyOutput {
157        result: String,
158    }
159
160    #[async_trait]
161    impl ProofHandler for MyHandler {
162        type Input = MyInput;
163        type Output = MyOutput;
164        type Error = String;
165        async fn prove(data: MyInput) -> Result<MyOutput, String> {
166            Ok(MyOutput { result: data.name })
167        }
168    }
169
170    #[async_trait]
171    impl ProofHandler for MyHandlerWithPanic {
172        type Input = MyInput;
173        type Output = MyOutput;
174        type Error = ();
175        async fn prove(_data: MyInput) -> Result<MyOutput, ()> {
176            panic!("Houston, we have a problem")
177        }
178    }
179
180    fn get_payload(event: Event) -> ProofResponse {
181        if let Some(Data::Json(v)) = event.data() {
182            serde_json::from_value::<ProofResponse>(v.clone()).unwrap()
183        } else {
184            panic!("Expected JSON data in the event.");
185        }
186    }
187
188    #[actix_rt::test]
189    async fn test_prover_handles_valid_event() {
190        let mock_event = EventBuilderV10::new()
191            .id("test_id")
192            .source("test://source")
193            .ty("12345678")
194            .data("application/json", json!({"name": "aloha"}))
195            .build()
196            .unwrap();
197        let result = MyHandler::handle(mock_event).await;
198        assert!(result.is_ok());
199
200        let event = result.unwrap();
201        // Check some properties of the returned event, e.g.:
202        assert_eq!(event.source().to_string(), "prover-12345678");
203        assert_eq!(event.ty().to_string(), "tenant-service-result");
204
205        let response = get_payload(event);
206        assert_eq!(response.status, ProofStatus::Success as u8);
207        assert_eq!(response.result, "{\"result\":\"aloha\"}");
208    }
209
210    #[actix_rt::test]
211    async fn test_prover_handles_event_wrong_input() {
212        let mock_event = EventBuilderV10::new()
213            .id("test_id")
214            .source("test://source")
215            .data("application/json", json!({"wrong_key": "aloha"}))
216            .ty("test.type")
217            .build()
218            .unwrap();
219        let result = MyHandler::handle(mock_event).await;
220        assert!(result.is_ok());
221
222        let response = get_payload(result.unwrap());
223        assert_eq!(response.status, ProofStatus::Failure as u8);
224        assert!(response.result.contains("missing field `name`"));
225    }
226
227    #[actix_rt::test]
228    async fn test_prover_handles_event_missing_input() {
229        let mock_event = EventBuilderV10::new()
230            .id("test_id")
231            .source("test://source")
232            .ty("test.type")
233            .build()
234            .unwrap();
235        let result = MyHandler::handle(mock_event).await;
236        assert!(result.is_ok());
237
238        let response = get_payload(result.unwrap());
239        assert_eq!(response.status, ProofStatus::Failure as u8);
240        assert!(response.result.contains("Event payload is missing"));
241    }
242
243    #[actix_rt::test]
244    async fn test_prover_handles_invalid_json_input() {
245        let mock_event = EventBuilderV10::new()
246            .id("test_id")
247            .source("test://source")
248            .data("application/json", "invalid json {") // Invalid JSON string
249            .ty("test.type")
250            .build()
251            .unwrap();
252        let result = MyHandler::handle(mock_event).await;
253        assert!(result.is_ok());
254
255        let response = get_payload(result.unwrap());
256        assert_eq!(response.status, ProofStatus::Failure as u8);
257        assert!(
258            response.result.contains("Failed to parse Json from event payload"),
259            "Expected JSON parsing error message, got: {}",
260            response.result
261        );
262    }
263
264    #[actix_rt::test]
265    async fn test_prover_handles_panic() {
266        let mock_event = EventBuilderV10::new()
267            .id("test_id")
268            .source("test://source")
269            .ty("12345678")
270            .data("application/json", json!({"name": "aloha"}))
271            .build()
272            .unwrap();
273        let result = MyHandlerWithPanic::handle(mock_event).await;
274        assert!(result.is_ok());
275
276        let response = get_payload(result.unwrap());
277        assert_eq!(response.status, ProofStatus::Failure as u8);
278        assert_eq!(response.result, "Houston, we have a problem");
279    }
280
281    #[actix_rt::test]
282    async fn test_prover_handles_url_input() {
283        // Create the mock server in a separate thread to avoid runtime conflicts
284        let server = std::thread::spawn(|| {
285            let mut server = Server::new();
286            let mock = server
287                .mock("GET", "/input")
288                .with_status(200)
289                .with_header("content-type", "application/json")
290                .with_body(r#"{"name": "from_url"}"#)
291                .expect(1)
292                .create();
293
294            (server, mock)
295        })
296        .join()
297        .unwrap();
298
299        let test_url = format!("{}/input", server.0.url());
300        let mock_event = EventBuilderV10::new()
301            .id("test_id")
302            .source("test://source")
303            .ty("12345678")
304            .extension(READ_INPUT_FROM_URL_HEADER, test_url)
305            .build()
306            .unwrap();
307
308        let result = MyHandler::handle(mock_event).await;
309
310        assert!(result.is_ok());
311        let event = result.unwrap();
312        assert_eq!(event.source().to_string(), "prover-12345678");
313        assert_eq!(event.ty().to_string(), "tenant-service-result");
314
315        let response = get_payload(event);
316        assert_eq!(response.status, ProofStatus::Success as u8);
317        assert_eq!(response.result, "{\"result\":\"from_url\"}");
318
319        // Verify that the mock was called
320        server.1.assert();
321    }
322
323    #[actix_rt::test]
324    async fn test_prover_handles_url_fetch_failure() {
325        let test_url = "http://localhost:12345/nonexistent";
326        let mock_event = EventBuilderV10::new()
327            .id("test_id")
328            .source("test://source")
329            .ty("12345678")
330            .extension(READ_INPUT_FROM_URL_HEADER, test_url)
331            .build()
332            .unwrap();
333
334        let result = MyHandler::handle(mock_event).await;
335        assert!(result.is_ok(), "Handler should return Ok even for URL fetch failures");
336
337        let event = result.unwrap();
338        let response = get_payload(event);
339        assert_eq!(response.status, ProofStatus::Failure as u8);
340        assert!(
341            response.result.contains("Failed to fetch prover input from URL")
342                && response.result.contains("Network error")
343                && response.result.contains(test_url),
344            "Expected detailed error message about URL fetch failure, got: {}",
345            response.result
346        );
347    }
348
349    #[actix_rt::test]
350    async fn test_prover_handles_url_invalid_json() {
351        let server = std::thread::spawn(|| {
352            let mut server = Server::new();
353            let mock = server
354                .mock("GET", "/invalid-json")
355                .with_status(200)
356                .with_header("content-type", "application/json")
357                .with_body(r#"{"invalid json format"#)
358                .expect(1)
359                .create();
360
361            (server, mock)
362        })
363        .join()
364        .unwrap();
365
366        let test_url = format!("{}/invalid-json", server.0.url());
367        let mock_event = EventBuilderV10::new()
368            .id("test_id")
369            .source("test://source")
370            .ty("12345678")
371            .extension(READ_INPUT_FROM_URL_HEADER, test_url.clone())
372            .build()
373            .unwrap();
374
375        let result = MyHandler::handle(mock_event).await;
376        assert!(result.is_ok());
377
378        let response = get_payload(result.unwrap());
379        assert_eq!(response.status, ProofStatus::Failure as u8);
380        assert!(
381            response.result.contains("Failed to parse JSON from prover input URL")
382                && response.result.contains(&test_url),
383            "Expected detailed error about JSON parsing failure, got: {}",
384            response.result
385        );
386
387        server.1.assert();
388    }
389
390    #[actix_rt::test]
391    async fn test_prover_handles_url_wrong_json_schema() {
392        let server = std::thread::spawn(|| {
393            let mut server = Server::new();
394            let mock = server
395                .mock("GET", "/wrong-schema")
396                .with_status(200)
397                .with_header("content-type", "application/json")
398                .with_body(r#"{"wrong_field": "value"}"#) // Valid JSON but wrong schema
399                .expect(1)
400                .create();
401
402            (server, mock)
403        })
404        .join()
405        .unwrap();
406
407        let test_url = format!("{}/wrong-schema", server.0.url());
408        let mock_event = EventBuilderV10::new()
409            .id("test_id")
410            .source("test://source")
411            .ty("12345678")
412            .extension(READ_INPUT_FROM_URL_HEADER, test_url)
413            .build()
414            .unwrap();
415
416        let result = MyHandler::handle(mock_event).await;
417        assert!(result.is_ok());
418
419        let response = get_payload(result.unwrap());
420        assert_eq!(response.status, ProofStatus::Failure as u8);
421        assert!(response.result.contains("missing field `name`"));
422
423        server.1.assert();
424    }
425}