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                invoice_description: b.invoice_description,
205            }),
206            onchain: settings.onchain.map(|o| super::OnchainSettings {
207                confirmations: o.confirmations,
208                min_receive_amount_sat: o.min_receive_amount_sat,
209                min_send_amount_sat: o.min_send_amount_sat,
210            }),
211            custom: settings.custom,
212        }))
213    }
214
215    async fn create_payment(
216        &self,
217        request: Request<CreatePaymentRequest>,
218    ) -> Result<Response<CreatePaymentResponse>, Status> {
219        let CreatePaymentRequest { options, .. } = request.into_inner();
220
221        let options = options.ok_or_else(|| Status::invalid_argument("Missing payment options"))?;
222
223        let proto_options = match options
224            .options
225            .ok_or_else(|| Status::invalid_argument("Missing options"))?
226        {
227            incoming_payment_options::Options::Custom(opts) => {
228                let amount: Option<cdk_common::Amount<CurrencyUnit>> = match opts.amount {
229                    Some(a) => Some(
230                        a.try_into()
231                            .map_err(|_| Status::invalid_argument("Invalid amount"))?,
232                    ),
233                    None => None,
234                };
235                if opts.quote_id.is_empty() {
236                    return Err(Status::invalid_argument(
237                        "Missing quote_id in Custom options",
238                    ));
239                }
240                let quote_id = parse_quote_id(&opts.quote_id)?;
241                let pubkey = opts
242                    .pubkey
243                    .as_deref()
244                    .map(PublicKey::from_str)
245                    .transpose()
246                    .map_err(|_| Status::invalid_argument("Invalid pubkey in Custom options"))?;
247                IncomingPaymentOptions::Custom(Box::new(
248                    cdk_common::payment::CustomIncomingPaymentOptions {
249                        method: "".to_string(),
250                        description: opts.description,
251                        amount,
252                        unix_expiry: opts.unix_expiry,
253                        extra_json: opts.extra_json,
254                        quote_id,
255                        pubkey,
256                    },
257                ))
258            }
259            incoming_payment_options::Options::Bolt11(opts) => {
260                let amount = opts
261                    .amount
262                    .ok_or_else(|| Status::invalid_argument("Missing amount"))?
263                    .try_into()
264                    .map_err(|_| Status::invalid_argument("Invalid amount"))?;
265                IncomingPaymentOptions::Bolt11(cdk_common::payment::Bolt11IncomingPaymentOptions {
266                    description: opts.description,
267                    amount,
268                    unix_expiry: opts.unix_expiry,
269                })
270            }
271            incoming_payment_options::Options::Bolt12(opts) => {
272                let amount: Option<cdk_common::Amount<CurrencyUnit>> = match opts.amount {
273                    Some(a) => Some(
274                        a.try_into()
275                            .map_err(|_| Status::invalid_argument("Invalid amount"))?,
276                    ),
277                    None => None,
278                };
279                IncomingPaymentOptions::Bolt12(Box::new(
280                    cdk_common::payment::Bolt12IncomingPaymentOptions {
281                        description: opts.description,
282                        amount,
283                        unix_expiry: opts.unix_expiry,
284                    },
285                ))
286            }
287            incoming_payment_options::Options::Onchain(opts) => IncomingPaymentOptions::Onchain(
288                cdk_common::payment::OnchainIncomingPaymentOptions {
289                    quote_id: opts.quote_id.parse().map_err(|_| {
290                        Status::invalid_argument("Invalid quote_id in Onchain options")
291                    })?,
292                },
293            ),
294        };
295
296        let invoice_response = self
297            .inner
298            .create_incoming_payment_request(proto_options)
299            .await
300            .map_err(|_| Status::internal("Could not create invoice"))?;
301
302        Ok(Response::new(invoice_response.into()))
303    }
304
305    async fn get_payment_quote(
306        &self,
307        request: Request<PaymentQuoteRequest>,
308    ) -> Result<Response<PaymentQuoteResponse>, Status> {
309        let request = request.into_inner();
310
311        let unit = CurrencyUnit::from_str(&request.unit)
312            .map_err(|_| Status::invalid_argument("Invalid currency unit"))?;
313
314        let quote_id = parse_quote_id(&request.quote_id)?;
315
316        let options = match request.request_type() {
317            OutgoingPaymentRequestType::Bolt11Invoice => {
318                let bolt11: cdk_common::Bolt11Invoice =
319                    request.request.parse().map_err(Error::Invoice)?;
320
321                cdk_common::payment::OutgoingPaymentOptions::Bolt11(Box::new(
322                    cdk_common::payment::Bolt11OutgoingPaymentOptions {
323                        bolt11,
324                        max_fee_amount: None,
325                        timeout_secs: None,
326                        melt_options: request.options.map(TryInto::try_into).transpose()?,
327                        quote_id,
328                    },
329                ))
330            }
331            OutgoingPaymentRequestType::Bolt12Offer => {
332                // Parse offer to verify it's valid, but store as string
333                let _: Offer = request.request.parse().map_err(|_| Error::Bolt12Parse)?;
334
335                cdk_common::payment::OutgoingPaymentOptions::Bolt12(Box::new(
336                    cdk_common::payment::Bolt12OutgoingPaymentOptions {
337                        offer: Offer::from_str(&request.request)
338                            .expect("Already validated offer above"),
339                        max_fee_amount: None,
340                        timeout_secs: None,
341                        melt_options: request.options.map(TryInto::try_into).transpose()?,
342                        quote_id,
343                    },
344                ))
345            }
346            OutgoingPaymentRequestType::Custom => {
347                let amount = request
348                    .amount
349                    .try_from_proto()
350                    .map_err(|_| Status::invalid_argument("Invalid amount"))?;
351
352                // Custom payment method - pass request as-is with no validation
353                cdk_common::payment::OutgoingPaymentOptions::Custom(Box::new(
354                    cdk_common::payment::CustomOutgoingPaymentOptions {
355                        method: String::new(), // Will be set from variant
356                        request: request.request.clone(),
357                        amount,
358                        max_fee_amount: None,
359                        timeout_secs: None,
360                        melt_options: request.options.map(TryInto::try_into).transpose()?,
361                        extra_json: request.extra_json.clone(),
362                        quote_id,
363                    },
364                ))
365            }
366            OutgoingPaymentRequestType::Onchain => {
367                let opts = request.onchain_options.ok_or_else(|| {
368                    Status::invalid_argument("Missing onchain_options for onchain quote")
369                })?;
370                let amount = opts
371                    .amount
372                    .ok_or_else(|| Status::invalid_argument("Missing amount in onchain quote"))?
373                    .try_into()
374                    .map_err(|_| Status::invalid_argument("Invalid amount"))?;
375                let max_fee_amount = opts
376                    .max_fee_amount
377                    .try_from_proto()
378                    .map_err(|_| Status::invalid_argument("Invalid max_fee_amount"))?;
379                let onchain_quote_id = parse_quote_id(&opts.quote_id)?;
380                if onchain_quote_id != quote_id {
381                    return Err(Status::invalid_argument(
382                        "quote_id does not match onchain_options quote_id",
383                    ));
384                }
385
386                cdk_common::payment::OutgoingPaymentOptions::Onchain(Box::new(
387                    cdk_common::payment::OnchainOutgoingPaymentOptions {
388                        address: opts.address,
389                        amount,
390                        max_fee_amount,
391                        quote_id,
392                        fee_index: opts.fee_index,
393                        metadata: opts.metadata,
394                    },
395                ))
396            }
397            OutgoingPaymentRequestType::Unspecified => {
398                return Err(Status::invalid_argument("Unspecified payment request type"));
399            }
400        };
401
402        let payment_quote = self
403            .inner
404            .get_payment_quote(&unit, options)
405            .await
406            .map_err(|err| {
407                tracing::error!("Could not get payment quote: {}", err);
408                Status::internal("Could not get quote")
409            })?;
410
411        Ok(Response::new(payment_quote.into()))
412    }
413
414    async fn make_payment(
415        &self,
416        request: Request<MakePaymentRequest>,
417    ) -> Result<Response<MakePaymentResponse>, Status> {
418        let request = request.into_inner();
419
420        let unit = CurrencyUnit::from_str(&request.unit)
421            .map_err(|_| Status::invalid_argument("Invalid currency unit"))?;
422
423        let options = request
424            .payment_options
425            .ok_or_else(|| Status::invalid_argument("Missing payment options"))?;
426
427        let payment_options = match options
428            .options
429            .ok_or_else(|| Status::invalid_argument("Missing options"))?
430        {
431            outgoing_payment_variant::Options::Bolt11(opts) => {
432                let bolt11: cdk_common::Bolt11Invoice =
433                    opts.bolt11.parse().map_err(Error::Invoice)?;
434
435                let max_fee_amount = opts
436                    .max_fee_amount
437                    .try_from_proto()
438                    .map_err(|_| Status::invalid_argument("Invalid max_fee_amount"))?;
439                let quote_id = parse_quote_id(&opts.quote_id)?;
440
441                cdk_common::payment::OutgoingPaymentOptions::Bolt11(Box::new(
442                    cdk_common::payment::Bolt11OutgoingPaymentOptions {
443                        bolt11,
444                        max_fee_amount,
445                        timeout_secs: opts.timeout_secs,
446                        melt_options: opts.melt_options.map(TryInto::try_into).transpose()?,
447                        quote_id,
448                    },
449                ))
450            }
451            outgoing_payment_variant::Options::Bolt12(opts) => {
452                let offer = Offer::from_str(&opts.offer).map_err(|_| Error::Bolt12Parse)?;
453
454                let max_fee_amount = opts
455                    .max_fee_amount
456                    .try_from_proto()
457                    .map_err(|_| Status::invalid_argument("Invalid max_fee_amount"))?;
458                let quote_id = parse_quote_id(&opts.quote_id)?;
459
460                cdk_common::payment::OutgoingPaymentOptions::Bolt12(Box::new(
461                    cdk_common::payment::Bolt12OutgoingPaymentOptions {
462                        offer,
463                        max_fee_amount,
464                        timeout_secs: opts.timeout_secs,
465                        melt_options: opts.melt_options.map(TryInto::try_into).transpose()?,
466                        quote_id,
467                    },
468                ))
469            }
470            outgoing_payment_variant::Options::Custom(opts) => {
471                let max_fee_amount = opts
472                    .max_fee_amount
473                    .try_from_proto()
474                    .map_err(|_| Status::invalid_argument("Invalid max_fee_amount"))?;
475                let quote_id = parse_quote_id(&opts.quote_id)?;
476                let amount: Option<cdk_common::Amount<CurrencyUnit>> = match opts.amount {
477                    Some(a) => Some(
478                        a.try_into()
479                            .map_err(|_| Status::invalid_argument("Invalid amount"))?,
480                    ),
481                    None => None,
482                };
483
484                cdk_common::payment::OutgoingPaymentOptions::Custom(Box::new(
485                    cdk_common::payment::CustomOutgoingPaymentOptions {
486                        method: String::new(), // Method will be determined from context
487                        request: opts.offer,   // Reusing offer field for custom request string
488                        amount,
489                        max_fee_amount,
490                        timeout_secs: opts.timeout_secs,
491                        melt_options: opts.melt_options.map(TryInto::try_into).transpose()?,
492                        extra_json: opts.extra_json,
493                        quote_id,
494                    },
495                ))
496            }
497            outgoing_payment_variant::Options::Onchain(opts) => {
498                let amount = opts
499                    .amount
500                    .ok_or_else(|| Status::invalid_argument("Missing amount"))?
501                    .try_into()
502                    .map_err(|_| Status::invalid_argument("Invalid amount"))?;
503
504                let max_fee_amount = opts
505                    .max_fee_amount
506                    .try_from_proto()
507                    .map_err(|_| Status::invalid_argument("Invalid max_fee_amount"))?;
508
509                cdk_common::payment::OutgoingPaymentOptions::Onchain(Box::new(
510                    cdk_common::payment::OnchainOutgoingPaymentOptions {
511                        address: opts.address,
512                        amount,
513                        max_fee_amount,
514                        quote_id: opts.quote_id.parse().map_err(|_| {
515                            Status::invalid_argument("Invalid quote_id in Onchain options")
516                        })?,
517                        fee_index: opts.fee_index,
518                        metadata: opts.metadata,
519                    },
520                ))
521            }
522        };
523
524        let pay_response = self
525            .inner
526            .make_payment(&unit, payment_options)
527            .await
528            .map_err(|err| {
529                tracing::error!("Could not make payment: {}", err);
530
531                match err {
532                    cdk_common::payment::Error::InvoiceAlreadyPaid => {
533                        Status::already_exists("Payment request already paid")
534                    }
535                    cdk_common::payment::Error::InvoicePaymentPending => {
536                        Status::already_exists("Payment request pending")
537                    }
538                    _ => Status::internal("Could not pay invoice"),
539                }
540            })?;
541
542        Ok(Response::new(pay_response.into()))
543    }
544
545    async fn check_incoming_payment(
546        &self,
547        request: Request<CheckIncomingPaymentRequest>,
548    ) -> Result<Response<CheckIncomingPaymentResponse>, Status> {
549        let request = request.into_inner();
550
551        let payment_identifier = request
552            .request_identifier
553            .ok_or_else(|| Status::invalid_argument("Missing request identifier"))?
554            .try_into()
555            .map_err(|_| Status::invalid_argument("Invalid request identifier"))?;
556
557        let check_responses = self
558            .inner
559            .check_incoming_payment_status(&payment_identifier)
560            .await
561            .map_err(|_| Status::internal("Could not check incoming payment status"))?;
562
563        Ok(Response::new(CheckIncomingPaymentResponse {
564            payments: check_responses.into_iter().map(|r| r.into()).collect(),
565        }))
566    }
567
568    async fn check_outgoing_payment(
569        &self,
570        request: Request<CheckOutgoingPaymentRequest>,
571    ) -> Result<Response<MakePaymentResponse>, Status> {
572        let request = request.into_inner();
573
574        let payment_identifier = request
575            .request_identifier
576            .ok_or_else(|| Status::invalid_argument("Missing request identifier"))?
577            .try_into()
578            .map_err(|_| Status::invalid_argument("Invalid request identifier"))?;
579
580        let check_response = self
581            .inner
582            .check_outgoing_payment(&payment_identifier)
583            .await
584            .map_err(|_| Status::internal("Could not check outgoing payment status"))?;
585
586        Ok(Response::new(check_response.into()))
587    }
588
589    type WaitPaymentEventStream = ResponseStream;
590
591    #[allow(clippy::incompatible_msrv)]
592    #[instrument(skip_all)]
593    async fn wait_payment_event(
594        &self,
595        _request: Request<EmptyRequest>,
596    ) -> Result<Response<Self::WaitPaymentEventStream>, Status> {
597        tracing::debug!("Server waiting for payment stream");
598        let (tx, rx) = mpsc::channel(128);
599
600        let shutdown_clone = self.shutdown.clone();
601        let ln = self.inner.clone();
602        tokio::spawn(async move {
603            loop {
604                tokio::select! {
605                    _ = shutdown_clone.notified() => {
606                        tracing::info!("Shutdown signal received, stopping task");
607                        ln.cancel_payment_event_stream();
608                        break;
609                    }
610                    result = ln.wait_payment_event() => {
611                        match result {
612                            Ok(mut stream) => {
613                                while let Some(event) = stream.next().await {
614                                    match tx.send(Result::<_, Status>::Ok(event.into())).await {
615                                        Ok(_) => {
616                                            // Response was queued to be sent to client
617                                        }
618                                        Err(item) => {
619                                            tracing::error!("Error adding payment event to stream: {}", item);
620                                            break;
621                                        }
622                                    }
623                                }
624                            }
625                            Err(err) => {
626                                tracing::warn!("Could not get invoice stream: {}", err);
627                                tokio::time::sleep(std::time::Duration::from_secs(5)).await;
628                            }
629                        }
630                    }
631                }
632            }
633        });
634
635        let output_stream = ReceiverStream::new(rx);
636        Ok(Response::new(
637            Box::pin(output_stream) as Self::WaitPaymentEventStream
638        ))
639    }
640}
641
642fn parse_quote_id(s: &str) -> Result<QuoteId, Status> {
643    s.parse()
644        .map_err(|err| Status::invalid_argument(format!("Invalid quote_id: {err}")))
645}