1use std::io::ErrorKind;
4use std::path::{Path, PathBuf};
5
6use anyhow::{Result, bail};
7use serde::{Deserialize, Serialize};
8
9use crate::identity::{AgentId, AgentKey, MessageKind, SignedMessage};
10use crate::route::Route;
11use crate::state::{atomic_write, lock};
12
13const CAP: usize = 64;
14
15#[derive(Clone, Serialize, Deserialize)]
16pub struct Request {
17 pub key: String,
18 pub name: String,
19 pub request_id: String,
20 pub reply_to: String,
21}
22
23impl Request {
24 pub fn inbound(msg: &SignedMessage) -> Result<Self> {
25 let route = msg.reply_to.as_deref().map(Route::parse);
26 let Some(route) = route.filter(|route| route.key == msg.from && route.session.is_some())
27 else {
28 bail!("pairing request has no return session for its sender");
29 };
30 Ok(Self {
31 key: msg.from.clone(),
32 name: msg.text.clone(),
33 request_id: msg.msg_id.clone(),
34 reply_to: route.to_string(),
35 })
36 }
37}
38
39#[derive(Clone, Serialize, Deserialize)]
40pub struct ControlMessage {
41 pub route: String,
42 pub msg: SignedMessage,
43}
44
45impl ControlMessage {
46 pub fn for_delivery(&self, key: &AgentKey, now: u64) -> Result<SignedMessage> {
47 if self.msg.verify()? != key.id()
48 || !matches!(
49 self.msg.kind,
50 MessageKind::PairRequest | MessageKind::PairAccept
51 )
52 {
53 bail!("pairing outbox contains an invalid control message");
54 }
55 let mut msg = key.sign_full(
58 AgentId::from_b64(&self.msg.to)?,
59 &self.msg.text,
60 now,
61 &self.msg.msg_id,
62 self.msg.kind,
63 self.msg.task_id.as_deref(),
64 self.msg.status,
65 self.msg.in_reply_to.as_deref(),
66 );
67 msg.reply_to = self.msg.reply_to.clone();
68 Ok(msg)
69 }
70}
71
72#[derive(Default, Serialize, Deserialize)]
73struct State {
74 inbound: Vec<Request>,
75 outbound: Vec<Request>,
76 queued: Vec<ControlMessage>,
77}
78
79pub enum Acceptance {
80 Queued,
81 ConfirmationPending(String),
82}
83
84pub struct PairingStore {
85 path: PathBuf,
86}
87
88impl PairingStore {
89 pub fn new(path: &Path) -> Self {
90 Self {
91 path: path.to_owned(),
92 }
93 }
94
95 fn read(&self) -> Result<State> {
96 match std::fs::read(&self.path) {
97 Ok(bytes) => Ok(serde_json::from_slice(&bytes)?),
98 Err(e) if e.kind() == ErrorKind::NotFound => Ok(State::default()),
99 Err(e) => Err(e.into()),
100 }
101 }
102
103 fn update<T>(&self, change: impl FnOnce(&mut State) -> Result<T>) -> Result<T> {
104 let _lock = lock(&self.path.with_extension("lock"))?;
105 let mut state = self.read()?;
106 let result = change(&mut state)?;
107 atomic_write(&self.path, &serde_json::to_vec(&state)?)?;
108 Ok(result)
109 }
110
111 pub fn inbound(&self) -> Result<Vec<Request>> {
112 Ok(self.read()?.inbound)
113 }
114
115 pub fn find(&self, fingerprint: &str) -> Result<Option<Request>> {
116 let requests = self.inbound()?;
117 let mut matches = requests.into_iter().filter(|r| {
118 r.key == fingerprint || r.key.chars().take(8).collect::<String>() == fingerprint
119 });
120 let first = matches.next();
121 if matches.next().is_some() {
122 bail!("ambiguous fingerprint; use the full key");
123 }
124 Ok(first)
125 }
126
127 pub fn receive(&self, request: Request) -> Result<()> {
128 self.update(|s| {
129 put(&mut s.inbound, request);
130 Ok(())
131 })
132 }
133
134 pub fn request(
135 &self,
136 mut request: Request,
137 message: impl FnOnce(&Request) -> ControlMessage,
138 ) -> Result<()> {
139 self.update(|s| {
140 if let Some(old) = s.outbound.iter().find(|r| r.key == request.key) {
143 request.request_id = old.request_id.clone();
144 }
145 s.queued.retain(|job| job.msg.msg_id != request.request_id);
146 enqueue(s, message(&request))?;
147 put(&mut s.outbound, request);
148 Ok(())
149 })
150 }
151
152 pub fn accept(
153 &self,
154 request: &Request,
155 message: ControlMessage,
156 authorize: impl FnOnce() -> Result<()>,
157 ) -> Result<Acceptance> {
158 self.accept_with_commit(request, message, authorize, atomic_write)
159 }
160
161 fn accept_with_commit(
162 &self,
163 request: &Request,
164 message: ControlMessage,
165 authorize: impl FnOnce() -> Result<()>,
166 commit: impl FnOnce(&Path, &[u8]) -> Result<()>,
167 ) -> Result<Acceptance> {
168 let _lock = lock(&self.path.with_extension("lock"))?;
169 let mut state = self.read()?;
170 if !state
171 .inbound
172 .iter()
173 .any(|r| r.key == request.key && r.request_id == request.request_id)
174 {
175 bail!("pairing request changed; review the pending requests again");
176 }
177 enqueue(&mut state, message)?;
180 state.inbound.retain(|r| r.key != request.key);
181 let bytes = serde_json::to_vec(&state)?;
182 authorize()?;
183 match commit(&self.path, &bytes) {
184 Ok(()) => Ok(Acceptance::Queued),
185 Err(e) => Ok(Acceptance::ConfirmationPending(e.to_string())),
188 }
189 }
190
191 pub fn reject(&self, key: &str) -> Result<()> {
192 self.update(|s| {
193 s.inbound.retain(|r| r.key != key);
194 Ok(())
195 })
196 }
197
198 pub fn pending_accept(&self, key: &str, in_reply_to: Option<&str>) -> Result<Option<Request>> {
199 Ok(self
200 .read()?
201 .outbound
202 .into_iter()
203 .find(|r| r.key == key && in_reply_to.is_none_or(|id| id == r.request_id)))
204 }
205
206 pub fn complete(&self, request: &Request) -> Result<()> {
207 self.update(|s| {
208 s.outbound
209 .retain(|r| r.key != request.key || r.request_id != request.request_id);
210 Ok(())
211 })
212 }
213
214 pub fn queued(&self) -> Result<Vec<ControlMessage>> {
215 Ok(self.read()?.queued)
216 }
217
218 pub fn sent(&self, msg_id: &str) -> Result<()> {
219 self.update(|s| {
220 s.queued.retain(|job| job.msg.msg_id != msg_id);
221 Ok(())
222 })
223 }
224}
225
226fn put(requests: &mut Vec<Request>, request: Request) {
227 requests.retain(|r| r.key != request.key);
228 if requests.len() >= CAP {
229 requests.remove(0);
230 }
231 requests.push(request);
232}
233
234fn enqueue(state: &mut State, message: ControlMessage) -> Result<()> {
235 if state.queued.len() >= CAP * 2 {
236 bail!("pairing outbox is full; retry after the relay recovers");
237 }
238 state.queued.push(message);
239 Ok(())
240}
241
242#[cfg(test)]
243mod tests {
244 use super::*;
245 use crate::identity::{AgentKey, MessageKind};
246 use crate::policy_store::PolicyStore;
247
248 #[test]
249 fn renewing_control_messages_preserves_identity_and_metadata() {
250 let alice = AgentKey::generate().unwrap();
251 let bob = AgentKey::generate().unwrap();
252 for kind in [MessageKind::PairRequest, MessageKind::PairAccept] {
253 let mut msg = alice.sign_full(
254 bob.id(),
255 "alice",
256 1,
257 "original",
258 kind,
259 None,
260 None,
261 Some("request"),
262 );
263 msg.reply_to = Some(format!("{}#session", alice.id().to_b64()));
264 let job = ControlMessage {
265 route: format!("{}#recipient", bob.id().to_b64()),
266 msg,
267 };
268 let renewed = job.for_delivery(&alice, 200_000_000).unwrap();
269 assert_eq!(renewed.verify().unwrap(), alice.id());
270 assert_eq!(renewed.ts, 200_000_000);
271 assert_ne!(renewed.sig, job.msg.sig);
272 let mut fields = serde_json::to_value(&renewed).unwrap();
273 fields["ts"] = serde_json::json!(job.msg.ts);
274 fields["sig"] = serde_json::json!(job.msg.sig);
275 assert_eq!(fields, serde_json::to_value(&job.msg).unwrap());
276 assert!(job.for_delivery(&bob, 200_000_000).is_err());
277 }
278 }
279
280 #[test]
281 fn requests_and_acceptances_survive_restart_with_exact_return_route() {
282 let dir = tempfile::tempdir().unwrap();
283 let path = dir.path().join("pairing.json");
284 let alice = AgentKey::generate().unwrap();
285 let bob = AgentKey::generate().unwrap();
286 let mut msg = alice.sign_as(bob.id(), "alice", 1, "request", MessageKind::PairRequest);
287 msg.reply_to = Some(format!("{}#second-session", alice.id().to_b64()));
288 let request = Request::inbound(&msg).unwrap();
289 let store = PairingStore::new(&path);
290 store.receive(request.clone()).unwrap();
291 let reopened = PairingStore::new(&path);
292 let pending = reopened.find(&alice.id().to_b64()).unwrap().unwrap();
293 assert!(pending.reply_to.ends_with("#second-session"));
294 let accept = bob.sign_full(
295 alice.id(),
296 "bob",
297 2,
298 "accept",
299 MessageKind::PairAccept,
300 None,
301 None,
302 Some("request"),
303 );
304 reopened
305 .accept(
306 &pending,
307 ControlMessage {
308 route: pending.reply_to.clone(),
309 msg: accept,
310 },
311 || Ok(()),
312 )
313 .unwrap();
314 let reopened = PairingStore::new(&path);
315 assert!(reopened.inbound().unwrap().is_empty());
316 let queued = reopened.queued().unwrap();
317 assert_eq!(queued[0].route, request.reply_to);
318 assert_eq!(queued[0].msg.in_reply_to.as_deref(), Some("request"));
319 reopened.sent("accept").unwrap();
320 assert!(store.queued().unwrap().is_empty());
321 }
322 #[test]
323 fn failed_confirmation_is_explicit_and_can_be_retried() {
324 let dir = tempfile::tempdir().unwrap();
325 let pairing = PairingStore::new(&dir.path().join("pairing.json"));
326 let alice = AgentKey::generate().unwrap();
327 let bob = AgentKey::generate().unwrap();
328 let policy_path = dir.path().join("peers.json");
329 std::fs::write(&policy_path, "{}").unwrap();
330 let policy = PolicyStore::open(&policy_path).unwrap();
331 let request = Request {
332 key: alice.id().to_b64(),
333 name: "alice".into(),
334 request_id: "request".into(),
335 reply_to: format!("{}#session", alice.id().to_b64()),
336 };
337 pairing.receive(request.clone()).unwrap();
338 let message = ControlMessage {
339 route: request.reply_to.clone(),
340 msg: bob.sign_as(alice.id(), "bob", 1, "accept", MessageKind::PairAccept),
341 };
342 let outcome = pairing
343 .accept_with_commit(
344 &request,
345 message.clone(),
346 || policy.add(&request.name, &request.key),
347 |_, _| bail!("simulated storage failure after authorization"),
348 )
349 .unwrap();
350 assert!(matches!(outcome, Acceptance::ConfirmationPending(_)));
351 assert!(policy.read().unwrap().resolve("alice").is_ok());
352 assert!(pairing.find(&request.key).unwrap().is_some());
353 assert!(pairing.queued().unwrap().is_empty());
354 assert!(matches!(
355 pairing
356 .accept(&request, message, || policy
357 .add(&request.name, &request.key))
358 .unwrap(),
359 Acceptance::Queued
360 ));
361 assert!(pairing.inbound().unwrap().is_empty());
362 assert_eq!(pairing.queued().unwrap().len(), 1);
363 }
364
365 #[test]
366 fn full_outbox_never_runs_authorization() {
367 let dir = tempfile::tempdir().unwrap();
368 let pairing = PairingStore::new(&dir.path().join("pairing.json"));
369 let alice = AgentKey::generate().unwrap();
370 let bob = AgentKey::generate().unwrap();
371 let request = Request {
372 key: alice.id().to_b64(),
373 name: "alice".into(),
374 request_id: "request".into(),
375 reply_to: format!("{}#session", alice.id().to_b64()),
376 };
377 let message = ControlMessage {
378 route: request.reply_to.clone(),
379 msg: bob.sign_as(alice.id(), "bob", 1, "accept", MessageKind::PairAccept),
380 };
381 pairing.receive(request.clone()).unwrap();
382 pairing
383 .update(|s| {
384 s.queued = vec![message.clone(); CAP * 2];
385 Ok(())
386 })
387 .unwrap();
388 let mut authorized = false;
389 assert!(
390 pairing
391 .accept(&request, message, || {
392 authorized = true;
393 Ok(())
394 })
395 .is_err()
396 );
397 assert!(!authorized);
398 assert!(pairing.find(&request.key).unwrap().is_some());
399 }
400}