summaryrefslogtreecommitdiff
path: root/rust/tests/websocket.rs
blob: 97a4ae2bac78c821048b49f938036574cea4fa1e (plain)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
use binaryninja::headless::Session;
use binaryninja::rc::Ref;
use binaryninja::string::BnStrCompatible;
use binaryninja::websocket::{
    register_websocket_provider, CoreWebsocketClient, CoreWebsocketProvider, WebsocketClient,
    WebsocketClientCallback, WebsocketProvider,
};

struct MyWebsocketProvider {
    core: CoreWebsocketProvider,
}

impl WebsocketProvider for MyWebsocketProvider {
    type Client = MyWebsocketClient;

    fn handle(&self) -> CoreWebsocketProvider {
        self.core
    }

    fn from_core(core: CoreWebsocketProvider) -> Self {
        MyWebsocketProvider { core }
    }
}

struct MyWebsocketClient {
    core: Ref<CoreWebsocketClient>,
}

impl WebsocketClient for MyWebsocketClient {
    fn from_core(core: Ref<CoreWebsocketClient>) -> Self {
        Self { core }
    }

    fn connect<I, K, V>(&self, host: &str, _headers: I) -> bool
    where
        I: IntoIterator<Item = (K, V)>,
        K: BnStrCompatible,
        V: BnStrCompatible,
    {
        assert_eq!(host, "url");
        true
    }

    fn write(&self, data: &[u8]) -> bool {
        if !self.core.notify_read("sent: ".as_bytes()) {
            return false;
        }
        if !self.core.notify_read(data) {
            return false;
        }
        self.core.notify_read("\n".as_bytes())
    }

    fn disconnect(&self) -> bool {
        true
    }
}

#[derive(Default)]
struct MyClientCallbacks {
    data_read: Vec<u8>,
    did_disconnect: bool,
    did_error: bool,
}

impl WebsocketClientCallback for MyClientCallbacks {
    fn connected(&mut self) -> bool {
        true
    }

    fn disconnected(&mut self) {
        self.did_disconnect = true;
    }

    fn error(&mut self, msg: &str) {
        assert_eq!(msg, "error");
        self.did_error = true;
    }

    fn read(&mut self, data: &[u8]) -> bool {
        self.data_read.extend_from_slice(data);
        true
    }
}

#[test]
fn reg_websocket_provider() {
    let _session = Session::new().expect("Failed to initialize session");
    let provider = register_websocket_provider::<MyWebsocketProvider>("RustWebsocketProvider");
    let client = provider.create_client().unwrap();
    let mut callback = MyClientCallbacks::default();
    let success = client.initialize_connection("url", [("header", "value")], &mut callback);
    assert!(success, "Failed to initialize connection!");
}

#[test]
fn listen_websocket_provider() {
    let _session = Session::new().expect("Failed to initialize session");
    let provider = register_websocket_provider::<MyWebsocketProvider>("RustWebsocketProvider2");

    let client = provider.create_client().unwrap();
    let mut callback = MyClientCallbacks::default();
    client.initialize_connection("url", [("header", "value")], &mut callback);

    assert!(client.write("test1".as_bytes()));
    assert!(client.write("test2".as_bytes()));

    client.notify_error("error");
    client.disconnect();
    drop(client);

    assert_eq!(
        &callback.data_read[..],
        "sent: test1\nsent: test2\n".as_bytes()
    );
    // If we disconnected that means the error callback was not notified.
    assert!(!callback.did_disconnect);
    assert!(callback.did_error);
}