using System; using System.Globalization; using System.Linq; using System.Threading; using System.Threading.Tasks; using BTCPayServer.Client.Models; using BTCPayServer.Data; using BTCPayServer.Services.Invoices; using Dapper; using Microsoft.EntityFrameworkCore; namespace BTCPayServer.Services.PaymentRequests { public record PaymentRequestEvent { public const string Created = nameof(Created); public const string Updated = nameof(Updated); public const string Archived = nameof(Archived); public const string StatusChanged = nameof(StatusChanged); public const string Completed = nameof(Completed); public PaymentRequestData Data { get; set; } public string Type { get; set; } } public class PaymentRequestRepository { private readonly ApplicationDbContextFactory _ContextFactory; private readonly InvoiceRepository _InvoiceRepository; private readonly EventAggregator _eventAggregator; public PaymentRequestRepository(ApplicationDbContextFactory contextFactory, InvoiceRepository invoiceRepository, EventAggregator eventAggregator) { _ContextFactory = contextFactory; _InvoiceRepository = invoiceRepository; _eventAggregator = eventAggregator; } public async Task CreateOrUpdatePaymentRequest(PaymentRequestData entity) { await using var context = _ContextFactory.CreateContext(); var added = false; if (string.IsNullOrEmpty(entity.Id)) { entity.Id = Guid.NewGuid().ToString(); await context.PaymentRequests.AddAsync(entity); added = true; } else { context.PaymentRequests.Update(entity); } await context.SaveChangesAsync(); _eventAggregator.Publish(new PaymentRequestEvent() { Data = entity, Type = added ? PaymentRequestEvent.Created : PaymentRequestEvent.Updated }); return entity; } public async Task ArchivePaymentRequest(string id, bool toggle = false) { await using var context = _ContextFactory.CreateContext(); var pr = await context.PaymentRequests.FindAsync(id); if(pr == null) return null; if(pr.Archived && !toggle) return pr.Archived; pr.Archived = !pr.Archived; await context.SaveChangesAsync(); if (pr.Archived) { _eventAggregator.Publish(new PaymentRequestEvent() { Data = pr, Type = PaymentRequestEvent.Archived }); } return pr.Archived; } public async Task FindPaymentRequest(string id, string userId, CancellationToken cancellationToken = default) { if (string.IsNullOrEmpty(id)) { return null; } await using var context = _ContextFactory.CreateContext(); var result = await context.PaymentRequests.Include(x => x.StoreData) .Where(data => string.IsNullOrEmpty(userId) || (data.StoreData != null && data.StoreData.UserStores.Any(u => u.ApplicationUserId == userId))) .SingleOrDefaultAsync(x => x.Id == id, cancellationToken); return result; } public async Task UpdatePaymentRequestStatus(string paymentRequestId, Client.Models.PaymentRequestStatus status, CancellationToken cancellationToken = default) { await using var context = _ContextFactory.CreateContext(); var conn = context.Database.GetDbConnection(); var affectedRows = await conn.ExecuteAsync(""" UPDATE "PaymentRequests" SET "Status" = @status WHERE "Id" = @id AND "Status" != @status; """, new{ id = paymentRequestId, status = status.ToString()}); if (affectedRows == 0) return; var paymentRequestData = await context.FindAsync(paymentRequestId); if (status != paymentRequestData?.Status) return; _eventAggregator.Publish(new PaymentRequestEvent() { Data = paymentRequestData, Type = PaymentRequestEvent.StatusChanged }); if (status == PaymentRequestStatus.Completed) { _eventAggregator.Publish(new PaymentRequestEvent() { Data = paymentRequestData, Type = PaymentRequestEvent.Completed }); } } public async Task GetExpirablePaymentRequests(CancellationToken cancellationToken = default) { using var context = _ContextFactory.CreateContext(); var queryable = context.PaymentRequests.Include(data => data.StoreData).AsQueryable(); queryable = queryable .Where(data => (data.Status == Client.Models.PaymentRequestStatus.Pending || data.Status == Client.Models.PaymentRequestStatus.Processing) && data.Expiry != null); return await queryable.ToArrayAsync(cancellationToken); } public async Task FindPaymentRequests(PaymentRequestQuery query, CancellationToken cancellationToken = default) { await using var context = _ContextFactory.CreateContext(); IQueryable queryable; if (!string.IsNullOrEmpty(query.SearchText)) { if (string.IsNullOrEmpty(query.StoreId)) throw new InvalidOperationException("PaymentRequestQuery.StoreId should be specified"); var search = query.SearchText; // Escape LIKE wildcards to prevent SQL injection var escapedSearch = search.Replace("\\", "\\\\").Replace("%", "\\%").Replace("_", "\\_"); var likePattern = $"%{escapedSearch}%"; var amountOrNull = decimal.TryParse(search, NumberStyles.Number, CultureInfo.InvariantCulture, out var amount) ? amount : (decimal?)null; queryable = context.PaymentRequests .Where(a => a.StoreDataId == query.StoreId) .Where(a => a.ReferenceId == search || a.Id == search || EF.Functions.ILike(a.Title, likePattern, "\\") || (amountOrNull.HasValue && a.Amount == amountOrNull.Value) ); } else { queryable = context.PaymentRequests.AsQueryable(); if (!string.IsNullOrEmpty(query.StoreId)) { queryable = queryable.Where(data => data.StoreDataId == query.StoreId); } } if (query.LabelFilter is not null) { if (string.IsNullOrEmpty(query.StoreId)) throw new InvalidOperationException("PaymentRequestQuery.StoreId should be specified for label filtering"); var labels = query.LabelFilter; queryable = queryable.Where(pr => context.StoreLabelLinks.Any(l => l.StoreId == query.StoreId && l.ObjectId == pr.Id && l.StoreLabel.Type == WalletObjectData.Types.PaymentRequest && // ReSharper disable once CSharp14OverloadResolutionWithSpanBreakingChange labels.Contains(l.StoreLabel.Text))); } queryable = queryable.Include(data => data.StoreData); if (!query.IncludeArchived) queryable = queryable.Where(data => !data.Archived); if (query.Status != null && query.Status.Any()) queryable = queryable.Where(data => query.Status.Contains(data.Status)); if (query.Ids != null && query.Ids.Any()) queryable = queryable.Where(data => query.Ids.Contains(data.Id)); if (!string.IsNullOrEmpty(query.UserId)) queryable = queryable.Where(data => data.StoreData.UserStores.Any(u => u.ApplicationUserId == query.UserId)); if (query.StartDate.HasValue) queryable = queryable.Where(data => data.Created >= query.StartDate.Value); if (query.EndDate.HasValue) queryable = queryable.Where(data => data.Created <= query.EndDate.Value); queryable = queryable.OrderByDescending(u => u.Created); if (query.Skip.HasValue) queryable = queryable.Skip(query.Skip.Value); if (query.Count.HasValue) queryable = queryable.Take(query.Count.Value); var items = await queryable.ToArrayAsync(cancellationToken); return items; } public async Task GetInvoicesForPaymentRequest(string paymentRequestId, InvoiceQuery invoiceQuery = null) { if (invoiceQuery == null) { invoiceQuery = new InvoiceQuery(); } invoiceQuery.OrderId = new[] { GetOrderIdForPaymentRequest(paymentRequestId) }; return (await _InvoiceRepository.GetInvoices(invoiceQuery)) .Where(i => i.InternalTags.Contains(GetInternalTag(paymentRequestId))) .ToArray(); } public static string GetOrderIdForPaymentRequest(string paymentRequestId) { return $"PAY_REQUEST_{paymentRequestId}"; } public static string GetPaymentRequestIdFromOrderId(string invoiceOrderId) { if (string.IsNullOrEmpty(invoiceOrderId) || !invoiceOrderId.StartsWith("PAY_REQUEST_", StringComparison.InvariantCulture)) { return null; } return invoiceOrderId.Replace("PAY_REQUEST_", "", StringComparison.InvariantCulture); } public static string GetInternalTag(string id) { return $"PAYREQ#{id}"; } public static string[] GetPaymentIdsFromInternalTags(InvoiceEntity invoiceEntity) { return invoiceEntity.GetInternalTags("PAYREQ#"); } } public class PaymentRequestQuery { public string StoreId { get; set; } public bool IncludeArchived { get; set; } = true; public PaymentRequestStatus[] Status { get; set; } public string UserId { get; set; } public int? Skip { get; set; } public int? Count { get; set; } public string[] Ids { get; set; } public string SearchText { get; set; } public DateTimeOffset? StartDate { get; set; } public DateTimeOffset? EndDate { get; set; } public string[] LabelFilter { get; set; } } }