1use crate::cbor::{self, Value};
7use crate::node_key::{node_id_of, NodeKey};
8use crate::profile::Profile;
9use crate::signed_object::{sign_object, verify_object, Object};
10
11use super::{
12 bounded_text, check_payload, entry, fixed, has_fields, identity_signer, names_request,
13 object_refusal, read_fields, received_frame, text_of, FrameError, Rule, VerifiedRequest,
14 MAX_ERROR_CODE_BYTES, MAX_ERROR_TEXT_BYTES, PROTOCOL_VERSION, RELAY_ERROR_LABEL, REPLY_LABEL,
15};
16
17const RESULT: &str = "result";
18const ERROR: &str = "error";
19const STREAM_ERROR: &str = "stream_error";
20
21const RELAY_CODES: &[&str] = &["unknown_next_peer"];
23
24#[derive(Debug, Clone, Copy, PartialEq, Eq)]
26pub enum ReplyType {
27 Result,
28 Error,
29}
30
31#[derive(Debug, Clone, Copy, PartialEq, Eq)]
33pub enum RelayErrorType {
34 Error,
35 StreamError,
36}
37
38impl RelayErrorType {
39 fn name(self) -> &'static str {
40 match self {
41 RelayErrorType::Error => ERROR,
42 RelayErrorType::StreamError => STREAM_ERROR,
43 }
44 }
45}
46
47#[derive(Debug, Clone, PartialEq)]
50pub struct VerifiedReply {
51 pub frame_type: ReplyType,
52 pub responded_by: [u8; 32],
53 pub payload: Option<Value>,
54 pub code: Option<String>,
55 pub detail: Option<String>,
56}
57
58#[derive(Debug, Clone, PartialEq)]
62pub struct RelayErrorSpec {
63 pub frame_type: RelayErrorType,
64 pub request: VerifiedRequest,
65 pub code: String,
66 pub offending_hop: Option<[u8; 32]>,
67 pub source_route_partial: Option<Vec<u8>>,
68}
69
70#[derive(Debug, Clone, PartialEq, Eq)]
73pub struct VerifiedRelayError {
74 pub frame_type: RelayErrorType,
75 pub reported_by: [u8; 32],
76 pub code: String,
77 pub offending_hop: Option<[u8; 32]>,
78}
79
80pub fn sign_result(
84 request: &VerifiedRequest,
85 payload: &Value,
86 source_route_reverse: Option<Vec<u8>>,
87 key: &NodeKey,
88) -> Result<Value, FrameError> {
89 reply_signer(request, key)?;
90 check_payload(payload)?;
91 sign_reply(
92 ReplyType::Result,
93 request,
94 vec![entry("payload", payload.clone())],
95 source_route_reverse,
96 key,
97 )
98}
99
100pub fn sign_provider_error(
103 request: &VerifiedRequest,
104 code: &str,
105 detail: Option<&str>,
106 source_route_reverse: Option<Vec<u8>>,
107 key: &NodeKey,
108) -> Result<Value, FrameError> {
109 reply_signer(request, key)?;
110 bounded_text("code", code.as_bytes(), MAX_ERROR_CODE_BYTES)?;
111 let mut fields = vec![entry("code", Value::text(code))];
112 if let Some(detail) = detail {
113 bounded_text("detail", detail.as_bytes(), MAX_ERROR_TEXT_BYTES)?;
114 fields.push(entry("detail", Value::text(detail)));
115 }
116 sign_reply(ReplyType::Error, request, fields, source_route_reverse, key)
117}
118
119fn reply_signer(request: &VerifiedRequest, key: &NodeKey) -> Result<(), FrameError> {
120 identity_signer(key)?;
121 if key.key_id() != request.target {
122 return Err(FrameError::Unsignable);
123 }
124 Ok(())
125}
126
127fn sign_reply(
128 frame_type: ReplyType,
129 request: &VerifiedRequest,
130 mut fields: Vec<(Value, Value)>,
131 source_route_reverse: Option<Vec<u8>>,
132 key: &NodeKey,
133) -> Result<Value, FrameError> {
134 let name = match frame_type {
135 ReplyType::Result => RESULT,
136 ReplyType::Error => ERROR,
137 };
138 fields.extend([
139 entry("frame_type", Value::text(name)),
140 entry("request_id", Value::Bytes(request.request_id.to_vec())),
141 entry("request_hash", Value::Bytes(request.request_hash.to_vec())),
142 entry("responded_by", Value::Bytes(key.key_id().to_vec())),
143 ]);
144 let reply = sign_object(REPLY_LABEL, &fields, key).map_err(object_refusal)?;
145 Ok(routed_frame(
146 name,
147 "reply",
148 &reply,
149 "source_route_reverse",
150 source_route_reverse,
151 ))
152}
153
154const REPLY_ROUTES: &[(&str, Rule)] = &[("source_route_reverse", Rule::AnyBytes)];
155const RELAY_ERROR_ROUTES: &[(&str, Rule)] = &[("source_route_partial", Rule::AnyBytes)];
156
157pub fn verify_reply(
162 frame: &Value,
163 request: &VerifiedRequest,
164 profile: Profile,
165) -> Result<VerifiedReply, FrameError> {
166 let (frame_type, object) = received_frame(
167 frame,
168 "reply",
169 Rule::CarriedObject,
170 REPLY_ROUTES,
171 &[RESULT, ERROR],
172 )
173 .ok_or(FrameError::Malformed)?;
174 let verified = verify_object(REPLY_LABEL, &object, profile).map_err(object_refusal)?;
175 let fields =
176 read_fields(&verified.fields, &reply_table(&frame_type)).ok_or(FrameError::Malformed)?;
177 let (has_payload, has_code, has_detail) = (
178 fields.contains_key("payload"),
179 fields.contains_key("code"),
180 fields.contains_key("detail"),
181 );
182 let shaped = if frame_type == RESULT {
183 has_payload && !has_code && !has_detail
184 } else {
185 has_code && !has_payload
186 };
187 if !has_fields(
188 &fields,
189 &["frame_type", "request_id", "request_hash", "responded_by"],
190 ) || !shaped
191 {
192 return Err(FrameError::Malformed);
193 }
194 let reply = VerifiedReply {
195 frame_type: if frame_type == RESULT {
196 ReplyType::Result
197 } else {
198 ReplyType::Error
199 },
200 responded_by: fixed(&fields["responded_by"]),
201 payload: fields.get("payload").cloned(),
202 code: fields.get("code").map(text_of),
203 detail: fields.get("detail").map(text_of),
204 };
205 if reply.responded_by != node_id_of(&verified.key, profile) {
206 return Err(FrameError::KeyIdMismatch);
207 }
208 if !names_request(&fields, request) {
209 return Err(FrameError::RequestMismatch);
210 }
211 if reply.responded_by != request.target {
212 return Err(FrameError::NotTheTarget);
213 }
214 Ok(reply)
215}
216
217fn reply_table(frame_type: &str) -> Vec<(&'static str, Rule)> {
218 vec![
219 (
220 "frame_type",
221 Rule::TextIn(if frame_type == RESULT {
222 &[RESULT]
223 } else {
224 &[ERROR]
225 }),
226 ),
227 ("alg", Rule::Any),
228 ("request_id", Rule::BytesOf(16)),
229 ("request_hash", Rule::BytesOf(48)),
230 ("responded_by", Rule::BytesOf(32)),
231 ("payload", Rule::Any),
232 ("code", Rule::TextWithin(MAX_ERROR_CODE_BYTES)),
233 ("detail", Rule::TextWithin(MAX_ERROR_TEXT_BYTES)),
234 ]
235}
236
237pub fn sign_relay_error(spec: &RelayErrorSpec, key: &NodeKey) -> Result<Value, FrameError> {
241 identity_signer(key)?;
242 if !RELAY_CODES.contains(&spec.code.as_str()) {
243 return Err(FrameError::RelayCodeOutsideItsSet);
244 }
245 let name = spec.frame_type.name();
246 let mut fields = vec![
247 entry("frame_type", Value::text(name)),
248 entry("request_id", Value::Bytes(spec.request.request_id.to_vec())),
249 entry(
250 "request_hash",
251 Value::Bytes(spec.request.request_hash.to_vec()),
252 ),
253 entry("reported_by", Value::Bytes(key.key_id().to_vec())),
254 entry("code", Value::text(spec.code.clone())),
255 ];
256 if let Some(hop) = spec.offending_hop {
257 fields.push(entry("offending_hop", Value::Bytes(hop.to_vec())));
258 }
259 let relay_error = sign_object(RELAY_ERROR_LABEL, &fields, key).map_err(object_refusal)?;
260 Ok(routed_frame(
261 name,
262 "relay_error",
263 &relay_error,
264 "source_route_partial",
265 spec.source_route_partial.clone(),
266 ))
267}
268
269pub fn verify_relay_error(
272 frame: &Value,
273 request: &VerifiedRequest,
274 profile: Profile,
275 expected_reporter: &[u8; 32],
276) -> Result<VerifiedRelayError, FrameError> {
277 let (frame_type, object) = received_frame(
278 frame,
279 "relay_error",
280 Rule::CarriedObject,
281 RELAY_ERROR_ROUTES,
282 &[ERROR, STREAM_ERROR],
283 )
284 .ok_or(FrameError::Malformed)?;
285 let verified = verify_object(RELAY_ERROR_LABEL, &object, profile).map_err(object_refusal)?;
286 let fields = read_fields(&verified.fields, &relay_error_table(&frame_type))
287 .ok_or(FrameError::Malformed)?;
288 if !has_fields(
289 &fields,
290 &[
291 "frame_type",
292 "request_id",
293 "request_hash",
294 "reported_by",
295 "code",
296 ],
297 ) {
298 return Err(FrameError::Malformed);
299 }
300 let relay_error = VerifiedRelayError {
301 frame_type: if frame_type == ERROR {
302 RelayErrorType::Error
303 } else {
304 RelayErrorType::StreamError
305 },
306 reported_by: fixed(&fields["reported_by"]),
307 code: text_of(&fields["code"]),
308 offending_hop: fields.get("offending_hop").map(fixed),
309 };
310 if relay_error.reported_by != node_id_of(&verified.key, profile) {
311 return Err(FrameError::KeyIdMismatch);
312 }
313 if !names_request(&fields, request) {
314 return Err(FrameError::RequestMismatch);
315 }
316 if &relay_error.reported_by != expected_reporter {
317 return Err(FrameError::NotTheConnection);
318 }
319 Ok(relay_error)
320}
321
322fn relay_error_table(frame_type: &str) -> Vec<(&'static str, Rule)> {
323 vec![
324 (
325 "frame_type",
326 Rule::TextIn(if frame_type == ERROR {
327 &[ERROR]
328 } else {
329 &[STREAM_ERROR]
330 }),
331 ),
332 ("alg", Rule::Any),
333 ("request_id", Rule::BytesOf(16)),
334 ("request_hash", Rule::BytesOf(48)),
335 ("reported_by", Rule::BytesOf(32)),
336 ("code", Rule::TextIn(RELAY_CODES)),
337 ("offending_hop", Rule::BytesOf(32)),
338 ]
339}
340
341pub fn claimed_reply_ids(frame: &Value) -> Result<([u8; 16], [u8; 48]), FrameError> {
347 let frame_type = frame.get("frame_type").map(text_of).unwrap_or_default();
348 let (object_name, routes, table) = match (frame.get("reply"), frame.get("relay_error")) {
349 (Some(_), _) if frame_type == RESULT || frame_type == ERROR => {
350 ("reply", REPLY_ROUTES, reply_table(&frame_type))
351 }
352 (_, Some(_)) if frame_type == ERROR || frame_type == STREAM_ERROR => (
353 "relay_error",
354 RELAY_ERROR_ROUTES,
355 relay_error_table(&frame_type),
356 ),
357 _ => return Err(FrameError::Malformed),
358 };
359 let types: &'static [&'static str] = match frame_type.as_str() {
360 RESULT => &[RESULT],
361 ERROR => &[ERROR],
362 _ => &[STREAM_ERROR],
363 };
364 let (_, object) = received_frame(frame, object_name, Rule::CarriedObject, routes, types)
365 .ok_or(FrameError::Malformed)?;
366 let parsed = Object::from_value(&object).map_err(|_| FrameError::Malformed)?;
367 let tbs = cbor::decode(&parsed.tbs).map_err(|_| FrameError::Malformed)?;
368 let fields = read_fields(&tbs, &table).ok_or(FrameError::Malformed)?;
369 if !has_fields(&fields, &["frame_type", "request_id", "request_hash"]) {
370 return Err(FrameError::Malformed);
371 }
372 Ok((fixed(&fields["request_id"]), fixed(&fields["request_hash"])))
373}
374
375fn routed_frame(
378 frame_type: &str,
379 object_name: &str,
380 object: &Object,
381 route_name: &str,
382 route: Option<Vec<u8>>,
383) -> Value {
384 let mut entries = vec![
385 entry("version", Value::Int(i128::from(PROTOCOL_VERSION))),
386 entry("frame_type", Value::text(frame_type)),
387 entry(object_name, object.to_value()),
388 ];
389 if let Some(route) = route {
390 entries.push(entry(route_name, Value::Bytes(route)));
391 }
392 Value::Map(entries)
393}