package noise

  1. Overview
  2. Docs

Source file handshake_state.ml

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
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
type role = Initiator | Responder

type t = {
  suite : Suite.t;
  role : role;
  turn : role; (* whose turn it is to write *)
  sym : Symmetric_state.t;
  s : Dh.keypair option;
  e : Dh.keypair option;
  rs : string option;
  re : string option;
  psks : string list;
  patterns : Pattern.message_pattern list;
  finished : bool;
  h_final : string option;
  psk_mode : bool;
  one_way : bool;
}

let ( let* ) = Result.bind
let peer = function Initiator -> Responder | Responder -> Initiator
let hash st = st.suite.Suite.hash
let cipher st = st.suite.Suite.cipher
let dh st = st.suite.Suite.dh

let create ~protocol_name ~role ?prologue ?s ?e ?rs ?re ?(psks = []) () =
  let* suite = Suite.of_protocol_name protocol_name in
  let* pattern =
    Pattern.of_name (List.nth (String.split_on_char '_' protocol_name) 1)
  in
  let h = suite.Suite.hash in
  let sym0 = Symmetric_state.initialize_symmetric ~hash:h protocol_name in
  let sym1 =
    Symmetric_state.mix_hash sym0 ~hash:h (Option.value prologue ~default:"")
  in
  (* Pre-message keys are mixed initiator-first (spec §5.3). *)
  let mix_pre_keys sym role_pre is_local =
    List.fold_left
      (fun sym tok ->
        match tok with
        | Pattern.S -> (
            let key =
              if is_local then
                match s with Some kp -> Some kp.Dh.pub | None -> None
              else rs
            in
            match key with
            | Some pub -> Symmetric_state.mix_hash sym ~hash:h pub
            | None -> sym)
        | Pattern.E -> (
            let key =
              if is_local then
                match e with Some kp -> Some kp.Dh.pub | None -> None
              else re
            in
            match key with
            | Some pub -> Symmetric_state.mix_hash sym ~hash:h pub
            | None -> sym)
        | _ -> sym)
      sym role_pre
  in
  let sym2 =
    let sym_a =
      mix_pre_keys sym1 pattern.Pattern.pre.initiator (role = Initiator)
    in
    mix_pre_keys sym_a pattern.Pattern.pre.responder (role = Responder)
  in
  let psk_mode = List.exists (List.mem Pattern.PSK) pattern.Pattern.messages in
  let one_way = List.length pattern.Pattern.messages = 1 in
  let psk_tokens =
    List.fold_left
      (fun acc msg ->
        acc + List.length (List.filter (fun tok -> tok = Pattern.PSK) msg))
      0 pattern.Pattern.messages
  in
  let* () =
    if List.length psks <> psk_tokens then
      Error
        (Error.Bad_psk
           (Printf.sprintf "pattern requires %d, got %d" psk_tokens
              (List.length psks)))
    else if List.exists (fun psk -> String.length psk <> 32) psks then
      Error (Error.Bad_psk "pre-shared keys must be 32 bytes")
    else Ok ()
  in
  Ok
    {
      suite;
      role;
      turn = Initiator;
      sym = sym2;
      s;
      e;
      rs;
      re;
      psks;
      patterns = pattern.Pattern.messages;
      finished = false;
      h_final = None;
      psk_mode;
      one_way;
    }

let do_dh st local_kp remote_pub =
  (dh st).Dh.dh ~local:local_kp ~remote:remote_pub

let get_e st =
  match st.e with
  | Some kp -> Ok kp
  | None -> Error (Error.Invalid_state "missing local ephemeral key")

let get_s st =
  match st.s with
  | Some kp -> Ok kp
  | None -> Error (Error.Invalid_state "missing local static key")

let get_re st =
  match st.re with
  | Some k -> Ok k
  | None -> Error (Error.Invalid_state "missing remote ephemeral key")

let get_rs st =
  match st.rs with
  | Some k -> Ok k
  | None -> Error (Error.Invalid_state "missing remote static key")

