1use core::fmt;
2use std::collections::BTreeMap;
3use std::sync::Arc;
4
5use fedimint_api_client::api::{FederationApiExt, ServerError};
6use fedimint_api_client::query::FilterMapThreshold;
7use fedimint_client_module::DynGlobalClientContext;
8use fedimint_client_module::sm::{ClientSMDatabaseTransaction, State, StateTransition};
9use fedimint_client_module::transaction::{ClientInput, ClientInputBundle};
10use fedimint_core::core::OperationId;
11use fedimint_core::encoding::{Decodable, Encodable};
12use fedimint_core::module::{Amounts, ApiRequestErased};
13use fedimint_core::secp256k1::Keypair;
14use fedimint_core::util::FmtCompact;
15use fedimint_core::{NumPeersExt, OutPoint, PeerId, runtime};
16use fedimint_lnv2_common::contracts::IncomingContract;
17use fedimint_lnv2_common::endpoint_constants::DECRYPTION_KEY_SHARE_ENDPOINT;
18use fedimint_lnv2_common::{LightningInput, LightningInputV0};
19use fedimint_logging::LOG_CLIENT_MODULE_GW;
20use serde::{Deserialize, Serialize};
21use tpe::{
22 AggregateDecryptionKey, AggregatePublicKey, DecryptionKeyShare, PublicKeyShare,
23 aggregate_dk_shares,
24};
25use tracing::warn;
26
27use super::events::{IncomingPaymentFailed, IncomingPaymentSucceeded};
28use crate::GatewayClientContextV2;
29
30#[derive(Debug, Clone, Eq, PartialEq, Hash, Decodable, Encodable)]
31pub struct ReceiveStateMachine {
32 pub common: ReceiveSMCommon,
33 pub state: ReceiveSMState,
34}
35
36impl ReceiveStateMachine {
37 pub fn update(&self, state: ReceiveSMState) -> Self {
38 Self {
39 common: self.common.clone(),
40 state,
41 }
42 }
43}
44
45impl fmt::Display for ReceiveStateMachine {
46 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
47 write!(
48 f,
49 "Receive State Machine Operation ID: {:?} State: {}",
50 self.common.operation_id, self.state
51 )
52 }
53}
54
55#[derive(Debug, Clone, Eq, PartialEq, Hash, Decodable, Encodable)]
56pub struct ReceiveSMCommon {
57 pub operation_id: OperationId,
58 pub contract: IncomingContract,
59 pub outpoint: OutPoint,
60 pub refund_keypair: Keypair,
61}
62
63#[derive(Debug, Clone, Eq, PartialEq, Hash, Decodable, Encodable)]
64pub enum ReceiveSMState {
65 Funding,
66 Rejected(String),
67 Success([u8; 32]),
68 Failure,
69 Refunding(Vec<OutPoint>),
70}
71
72impl fmt::Display for ReceiveSMState {
73 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
74 match self {
75 ReceiveSMState::Funding => write!(f, "Funding"),
76 ReceiveSMState::Rejected(_) => write!(f, "Rejected"),
77 ReceiveSMState::Success(_) => write!(f, "Success"),
78 ReceiveSMState::Failure => write!(f, "Failure"),
79 ReceiveSMState::Refunding(_) => write!(f, "Refunding"),
80 }
81 }
82}
83
84#[derive(Debug, Serialize, Deserialize)]
89pub enum DecryptionOutcome {
90 Rejected(String),
92 InconsistentKeys,
95 Decrypted(AggregateDecryptionKey, Option<[u8; 32]>),
98}
99
100#[cfg_attr(doc, aquamarine::aquamarine)]
101impl State for ReceiveStateMachine {
113 type ModuleContext = GatewayClientContextV2;
114
115 fn transitions(
116 &self,
117 context: &Self::ModuleContext,
118 global_context: &DynGlobalClientContext,
119 ) -> Vec<StateTransition<Self>> {
120 let gc = global_context.clone();
121 let tpe_agg_pk = context.tpe_agg_pk;
122 let gateway_context_ready = context.clone();
123
124 match &self.state {
125 ReceiveSMState::Funding => {
126 vec![StateTransition::new(
127 Self::await_decryption_shares(
128 global_context.clone(),
129 context.tpe_pks.clone(),
130 tpe_agg_pk,
131 self.common.outpoint,
132 self.common.contract.clone(),
133 ),
134 move |dbtx, outcome, old_state| {
135 Box::pin(Self::transition_decryption_shares(
136 dbtx,
137 outcome,
138 old_state,
139 gc.clone(),
140 gateway_context_ready.clone(),
141 ))
142 },
143 )]
144 }
145 ReceiveSMState::Success(..)
146 | ReceiveSMState::Rejected(..)
147 | ReceiveSMState::Refunding(..)
148 | ReceiveSMState::Failure => {
149 vec![]
150 }
151 }
152 }
153
154 fn operation_id(&self) -> OperationId {
155 self.common.operation_id
156 }
157}
158
159impl ReceiveStateMachine {
160 async fn await_decryption_shares(
161 global_context: DynGlobalClientContext,
162 tpe_pks: BTreeMap<PeerId, PublicKeyShare>,
163 tpe_agg_pk: AggregatePublicKey,
164 outpoint: OutPoint,
165 contract: IncomingContract,
166 ) -> DecryptionOutcome {
167 let num_peers = global_context.api().all_peers().to_num_peers();
168 let module_api = global_context.module_api();
169 let tpe_pks = Arc::new(tpe_pks);
170 let contract = Arc::new(contract);
171
172 let decryption_shares = module_api.request_with_strategy_retry(
178 FilterMapThreshold::new(
179 {
180 let tpe_pks = tpe_pks.clone();
181 let contract = contract.clone();
182
183 move |peer_id, share: DecryptionKeyShare| {
184 let tpe_pks = tpe_pks.clone();
185 let contract = contract.clone();
186
187 runtime::spawn_blocking(move || {
189 let pk =
190 tpe_pks
191 .get(&peer_id)
192 .ok_or(ServerError::InternalClientError(format!(
193 "Missing TPE PK for peer {peer_id}?!"
194 )))?;
195
196 if !contract.verify_decryption_share(pk, &share) {
197 return Err(ServerError::InvalidResponse(
198 "Invalid decryption share".to_string(),
199 ));
200 }
201
202 Ok(share)
203 })
204 }
205 },
206 num_peers,
207 ),
208 DECRYPTION_KEY_SHARE_ENDPOINT.to_owned(),
209 ApiRequestErased::new(outpoint),
210 );
211
212 let decryption_shares = std::pin::pin!(decryption_shares);
213 let tx_accepted = std::pin::pin!(global_context.await_tx_accepted(outpoint.txid));
214
215 let decryption_shares = match futures::future::select(decryption_shares, tx_accepted).await
216 {
217 futures::future::Either::Left((shares, _)) => shares,
218 futures::future::Either::Right((accepted, decryption_shares)) => {
219 if let Err(error) = accepted {
220 return DecryptionOutcome::Rejected(error);
221 }
222
223 decryption_shares.await
224 }
225 };
226
227 runtime::spawn_blocking(move || {
229 let agg_decryption_key = aggregate_dk_shares(
230 &decryption_shares
231 .into_iter()
232 .map(|(peer, share)| (peer.to_usize() as u64, share))
233 .collect(),
234 );
235
236 if !contract.verify_agg_decryption_key(&tpe_agg_pk, &agg_decryption_key) {
237 return DecryptionOutcome::InconsistentKeys;
238 }
239
240 let preimage = contract.decrypt_preimage(&agg_decryption_key);
241
242 DecryptionOutcome::Decrypted(agg_decryption_key, preimage)
243 })
244 .await
245 }
246
247 async fn transition_decryption_shares(
248 dbtx: &mut ClientSMDatabaseTransaction<'_, '_>,
249 outcome: DecryptionOutcome,
250 old_state: ReceiveStateMachine,
251 global_context: DynGlobalClientContext,
252 client_ctx: GatewayClientContextV2,
253 ) -> ReceiveStateMachine {
254 let payment_image = old_state.common.contract.commitment.payment_image.clone();
255
256 let (agg_decryption_key, preimage) = match outcome {
257 DecryptionOutcome::Rejected(error) => {
258 client_ctx
259 .module
260 .client_ctx
261 .log_event(
262 &mut dbtx.module_tx(),
263 IncomingPaymentFailed {
264 payment_image,
265 error: error.clone(),
266 },
267 )
268 .await;
269
270 return old_state.update(ReceiveSMState::Rejected(error));
271 }
272 DecryptionOutcome::InconsistentKeys => {
273 warn!(target: LOG_CLIENT_MODULE_GW, "Failed to obtain decryption key. Client config's public keys are inconsistent");
274
275 client_ctx
276 .module
277 .client_ctx
278 .log_event(
279 &mut dbtx.module_tx(),
280 IncomingPaymentFailed {
281 payment_image,
282 error: "Client config's public keys are inconsistent".to_string(),
283 },
284 )
285 .await;
286
287 return old_state.update(ReceiveSMState::Failure);
288 }
289 DecryptionOutcome::Decrypted(agg_decryption_key, preimage) => {
290 (agg_decryption_key, preimage)
291 }
292 };
293
294 if let Some(preimage) = preimage {
295 client_ctx
296 .module
297 .client_ctx
298 .log_event(
299 &mut dbtx.module_tx(),
300 IncomingPaymentSucceeded { payment_image },
301 )
302 .await;
303
304 return old_state.update(ReceiveSMState::Success(preimage));
305 }
306
307 let client_input = ClientInput::<LightningInput> {
308 input: LightningInput::V0(LightningInputV0::Incoming(
309 old_state.common.outpoint,
310 agg_decryption_key,
311 )),
312 amounts: Amounts::new_bitcoin(old_state.common.contract.commitment.amount),
313 keys: vec![old_state.common.refund_keypair],
314 };
315
316 let outpoints = match global_context
317 .claim_inputs(
318 dbtx,
319 ClientInputBundle::new_no_sm(vec![client_input]),
321 )
322 .await
323 {
324 Ok(outpoints) => outpoints.into_iter().collect(),
325 Err(err) => {
332 warn!(
333 target: LOG_CLIENT_MODULE_GW,
334 err = %err.fmt_compact(),
335 amount = %old_state.common.contract.commitment.amount,
336 "Not refunding incoming contract, its amount does not cover the refund fee"
337 );
338
339 client_ctx
340 .module
341 .client_ctx
342 .log_event(
343 &mut dbtx.module_tx(),
344 IncomingPaymentFailed {
345 payment_image: old_state
346 .common
347 .contract
348 .commitment
349 .payment_image
350 .clone(),
351 error: "Contract does not cover the refund fee".to_string(),
352 },
353 )
354 .await;
355
356 return old_state.update(ReceiveSMState::Failure);
357 }
358 };
359
360 client_ctx
361 .module
362 .client_ctx
363 .log_event(
364 &mut dbtx.module_tx(),
365 IncomingPaymentFailed {
366 payment_image: old_state.common.contract.commitment.payment_image.clone(),
367 error: "Failed to decrypt preimage".to_string(),
368 },
369 )
370 .await;
371
372 old_state.update(ReceiveSMState::Refunding(outpoints))
373 }
374}