use std::future::Future;
use std::path::PathBuf;
use std::pin::Pin;
use std::sync::Arc;
use anyhow::{Context, bail};
use tokio::io::BufReader;
use mcpmesh_local_api::PairResult;
use mcpmesh_net::framing::{FrameReader, Inbound, write_frame};
use crate::allowlist::{PeerEntry, PeerStore};
use crate::pairing::sas::short_auth_code;
use crate::pairing::{Invite, LiveInvites, Redeem};
use crate::util::epoch_now_u64 as epoch_now;
const MAX_PAIR_FRAME: usize = 64 * 1024;
const REASON_REFUSED: &str = "pairing refused";
const REASON_MALFORMED: &str = "malformed request";
const REASON_ID_MISMATCH: &str = "id mismatch";
pub(crate) const NO_LIVE_INVITE_CLOSE: &[u8] = b"no pairing in progress";
fn reason_nickname_taken(nickname: &str, invite_survived: bool) -> String {
let recovery = if invite_survived {
"the invite was NOT consumed — rename this node and redeem the same invite again"
} else {
"ask the inviter for a fresh invite"
};
format!("nickname '{nickname}' is already taken by another paired peer; {recovery}")
}
fn collision_refusal(nickname: &str, invite_survived: bool) -> PairReply {
PairReply::Refused {
reason: reason_nickname_taken(nickname, invite_survived),
code: invite_survived.then_some(RefusalCode::NicknameTaken),
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize)]
#[serde(rename_all = "snake_case")]
enum RefusalCode {
NicknameTaken,
Unknown,
}
impl<'de> serde::Deserialize<'de> for RefusalCode {
fn deserialize<D: serde::Deserializer<'de>>(d: D) -> Result<Self, D::Error> {
struct AnyCode;
impl<'de> serde::de::Visitor<'de> for AnyCode {
type Value = RefusalCode;
fn expecting(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result {
f.write_str("a refusal code")
}
fn visit_str<E: serde::de::Error>(self, s: &str) -> Result<Self::Value, E> {
Ok(match s {
"nickname_taken" => RefusalCode::NicknameTaken,
_ => RefusalCode::Unknown,
})
}
fn visit_unit<E: serde::de::Error>(self) -> Result<Self::Value, E> {
Ok(RefusalCode::Unknown)
}
fn visit_none<E: serde::de::Error>(self) -> Result<Self::Value, E> {
Ok(RefusalCode::Unknown)
}
fn visit_some<D: serde::Deserializer<'de>>(
self,
d: D,
) -> Result<Self::Value, D::Error> {
d.deserialize_any(AnyCode)
}
fn visit_bool<E: serde::de::Error>(self, _: bool) -> Result<Self::Value, E> {
Ok(RefusalCode::Unknown)
}
fn visit_i64<E: serde::de::Error>(self, _: i64) -> Result<Self::Value, E> {
Ok(RefusalCode::Unknown)
}
fn visit_u64<E: serde::de::Error>(self, _: u64) -> Result<Self::Value, E> {
Ok(RefusalCode::Unknown)
}
fn visit_f64<E: serde::de::Error>(self, _: f64) -> Result<Self::Value, E> {
Ok(RefusalCode::Unknown)
}
fn visit_map<A: serde::de::MapAccess<'de>>(
self,
mut m: A,
) -> Result<Self::Value, A::Error> {
while m
.next_entry::<serde::de::IgnoredAny, serde::de::IgnoredAny>()?
.is_some()
{}
Ok(RefusalCode::Unknown)
}
fn visit_seq<A: serde::de::SeqAccess<'de>>(
self,
mut s: A,
) -> Result<Self::Value, A::Error> {
while s.next_element::<serde::de::IgnoredAny>()?.is_some() {}
Ok(RefusalCode::Unknown)
}
}
d.deserialize_any(AnyCode)
}
}
#[derive(Debug)]
pub struct NicknameTaken(pub String);
impl std::fmt::Display for NicknameTaken {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(&self.0)
}
}
impl std::error::Error for NicknameTaken {}
fn no_live_invite_close(conn: &iroh::endpoint::Connection) -> bool {
matches!(
conn.close_reason(),
Some(iroh::endpoint::ConnectionError::ApplicationClosed(ac))
if ac.reason.as_ref() == NO_LIVE_INVITE_CLOSE
)
}
async fn resolve_and_check_collision(
store: &Arc<PeerStore>,
nickname: &str,
tls_id: [u8; 32],
) -> anyhow::Result<(Option<PeerEntry>, bool)> {
let store_c = store.clone();
let nickname_c = nickname.to_string();
tokio::task::spawn_blocking(move || {
let existing = store_c.resolve(&tls_id)?;
let collides = existing.is_none() && nickname_collision(&store_c, &nickname_c, &tls_id)?;
anyhow::Ok((existing, collides))
})
.await
.context("join nickname collision check")?
}
#[derive(serde::Serialize, serde::Deserialize)]
struct RedeemerHello {
secret: [u8; 32],
redeemer_id: [u8; 32],
redeemer_nickname: String,
#[serde(default)]
user_pk: Option<String>,
#[serde(default)]
binding_sig: Option<String>,
}
#[derive(serde::Serialize, serde::Deserialize)]
#[serde(tag = "result", rename_all = "snake_case")]
enum PairReply {
Ok {
inviter_id: [u8; 32],
inviter_nickname: String,
#[serde(default)]
user_pk: Option<String>,
#[serde(default)]
binding_sig: Option<String>,
},
Refused {
reason: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
code: Option<RefusalCode>,
},
}
#[derive(Clone, Debug)]
pub struct SelfBinding {
pub user_pk: String,
pub sig: String,
}
pub type GrantFn = Box<
dyn Fn(String, String, Vec<String>) -> Pin<Box<dyn Future<Output = anyhow::Result<()>> + Send>>
+ Send
+ Sync,
>;
pub type GrantBackFn = Box<
dyn Fn(String, String) -> Pin<Box<dyn Future<Output = anyhow::Result<()>> + Send>>
+ Send
+ Sync,
>;
pub type RecordPairingFn = Box<dyn Fn(String, String, u64) + Send + Sync>;
pub struct InviterCtx {
pub store: Arc<PeerStore>,
pub invites: Arc<LiveInvites>,
pub config_path: PathBuf,
pub self_binding: Option<SelfBinding>,
pub grant: GrantFn,
pub record_pairing: RecordPairingFn,
}
fn verified_user_id(
user_pk: &Option<String>,
binding_sig: &Option<String>,
authenticated_id: &[u8; 32],
) -> Option<String> {
match (user_pk, binding_sig) {
(Some(pk), Some(sig)) => {
match mcpmesh_trust::binding::verify_presented(pk, sig, authenticated_id) {
Ok(uid) => Some(uid),
Err(e) => {
tracing::warn!(
%e,
"peer presented an invalid device->user binding; storing entry without a user_id"
);
None
}
}
}
_ => None,
}
}
pub async fn handle_inviter_side(
conn: iroh::endpoint::Connection,
ctx: InviterCtx,
) -> anyhow::Result<()> {
let (mut send, recv) = conn.accept_bi().await?;
let mut reader = FrameReader::new(BufReader::new(recv), MAX_PAIR_FRAME);
let hello: RedeemerHello = match reader.next().await? {
Some(Inbound::Frame(v)) => match serde_json::from_value(v) {
Ok(h) => h,
Err(_) => return refuse(&mut send, REASON_MALFORMED, "malformed hello").await,
},
_ => return refuse(&mut send, REASON_MALFORMED, "malformed hello").await,
};
let tls_id = *conn.remote_id().as_bytes();
if tls_id != hello.redeemer_id {
return refuse(&mut send, REASON_ID_MISMATCH, "id mismatch").await;
}
let now = epoch_now();
if ctx.invites.peek_live(&hello.secret, now) {
let (_, collides) =
resolve_and_check_collision(&ctx.store, &hello.redeemer_nickname, tls_id).await?;
if collides {
tracing::warn!(
nickname = %hello.redeemer_nickname,
"pairing refused: nickname collision (invite preserved)"
);
let _ = send_reply(
&mut send,
&collision_refusal(&hello.redeemer_nickname, true),
)
.await;
return Ok(());
}
}
match ctx.invites.try_redeem(&hello.secret, now) {
Redeem::Ok(invite) => {
let (existing, collides) =
resolve_and_check_collision(&ctx.store, &hello.redeemer_nickname, tls_id).await?;
if collides {
tracing::warn!(
nickname = %hello.redeemer_nickname,
"pairing refused: nickname collision (post-redeem race guard; invite burned)"
);
let _ = send_reply(
&mut send,
&collision_refusal(&hello.redeemer_nickname, false),
)
.await;
return Ok(());
}
let nickname = existing
.as_ref()
.map_or_else(|| hello.redeemer_nickname.clone(), |e| e.nickname.clone());
let observed_addr = {
let addrs: Vec<iroh::TransportAddr> = conn
.paths()
.iter()
.map(|p| p.remote_addr().clone())
.collect();
if addrs.is_empty() {
None
} else {
serde_json::to_string(&iroh::EndpointAddr::from_parts(conn.remote_id(), addrs))
.ok()
}
};
let last_addr =
observed_addr.or_else(|| existing.as_ref().and_then(|e| e.last_addr.clone()));
let entry = PeerEntry {
endpoint_id: tls_id,
nickname: nickname.clone(),
services: existing
.as_ref()
.map(|e| e.services.clone())
.unwrap_or_default(),
paired_at: existing
.as_ref()
.and_then(|e| e.paired_at.clone())
.or_else(|| Some(now.to_string())),
user_id: verified_user_id(&hello.user_pk, &hello.binding_sig, &tls_id)
.or_else(|| existing.and_then(|e| e.user_id)),
last_addr,
};
let principal = entry
.user_id
.clone()
.unwrap_or_else(|| mcpmesh_net::EndpointId::from_bytes(tls_id).principal());
let store2 = ctx.store.clone();
tokio::task::spawn_blocking(move || store2.add(entry))
.await
.context("join pair store write")??;
(ctx.grant)(principal, nickname.clone(), invite.services.clone()).await?;
let sas = short_auth_code(&invite.inviter_id, &tls_id, &hello.secret);
tracing::info!(peer = %nickname, code = %sas, "paired");
(ctx.record_pairing)(nickname, sas, now);
let (inviter_pk, inviter_sig) = match ctx.self_binding {
Some(b) => (Some(b.user_pk), Some(b.sig)),
None => (None, None),
};
let _ = send_reply(
&mut send,
&PairReply::Ok {
inviter_id: invite.inviter_id,
inviter_nickname: invite.nickname.clone(),
user_pk: inviter_pk,
binding_sig: inviter_sig,
},
)
.await;
Ok(())
}
other => {
tracing::info!(outcome = ?other, "pair attempt refused");
let _ = send_reply(
&mut send,
&PairReply::Refused {
reason: REASON_REFUSED.into(),
code: None,
},
)
.await;
Ok(())
}
}
}
async fn refuse(
send: &mut iroh::endpoint::SendStream,
reason: &str,
log: &str,
) -> anyhow::Result<()> {
tracing::info!("pair attempt refused: {log}");
let _ = send_reply(
send,
&PairReply::Refused {
reason: reason.into(),
code: None,
},
)
.await;
Ok(())
}
async fn send_reply(
send: &mut iroh::endpoint::SendStream,
reply: &PairReply,
) -> anyhow::Result<()> {
write_frame(send, &serde_json::to_value(reply)?).await?;
let _ = send.finish();
let _ = send.stopped().await;
Ok(())
}
pub async fn redeem_invite(
endpoint: iroh::Endpoint,
self_nickname: String,
invite_line: String,
store: Arc<PeerStore>,
self_binding: Option<SelfBinding>,
grant_back: Option<GrantBackFn>,
) -> anyhow::Result<PairResult> {
let invite = Invite::decode(&invite_line)?;
if invite.expires_at_epoch < epoch_now() {
bail!("invite expired");
}
if let Some(conflict) = nickname_squat(&store, &invite.nickname, &invite.inviter_id)? {
bail!(
"this invite asks to be called '{}', but {conflict} \
Ask them for an invite suggesting a different name.",
invite.nickname,
);
}
let addr: iroh::EndpointAddr = serde_json::from_str(&invite.inviter_addr_json)
.context("invite carries an undecodable inviter address")?;
let conn = endpoint
.connect(addr, mcpmesh_net::ALPN_PAIR)
.await
.context("could not dial the inviter's machine")?;
if *conn.remote_id().as_bytes() != invite.inviter_id {
bail!("inviter id mismatch — refusing (address-swap defense)");
}
let (redeemer_pk, redeemer_sig) = match self_binding {
Some(b) => (Some(b.user_pk), Some(b.sig)),
None => (None, None),
};
let hello = RedeemerHello {
secret: invite.secret,
redeemer_id: *endpoint.id().as_bytes(),
redeemer_nickname: self_nickname,
user_pk: redeemer_pk,
binding_sig: redeemer_sig,
};
let exchange = async {
let (mut send, recv) = conn.open_bi().await.context("open the pairing bi-stream")?;
write_frame(&mut send, &serde_json::to_value(&hello)?)
.await
.context("send the pairing hello")?;
let mut reader = FrameReader::new(BufReader::new(recv), MAX_PAIR_FRAME);
match reader.next().await.context("read the pairing reply")? {
Some(Inbound::Frame(v)) => {
serde_json::from_value::<PairReply>(v).context("inviter reply is not a PairReply")
}
_ => bail!("no reply from the inviter (connection closed before a reply)"),
}
};
let reply: PairReply = match exchange.await {
Ok(reply) => reply,
Err(_) if no_live_invite_close(&conn) => {
bail!(
"the invite is no longer live on the inviter: it expired, was already \
redeemed, or the inviter's daemon restarted since minting it (invites do \
not survive a restart) — ask for a fresh invite"
);
}
Err(e) => return Err(e),
};
let inviter_user_id = match &reply {
PairReply::Refused { reason, code } => return Err(refusal_error(reason, *code)),
PairReply::Ok {
user_pk,
binding_sig,
..
} => verified_user_id(user_pk, binding_sig, &invite.inviter_id),
};
let peer_user_id = inviter_user_id.clone();
let inviter_id = invite.inviter_id;
let nickname = invite.nickname.clone();
let granted = invite.services.clone();
let paired_at = Some(epoch_now().to_string());
let last_addr = Some(invite.inviter_addr_json.clone());
tokio::task::spawn_blocking(move || {
let existing = store.resolve(&inviter_id)?;
let mut services = existing
.as_ref()
.map(|e| e.services.clone())
.unwrap_or_default();
for svc in granted {
if !services.contains(&svc) {
services.push(svc);
}
}
store.add(PeerEntry {
endpoint_id: inviter_id,
nickname,
services,
paired_at,
user_id: inviter_user_id.or_else(|| existing.and_then(|e| e.user_id)),
last_addr,
})
})
.await
.context("join redeemer store write")??;
if let Some(grant_back) = grant_back {
let inviter_principal = peer_user_id
.clone()
.unwrap_or_else(|| mcpmesh_net::EndpointId::from_bytes(invite.inviter_id).principal());
grant_back(inviter_principal, invite.nickname.clone()).await?;
}
let self_id = *endpoint.id().as_bytes();
let sas_code = short_auth_code(&invite.inviter_id, &self_id, &invite.secret);
Ok(PairResult {
peer_nickname: invite.nickname,
sas_code,
services: invite.services,
app_label: invite.app_label,
peer_user_id,
})
}
fn refusal_error(reason: &str, code: Option<RefusalCode>) -> anyhow::Error {
let msg = format!("pairing refused: {reason}");
match code {
Some(RefusalCode::NicknameTaken) => anyhow::Error::new(NicknameTaken(msg)),
_ => anyhow::anyhow!(msg),
}
}
fn nickname_collision(
store: &PeerStore,
nickname: &str,
tls_id: &[u8; 32],
) -> anyhow::Result<bool> {
Ok(store
.list()?
.into_iter()
.any(|e| e.nickname == nickname && &e.endpoint_id != tls_id))
}
fn nickname_squat(
store: &PeerStore,
nickname: &str,
inviter_id: &[u8; 32],
) -> anyhow::Result<Option<String>> {
let clashes = store
.list()?
.into_iter()
.any(|e| e.nickname == nickname && &e.endpoint_id != inviter_id);
Ok(clashes.then(|| {
"you already use that name for a different peer — \
accepting it would make your own dials to that name ambiguous. \
Unpair the existing peer first if you no longer need it."
.to_string()
}))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn the_refusal_names_an_action_not_a_control_verb() {
let survived = reason_nickname_taken("studio-mac", true);
assert!(
!survived.contains("set_nickname"),
"no control verb may appear in a string a human is shown: {survived}"
);
assert!(survived.contains("rename this node"), "got {survived}");
assert!(
survived.contains("the invite was NOT consumed"),
"the recoverable case must still say the invite survived — that is what makes the \
advice actionable (#87): {survived}"
);
assert!(survived.contains("studio-mac"), "got {survived}");
let burned = reason_nickname_taken("studio-mac", false);
assert!(
burned.contains("ask the inviter for a fresh invite"),
"the burned-invite clause names an action already; it must not regress: {burned}"
);
assert!(!burned.contains("set_nickname"), "got {burned}");
}
#[test]
fn a_refusal_code_is_additive_and_degrades() {
let old: PairReply = serde_json::from_value(
serde_json::json!({"result": "refused", "reason": "pairing refused"}),
)
.expect("an older inviter's refusal must still parse");
let PairReply::Refused { code, reason } = old else {
panic!("expected a refusal");
};
assert_eq!(code, None, "absent means absent — never a guessed kind");
assert_eq!(reason, "pairing refused");
for bad in [
serde_json::json!("invite_expired"),
serde_json::Value::Null,
serde_json::json!(7),
serde_json::json!(true),
serde_json::json!({"kind": "nickname_taken", "nested": [1, 2]}),
serde_json::json!(["nickname_taken"]),
] {
let v = serde_json::json!({"result": "refused", "reason": "r", "code": bad});
let reply: PairReply = serde_json::from_value(v)
.unwrap_or_else(|e| panic!("`code: {bad}` must not fail the whole reply: {e}"));
let PairReply::Refused { code, reason } = reply else {
panic!("expected a refusal");
};
assert_eq!(reason, "r", "the rest of the reply survives: {bad}");
assert!(
matches!(code, Some(RefusalCode::Unknown) | None),
"an unreadable code must degrade, not claim a kind: {bad} -> {code:?}"
);
}
}
#[test]
fn only_the_collision_refusal_is_coded() {
let coded = serde_json::to_value(PairReply::Refused {
reason: reason_nickname_taken("bob", true),
code: Some(RefusalCode::NicknameTaken),
})
.unwrap();
assert_eq!(coded["code"], "nickname_taken", "got {coded}");
let generic = serde_json::to_value(PairReply::Refused {
reason: REASON_REFUSED.into(),
code: None,
})
.unwrap();
assert!(
generic.get("code").is_none(),
"the opaque refusal must carry NO code — one would make it a redemption oracle: \
{generic}"
);
assert_eq!(
generic["reason"], REASON_REFUSED,
"and its reason stays opaque: {generic}"
);
}
#[test]
fn only_a_coded_collision_refusal_becomes_the_typed_error() {
let wire = reason_nickname_taken("studio-mac", true);
let coded = refusal_error(&wire, Some(RefusalCode::NicknameTaken));
assert!(
coded.downcast_ref::<NicknameTaken>().is_some(),
"respond's downcast arm is what maps this to ERR_NICKNAME_TAKEN: {coded}"
);
assert!(coded.to_string().contains("rename this node"), "{coded}");
for opaque in [None, Some(RefusalCode::Unknown)] {
let e = refusal_error(REASON_REFUSED, opaque);
assert!(
e.downcast_ref::<NicknameTaken>().is_none(),
"an opaque refusal must NOT claim the collision code — an embedder would tell a user to rename after a wrong or expired secret: {opaque:?} -> {e}"
);
assert_eq!(e.to_string(), format!("pairing refused: {REASON_REFUSED}"));
}
}
#[test]
fn only_a_surviving_invite_earns_the_rename_and_retry_code() {
let PairReply::Refused { reason, code } = collision_refusal("studio-mac", true) else {
panic!("expected a refusal");
};
assert!(reason.contains("redeem the same invite again"), "{reason}");
assert_eq!(
code,
Some(RefusalCode::NicknameTaken),
"the recoverable collision is the one that earns the code"
);
let PairReply::Refused { reason, code } = collision_refusal("studio-mac", false) else {
panic!("expected a refusal");
};
assert!(
reason.contains("ask the inviter for a fresh invite"),
"{reason}"
);
assert!(
!reason.contains("redeem the same invite again"),
"the two remedies must stay distinguishable: {reason}"
);
assert_eq!(
code, None,
"a burned invite must NOT carry the rename-and-retry code — an embedder writing copy \
off it would send the user back to an invite that no longer exists"
);
}
#[test]
fn the_typed_error_displays_the_inviters_reason_verbatim() {
let wire = reason_nickname_taken("studio-mac", true);
let e = NicknameTaken(format!("pairing refused: {wire}"));
assert_eq!(e.to_string(), format!("pairing refused: {wire}"));
assert!(e.to_string().contains("rename this node"));
}
}