Skip to main content

cdk_payment_processor/proto/
server.rs

1use std::net::SocketAddr;
2use std::path::PathBuf;
3use std::pin::Pin;
4use std::str::FromStr;
5use std::sync::Arc;
6use std::time::Duration;
7
8use cdk_common::grpc::create_version_check_interceptor;
9use cdk_common::payment::{IncomingPaymentOptions, MintPayment};
10use cdk_common::{CurrencyUnit, PublicKey, QuoteId};
11use futures::{Stream, StreamExt};
12use lightning::offers::offer::Offer;
13use tokio::sync::{mpsc, Notify};
14use tokio::task::JoinHandle;
15use tokio::time::{sleep, Instant};
16use tokio_stream::wrappers::ReceiverStream;
17use tonic::transport::{Certificate, Identity, Server, ServerTlsConfig};
18use tonic::{async_trait, Request, Response, Status};
19use tracing::instrument;
20
21use super::cdk_payment_processor_server::{CdkPaymentProcessor, CdkPaymentProcessorServer};
22use crate::error::Error;
23use crate::proto::{TryFromProtoAmount, *};
24
25type ResponseStream = Pin<Box<dyn Stream<Item = Result<PaymentEventResponse, Status>> + Send>>;
26
27/// Payment Processor
28#[derive(Clone)]
29pub struct PaymentProcessorServer {
30    inner: Arc<dyn MintPayment<Err = cdk_common::payment::Error> + Send + Sync>,
31    socket_addr: SocketAddr,
32    shutdown: Arc<Notify>,
33    handle: Option<Arc<JoinHandle<anyhow::Result<()>>>>,
34}
35
36impl std::fmt::Debug for PaymentProcessorServer {
37    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
38        f.debug_struct("PaymentProcessorServer")
39            .field("socket_addr", &self.socket_addr)
40            .finish_non_exhaustive()
41    }
42}
43
44impl PaymentProcessorServer {
45    /// Create new [`PaymentProcessorServer`]
46    pub fn new(
47        payment_processor: Arc<dyn MintPayment<Err = cdk_common::payment::Error> + Send + Sync>,
48        addr: &str,
49        port: u16,
50    ) -> anyhow::Result<Self> {
51        let socket_addr = SocketAddr::new(addr.parse()?, port);
52        Ok(Self {
53            inner: payment_processor,
54            socket_addr,
55            shutdown: Arc::new(Notify::new()),
56            handle: None,
57        })
58    }
59
60    /// Start fake wallet grpc server
61    pub async fn start(&mut self, tls_dir: Option<PathBuf>) -> anyhow::Result<()> {
62        tracing::info!("Starting RPC server {}", self.socket_addr);
63
64        let server = match tls_dir {
65            Some(tls_dir) => {
66                tracing::info!("TLS configuration found, starting secure server");
67
68                // Check for server.pem
69                let server_pem_path = tls_dir.join("server.pem");
70                if !server_pem_path.exists() {
71                    let err_msg = format!(
72                        "TLS certificate file not found: {}",
73                        server_pem_path.display()
74                    );
75                    tracing::error!("{}", err_msg);
76                    return Err(anyhow::anyhow!(err_msg));
77                }
78
79                // Check for server.key
80                let server_key_path = tls_dir.join("server.key");
81                if !server_key_path.exists() {
82                    let err_msg = format!("TLS key file not found: {}", server_key_path.display());
83                    tracing::error!("{}", err_msg);
84                    return Err(anyhow::anyhow!(err_msg));
85                }
86
87                // Check for ca.pem
88                let ca_pem_path = tls_dir.join("ca.pem");
89                if !ca_pem_path.exists() {
90                    let err_msg =
91                        format!("CA certificate file not found: {}", ca_pem_path.display());
92                    tracing::error!("{}", err_msg);
93                    return Err(anyhow::anyhow!(err_msg));
94                }
95
96                let cert = std::fs::read_to_string(&server_pem_path)?;
97                let key = std::fs::read_to_string(&server_key_path)?;
98                let client_ca_cert = std::fs::read_to_string(&ca_pem_path)?;
99
100                let client_ca_cert = Certificate::from_pem(client_ca_cert);
101                let server_identity = Identity::from_pem(cert, key);
102                let tls_config = ServerTlsConfig::new()
103                    .identity(server_identity)
104                    .client_ca_root(client_ca_cert);
105
106                Server::builder().tls_config(tls_config)?.add_service(
107                    CdkPaymentProcessorServer::with_interceptor(
108                        self.clone(),
109                        create_version_check_interceptor(
110                            cdk_common::grpc::VERSION_HEADER,
111                            cdk_common::PAYMENT_PROCESSOR_PROTOCOL_VERSION,
112                        ),
113                    ),
114                )
115            }
116            None => {
117                tracing::warn!("No valid TLS configuration found, starting insecure server");
118                Server::builder().add_service(CdkPaymentProcessorServer::with_interceptor(
119                    self.clone(),
120                    create_version_check_interceptor(
121                        cdk_common::grpc::VERSION_HEADER,
122                        cdk_common::PAYMENT_PROCESSOR_PROTOCOL_VERSION,
123                    ),
124                ))
125            }
126        };
127
128        let shutdown = self.shutdown.clone();
129        let addr = self.socket_addr;
130
131        self.handle = Some(Arc::new(tokio::spawn(async move {
132            let server = server.serve_with_shutdown(addr, async {
133                shutdown.notified().await;
134            });
135
136            server.await?;
137            Ok(())
138        })));
139
140        Ok(())
141    }
142
143    /// Stop fake wallet grpc server
144    pub async fn stop(&self) -> anyhow::Result<()> {
145        const SHUTDOWN_TIMEOUT: Duration = Duration::from_secs(5);
146
147        if let Some(handle) = &self.handle {
148            tracing::info!("Initiating server shutdown");
149            self.shutdown.notify_waiters();
150
151            let start = Instant::now();
152
153            while !handle.is_finished() {
154                if start.elapsed() >= SHUTDOWN_TIMEOUT {
155                    tracing::error!(
156                        "Server shutdown timed out after {} seconds, aborting handle",
157                        SHUTDOWN_TIMEOUT.as_secs()
158                    );
159                    handle.abort();
160                    break;
161                }
162                sleep(Duration::from_millis(100)).await;
163            }
164
165            if handle.is_finished() {
166                tracing::info!("Server shutdown completed successfully");
167            }
168        } else {
169            tracing::info!("No server handle found, nothing to stop");
170        }
171
172        Ok(())
173    }
174}
175
176impl Drop for PaymentProcessorServer {
177    fn drop(&mut self) {
178        tracing::debug!("Dropping payment process server");
179        self.shutdown.notify_one();
180    }
181}
182
183#[async_trait]
184impl CdkPaymentProcessor for PaymentProcessorServer {
185    async fn get_settings(
186        &self,
187        _request: Request<EmptyRequest>,
188    ) -> Result<Response<SettingsResponse>, Status> {
189        let settings = self
190            .inner
191            .get_settings()
192            .await
193            .map_err(|_| Status::internal("Could not get settings"))?;
194
195        Ok(Response::new(SettingsResponse {
196            unit: settings.unit,
197            bolt11: settings.bolt11.map(|b| super::Bolt11Settings {
198                mpp: b.mpp,
199                amountless: b.amountless,
200                invoice_description: b.invoice_description,
201            }),
202            bolt12: settings.bolt12.map(|b| super::Bolt12Settings {
203                amountless: b.amountless,
204            }),
205            onchain: settings.onchain.map(|o| super::OnchainSettings {
206                confirmations: o.confirmations,
207                min_receive_amount_sat: o.min_receive_amount_sat,
208                min_send_amount_sat: o.min_send_amount_sat,
209            }),
210            custom: settings.custom,
211        }))
212    }
213
214    async fn create_payment(
215        &self,
216        request: Request<CreatePaymentRequest>,
217    ) -> Result<Response<CreatePaymentResponse>, Status> {
218        let CreatePaymentRequest { options, .. } = request.into_inner();
219
220        let options = options.ok_or_else(|| Status::invalid_argument("Missing payment options"))?;
221
222        let proto_options = match options
223            .options
224            .ok_or_else(|| Status::invalid_argument("Missing options"))?
225        {
226            incoming_payment_options::Options::Custom(opts) => {
227                let amount: Option<cdk_common::Amount<CurrencyUnit>> = match opts.amount {
228                    Some(a) => Some(
229                        a.try_into()
230                            .map_err(|_| Status::invalid_argument("Invalid amount"))?,
231                    ),
232                    None => None,
233                };
234                if opts.quote_id.is_empty() {
235                    return Err(Status::invalid_argument(
236                        "Missing quote_id in Custom options",
237                    ));
238                }
239                let quote_id = parse_quote_id(&opts.quote_id)?;
240                let pubkey = opts
241                    .pubkey
242                    .as_deref()
243                    .map(PublicKey::from_str)
244                    .transpose()
245                    .map_err(|_| Status::invalid_argument("Invalid pubkey in Custom options"))?;
246                IncomingPaymentOptions::Custom(Box::new(
247                    cdk_common::payment::CustomIncomingPaymentOptions {
248                        method: "".to_string(),
249                        description: opts.description,
250                        amount,
251                        unix_expiry: opts.unix_expiry,
252                        extra_json: opts.extra_json,
253                        quote_id,
254                        pubkey,
255                    },
256                ))
257            }
258            incoming_payment_options::Options::Bolt11(opts) => {
259                let amount = opts
260                    .amount
261                    .ok_or_else(|| Status::invalid_argument("Missing amount"))?
262                    .try_into()
263                    .map_err(|_| Status::invalid_argument("Invalid amount"))?;
264                IncomingPaymentOptions::Bolt11(cdk_common::payment::Bolt11IncomingPaymentOptions {
265                    description: opts.description,
266                    amount,
267                    unix_expiry: opts.unix_expiry,
268                })
269            }
270            incoming_payment_options::Options::Bolt12(opts) => {
271                let amount: Option<cdk_common::Amount<CurrencyUnit>> = match opts.amount {
272                    Some(a) => Some(
273                        a.try_into()
274                            .map_err(|_| Status::invalid_argument("Invalid amount"))?,
275                    ),
276                    None => None,
277                };
278                IncomingPaymentOptions::Bolt12(Box::new(
279                    cdk_common::payment::Bolt12IncomingPaymentOptions {
280                        description: opts.description,
281                        amount,
282                        unix_expiry: opts.unix_expiry,
283                    },
284                ))
285            }
286            incoming_payment_options::Options::Onchain(opts) => IncomingPaymentOptions::Onchain(
287                cdk_common::payment::OnchainIncomingPaymentOptions {
288                    quote_id: opts.quote_id.parse().map_err(|_| {
289                        Status::invalid_argument("Invalid quote_id in Onchain options")
290                    })?,
291                },
292            ),
293        };
294
295        let invoice_response = self
296            .inner
297            .create_incoming_payment_request(proto_options)
298            .await
299            .map_err(|_| Status::internal("Could not create invoice"))?;
300
301        Ok(Response::new(invoice_response.into()))
302    }
303
304    async fn get_payment_quote(
305        &self,
306        request: Request<PaymentQuoteRequest>,
307    ) -> Result<Response<PaymentQuoteResponse>, Status> {
308        let request = request.into_inner();
309
310        let unit = CurrencyUnit::from_str(&request.unit)
311            .map_err(|_| Status::invalid_argument("Invalid currency unit"))?;
312
313        let quote_id = parse_quote_id(&request.quote_id)?;
314
315        let options = match request.request_type() {
316            OutgoingPaymentRequestType::Bolt11Invoice => {
317                let bolt11: cdk_common::Bolt11Invoice =
318                    request.request.parse().map_err(Error::Invoice)?;
319
320                cdk_common::payment::OutgoingPaymentOptions::Bolt11(Box::new(
321                    cdk_common::payment::Bolt11OutgoingPaymentOptions {
322                        bolt11,
323                        max_fee_amount: None,
324                        timeout_secs: None,
325                        melt_options: request.options.map(TryInto::try_into).transpose()?,
326                        quote_id,
327                    },
328                ))
329            }
330            OutgoingPaymentRequestType::Bolt12Offer => {
331                // Parse offer to verify it's valid, but store as string
332                let _: Offer = request.request.parse().map_err(|_| Error::Bolt12Parse)?;
333
334                cdk_common::payment::OutgoingPaymentOptions::Bolt12(Box::new(
335                    cdk_common::payment::Bolt12OutgoingPaymentOptions {
336                        offer: Offer::from_str(&request.request)
337                            .expect("Already validated offer above"),
338                        max_fee_amount: None,
339                        timeout_secs: None,
340                        melt_options: request.options.map(TryInto::try_into).transpose()?,
341                        quote_id,
342                    },
343                ))
344            }
345            OutgoingPaymentRequestType::Custom => {
346                let amount = request
347                    .amount
348                    .try_from_proto()
349                    .map_err(|_| Status::invalid_argument("Invalid amount"))?;
350
351                // Custom payment method - pass request as-is with no validation
352                cdk_common::payment::OutgoingPaymentOptions::Custom(Box::new(
353                    cdk_common::payment::CustomOutgoingPaymentOptions {
354                        method: String::new(), // Will be set from variant
355                        request: request.request.clone(),
356                        amount,
357                        max_fee_amount: None,
358                        timeout_secs: None,
359                        melt_options: request.options.map(TryInto::try_into).transpose()?,
360                        extra_json: request.extra_json.clone(),
361                        quote_id,
362                    },
363                ))
364            }
365            OutgoingPaymentRequestType::Onchain => {
366                let opts = request.onchain_options.ok_or_else(|| {
367                    Status::invalid_argument("Missing onchain_options for onchain quote")
368                })?;
369                let amount = opts
370                    .amount
371                    .ok_or_else(|| Status::invalid_argument("Missing amount in onchain quote"))?
372                    .try_into()
373                    .map_err(|_| Status::invalid_argument("Invalid amount"))?;
374                let max_fee_amount = opts
375                    .max_fee_amount
376                    .try_from_proto()
377                    .map_err(|_| Status::invalid_argument("Invalid max_fee_amount"))?;
378                let onchain_quote_id = parse_quote_id(&opts.quote_id)?;
379                if onchain_quote_id != quote_id {
380                    return Err(Status::invalid_argument(
381                        "quote_id does not match onchain_options quote_id",
382                    ));
383                }
384
385                cdk_common::payment::OutgoingPaymentOptions::Onchain(Box::new(
386                    cdk_common::payment::OnchainOutgoingPaymentOptions {
387                        address: opts.address,
388                        amount,
389                        max_fee_amount,
390                        quote_id,
391                        fee_index: opts.fee_index,
392                        metadata: opts.metadata,
393                    },
394                ))
395            }
396            OutgoingPaymentRequestType::Unspecified => {
397                return Err(Status::invalid_argument("Unspecified payment request type"));
398            }
399        };
400
401        let payment_quote = self
402            .inner
403            .get_payment_quote(&unit, options)
404            .await
405            .map_err(|err| {
406                tracing::error!("Could not get payment quote: {}", err);
407                Status::internal("Could not get quote")
408            })?;
409
410        Ok(Response::new(payment_quote.into()))
411    }
412
413    async fn make_payment(
414        &self,
415        request: Request<MakePaymentRequest>,
416    ) -> Result<Response<MakePaymentResponse>, Status> {
417        let request = request.into_inner();
418
419        let unit = CurrencyUnit::from_str(&request.unit)
420            .map_err(|_| Status::invalid_argument("Invalid currency unit"))?;
421
422        let options = request
423            .payment_options
424            .ok_or_else(|| Status::invalid_argument("Missing payment options"))?;
425
426        let payment_options = match options
427            .options
428            .ok_or_else(|| Status::invalid_argument("Missing options"))?
429        {
430            outgoing_payment_variant::Options::Bolt11(opts) => {
431                let bolt11: cdk_common::Bolt11Invoice =
432                    opts.bolt11.parse().map_err(Error::Invoice)?;
433
434                let max_fee_amount = opts
435                    .max_fee_amount
436                    .try_from_proto()
437                    .map_err(|_| Status::invalid_argument("Invalid max_fee_amount"))?;
438                let quote_id = parse_quote_id(&opts.quote_id)?;
439
440                cdk_common::payment::OutgoingPaymentOptions::Bolt11(Box::new(
441                    cdk_common::payment::Bolt11OutgoingPaymentOptions {
442                        bolt11,
443                        max_fee_amount,
444                        timeout_secs: opts.timeout_secs,
445                        melt_options: opts.melt_options.map(TryInto::try_into).transpose()?,
446                        quote_id,
447                    },
448                ))
449            }
450            outgoing_payment_variant::Options::Bolt12(opts) => {
451                let offer = Offer::from_str(&opts.offer).map_err(|_| Error::Bolt12Parse)?;
452
453                let max_fee_amount = opts
454                    .max_fee_amount
455                    .try_from_proto()
456                    .map_err(|_| Status::invalid_argument("Invalid max_fee_amount"))?;
457                let quote_id = parse_quote_id(&opts.quote_id)?;
458
459                cdk_common::payment::OutgoingPaymentOptions::Bolt12(Box::new(
460                    cdk_common::payment::Bolt12OutgoingPaymentOptions {
461                        offer,
462                        max_fee_amount,
463                        timeout_secs: opts.timeout_secs,
464                        melt_options: opts.melt_options.map(TryInto::try_into).transpose()?,
465                        quote_id,
466                    },
467                ))
468            }
469            outgoing_payment_variant::Options::Custom(opts) => {
470                let max_fee_amount = opts
471                    .max_fee_amount
472                    .try_from_proto()
473                    .map_err(|_| Status::invalid_argument("Invalid max_fee_amount"))?;
474                let quote_id = parse_quote_id(&opts.quote_id)?;
475                let amount: Option<cdk_common::Amount<CurrencyUnit>> = match opts.amount {
476                    Some(a) => Some(
477                        a.try_into()
478                            .map_err(|_| Status::invalid_argument("Invalid amount"))?,
479                    ),
480                    None => None,
481                };
482
483                cdk_common::payment::OutgoingPaymentOptions::Custom(Box::new(
484                    cdk_common::payment::CustomOutgoingPaymentOptions {
485                        method: String::new(), // Method will be determined from context
486                        request: opts.offer,   // Reusing offer field for custom request string
487                        amount,
488                        max_fee_amount,
489                        timeout_secs: opts.timeout_secs,
490                        melt_options: opts.melt_options.map(TryInto::try_into).transpose()?,
491                        extra_json: opts.extra_json,
492                        quote_id,
493                    },
494                ))
495            }
496            outgoing_payment_variant::Options::Onchain(opts) => {
497                let amount = opts
498                    .amount
499                    .ok_or_else(|| Status::invalid_argument("Missing amount"))?
500                    .try_into()
501                    .map_err(|_| Status::invalid_argument("Invalid amount"))?;
502
503                let max_fee_amount = opts
504                    .max_fee_amount
505                    .try_from_proto()
506                    .map_err(|_| Status::invalid_argument("Invalid max_fee_amount"))?;
507
508                cdk_common::payment::OutgoingPaymentOptions::Onchain(Box::new(
509                    cdk_common::payment::OnchainOutgoingPaymentOptions {
510                        address: opts.address,
511                        amount,
512                        max_fee_amount,
513                        quote_id: opts.quote_id.parse().map_err(|_| {
514                            Status::invalid_argument("Invalid quote_id in Onchain options")
515                        })?,
516                        fee_index: opts.fee_index,
517                        metadata: opts.metadata,
518                    },
519                ))
520            }
521        };
522
523        let pay_response = self
524            .inner
525            .make_payment(&unit, payment_options)
526            .await
527            .map_err(|err| {
528                tracing::error!("Could not make payment: {}", err);
529
530                match err {
531                    cdk_common::payment::Error::InvoiceAlreadyPaid => {
532                        Status::already_exists("Payment request already paid")
533                    }
534                    cdk_common::payment::Error::InvoicePaymentPending => {
535                        Status::already_exists("Payment request pending")
536                    }
537                    _ => Status::internal("Could not pay invoice"),
538                }
539            })?;
540
541        Ok(Response::new(pay_response.into()))
542    }
543
544    async fn check_incoming_payment(
545        &self,
546        request: Request<CheckIncomingPaymentRequest>,
547    ) -> Result<Response<CheckIncomingPaymentResponse>, Status> {
548        let request = request.into_inner();
549
550        let payment_identifier = request
551            .request_identifier
552            .ok_or_else(|| Status::invalid_argument("Missing request identifier"))?
553            .try_into()
554            .map_err(|_| Status::invalid_argument("Invalid request identifier"))?;
555
556        let check_responses = self
557            .inner
558            .check_incoming_payment_status(&payment_identifier)
559            .await
560            .map_err(|_| Status::internal("Could not check incoming payment status"))?;
561
562        Ok(Response::new(CheckIncomingPaymentResponse {
563            payments: check_responses.into_iter().map(|r| r.into()).collect(),
564        }))
565    }
566
567    async fn check_outgoing_payment(
568        &self,
569        request: Request<CheckOutgoingPaymentRequest>,
570    ) -> Result<Response<MakePaymentResponse>, Status> {
571        let request = request.into_inner();
572
573        let payment_identifier = request
574            .request_identifier
575            .ok_or_else(|| Status::invalid_argument("Missing request identifier"))?
576            .try_into()
577            .map_err(|_| Status::invalid_argument("Invalid request identifier"))?;
578
579        let check_response = self
580            .inner
581            .check_outgoing_payment(&payment_identifier)
582            .await
583            .map_err(|_| Status::internal("Could not check outgoing payment status"))?;
584
585        Ok(Response::new(check_response.into()))
586    }
587
588    type WaitPaymentEventStream = ResponseStream;
589
590    #[allow(clippy::incompatible_msrv)]
591    #[instrument(skip_all)]
592    async fn wait_payment_event(
593        &self,
594        _request: Request<EmptyRequest>,
595    ) -> Result<Response<Self::WaitPaymentEventStream>, Status> {
596        tracing::debug!("Server waiting for payment stream");
597        let (tx, rx) = mpsc::channel(128);
598
599        let shutdown_clone = self.shutdown.clone();
600        let ln = self.inner.clone();
601        tokio::spawn(async move {
602            loop {
603                tokio::select! {
604                    _ = shutdown_clone.notified() => {
605                        tracing::info!("Shutdown signal received, stopping task");
606                        ln.cancel_payment_event_stream();
607                        break;
608                    }
609                    result = ln.wait_payment_event() => {
610                        match result {
611                            Ok(mut stream) => {
612                                while let Some(event) = stream.next().await {
613                                    match tx.send(Result::<_, Status>::Ok(event.into())).await {
614                                        Ok(_) => {
615                                            // Response was queued to be sent to client
616                                        }
617                                        Err(item) => {
618                                            tracing::error!("Error adding payment event to stream: {}", item);
619                                            break;
620                                        }
621                                    }
622                                }
623                            }
624                            Err(err) => {
625                                tracing::warn!("Could not get invoice stream: {}", err);
626                                tokio::time::sleep(std::time::Duration::from_secs(5)).await;
627                            }
628                        }
629                    }
630                }
631            }
632        });
633
634        let output_stream = ReceiverStream::new(rx);
635        Ok(Response::new(
636            Box::pin(output_stream) as Self::WaitPaymentEventStream
637        ))
638    }
639}
640
641fn parse_quote_id(s: &str) -> Result<QuoteId, Status> {
642    s.parse()
643        .map_err(|err| Status::invalid_argument(format!("Invalid quote_id: {err}")))
644}