let write_token st buf tok =
  let h = hash st in
  let c = cipher st in
  match tok with
  | Pattern.E ->
      let* e_kp =
        match st.e with Some kp -> Ok kp | None -> (dh st).Dh.generate ()
      in
      let pub = e_kp.Dh.pub in
      Buffer.add_string buf pub;
      let sym = Symmetric_state.mix_hash st.sym ~hash:h pub in
      let sym =
        if st.psk_mode then Symmetric_state.mix_key sym ~hash:h pub else sym
      in
      Ok ({ st with sym; e = Some e_kp }, buf)
  | Pattern.S ->
      let* s_kp = get_s st in
      let* ct, sym =
        Symmetric_state.encrypt_and_hash st.sym ~hash:h ~cipher:c s_kp.Dh.pub
      in
      Buffer.add_string buf ct;
      Ok ({ st with sym }, buf)
  | Pattern.EE ->
      let* e_kp = get_e st in
      let* re_pk = get_re st in
      let* shared = do_dh st e_kp re_pk in
      let sym = Symmetric_state.mix_key st.sym ~hash:h shared in
      Ok ({ st with sym }, buf)
  | Pattern.ES ->
      (* Initiator: DH(e, rs);  Responder: DH(s, re) *)
      let* shared =
        match st.role with
        | Initiator ->
            let* e_kp = get_e st in
            let* rs_pk = get_rs st in
            do_dh st e_kp rs_pk
        | Responder ->
            let* s_kp = get_s st in
            let* re_pk = get_re st in
            do_dh st s_kp re_pk
      in
      let sym = Symmetric_state.mix_key st.sym ~hash:h shared in
      Ok ({ st with sym }, buf)
  | Pattern.SE ->
      (* Initiator: DH(s, re);  Responder: DH(e, rs) *)
      let* shared =
        match st.role with
        | Initiator ->
            let* s_kp = get_s st in
            let* re_pk = get_re st in
            do_dh st s_kp re_pk
        | Responder ->
            let* e_kp = get_e st in
            let* rs_pk = get_rs st in
            do_dh st e_kp rs_pk
      in
      let sym = Symmetric_state.mix_key st.sym ~hash:h shared in
      Ok ({ st with sym }, buf)
  | Pattern.SS ->
      let* s_kp = get_s st in
      let* rs_pk = get_rs st in
      let* shared = do_dh st s_kp rs_pk in
      let sym = Symmetric_state.mix_key st.sym ~hash:h shared in
      Ok ({ st with sym }, buf)
  | Pattern.PSK -> (
      match st.psks with
      | [] -> Error (Error.Bad_psk "no pre-shared key left for PSK token")
      | psk :: psks' ->
          let sym = Symmetric_state.mix_key_and_hash st.sym ~hash:h psk in
          Ok ({ st with sym; psks = psks' }, buf))

let read_token st msg pos tok =
  let h = hash st in
  let c = cipher st in
  match tok with
  | Pattern.E ->
      let dhlen = (dh st).Dh.dhlen in
      if String.length msg - !pos < dhlen then
        Error (Error.Invalid_state "message too short for E token")
      else begin
        let pub = String.sub msg !pos dhlen in
        pos := !pos + dhlen;
        let sym = Symmetric_state.mix_hash st.sym ~hash:h pub in
        let sym =
          if st.psk_mode then Symmetric_state.mix_key sym ~hash:h pub else sym
        in
        Ok { st with sym; re = Some pub }
      end
  | Pattern.S ->
      let dhlen = (dh st).Dh.dhlen in
      let tag_len =
        if Cipher_state.has_key st.sym.Symmetric_state.cs then 16 else 0
      in
      let needed = dhlen + tag_len in
      if String.length msg - !pos < needed then
        Error (Error.Invalid_state "message too short for S token")
      else begin
        let ct = String.sub msg !pos needed in
        pos := !pos + needed;
        let* pub, sym =
          Symmetric_state.decrypt_and_hash st.sym ~hash:h ~cipher:c ct
        in
        Ok { st with sym; rs = Some pub }
      end
  | Pattern.EE ->
      let* re_pk = get_re st in
      let* e_kp = get_e st in
      let* shared = do_dh st e_kp re_pk in
      let sym = Symmetric_state.mix_key st.sym ~hash:h shared in
      Ok { st with sym }
  | Pattern.ES ->
      let* shared =
        match st.role with
        | Initiator ->
            let* e_kp = get_e st in
            let* rs_pk = get_rs st in
            do_dh st e_kp rs_pk
        | Responder ->
            let* s_kp = get_s st in
            let* re_pk = get_re st in
            do_dh st s_kp re_pk
      in
      let sym = Symmetric_state.mix_key st.sym ~hash:h shared in
      Ok { st with sym }
  | Pattern.SE ->
      let* shared =
        match st.role with
        | Initiator ->
            let* s_kp = get_s st in
            let* re_pk = get_re st in
            do_dh st s_kp re_pk
        | Responder ->
            let* e_kp = get_e st in
            let* rs_pk = get_rs st in
            do_dh st e_kp rs_pk
      in
      let sym = Symmetric_state.mix_key st.sym ~hash:h shared in
      Ok { st with sym }
  | Pattern.SS ->
      let* s_kp = get_s st in
      let* rs_pk = get_rs st in
      let* shared = do_dh st s_kp rs_pk in
      let sym = Symmetric_state.mix_key st.sym ~hash:h shared in
      Ok { st with sym }
  | Pattern.PSK -> (
      match st.psks with
      | [] -> Error (Error.Bad_psk "no pre-shared key left for PSK token")
      | psk :: psks' ->
          let sym = Symmetric_state.mix_key_and_hash st.sym ~hash:h psk in
          Ok { st with sym; psks = psks' })

let check_message_length bytes =
  let n = String.length bytes in
  if n > 65535 then Error (Error.Bad_message_length { max = 65535; got = n })
  else Ok ()

let write_message st payload =
  if st.finished then Error (Error.Invalid_state "handshake already finished")
  else if st.role <> st.turn then
    Error (Error.Invalid_state "out of turn: expected read_message")
  else
    match st.patterns with
    | [] -> Error (Error.Invalid_state "no more message patterns")
    | msg_pattern :: rest ->
        let buf = Buffer.create 256 in
        let* st', _ =
          List.fold_left
            (fun acc tok ->
              let* st_acc, buf_acc = acc in
              write_token st_acc buf_acc tok)
            (Ok (st, buf))
            msg_pattern
        in
        let c = cipher st' in
        let h = hash st' in
        let* ct_payload, sym' =
          Symmetric_state.encrypt_and_hash st'.sym ~hash:h ~cipher:c payload
        in
        Buffer.add_string buf ct_payload;
        let message = Buffer.contents buf in
        let* () = check_message_length message in
        (* The handshake hash is h after the final payload encryption: "h at the
       time Split() is called" (spec §11.2). *)
        let finished, h_final, patterns' =
          match rest with
          | [] -> (true, Some sym'.Symmetric_state.h, [])
          | _ -> (false, None, rest)
        in
        let st'' =
          {
            st' with
            sym = sym';
            patterns = patterns';
            finished;
            h_final;
            turn = peer st.turn;
          }
        in
        Ok (message, st'')

let read_message st message =
  if st.finished then Error (Error.Invalid_state "handshake already finished")
  else if st.role = st.turn then
    Error (Error.Invalid_state "out of turn: expected write_message")
  else begin
    let* () = check_message_length message in
    match st.patterns with
    | [] -> Error (Error.Invalid_state "no more message patterns")
    | msg_pattern :: rest ->
        let pos = ref 0 in
        let* st' =
          List.fold_left
            (fun acc tok ->
              let* st_acc = acc in
              read_token st_acc message pos tok)
            (Ok st) msg_pattern
        in
        let ct_payload =
          String.sub message !pos (String.length message - !pos)
        in
        let c = cipher st' in
        let h = hash st' in
        let* payload, sym' =
          Symmetric_state.decrypt_and_hash st'.sym ~hash:h ~cipher:c ct_payload
        in
        let finished, h_final, patterns' =
          match rest with
          | [] -> (true, Some sym'.Symmetric_state.h, [])
          | _ -> (false, None, rest)
        in
        let st'' =
          {
            st' with
            sym = sym';
            patterns = patterns';
            finished;
            h_final;
            turn = peer st.turn;
          }
        in
        Ok (payload, st'')
  end

let is_finished st = st.finished
let handshake_hash st = st.h_final

let split st =
  if not st.finished then
    Error (Error.Invalid_state "handshake not yet finished")
  else
    let cs1, cs2 = Symmetric_state.split st.sym ~hash:(hash st) in
    Ok (cs1, cs2)

let role_of st = st.role
let cipher_of st = st.suite.Suite.cipher
let one_way_of st = st.one_way