From dffe0d1d0880e901f21d07a0924af1c96b304ebe Mon Sep 17 00:00:00 2001 From: Jordan Dominion Date: Thu, 23 Nov 2023 20:24:23 -0500 Subject: [PATCH] Nullify `FileTransferService` --- .../Transfer/FileTransferService.cs | 95 +++++++++++-------- .../Transfer/IFileTransferStreamHandler.cs | 12 +-- 2 files changed, 59 insertions(+), 48 deletions(-) diff --git a/src/Tgstation.Server.Host/Transfer/FileTransferService.cs b/src/Tgstation.Server.Host/Transfer/FileTransferService.cs index cbdc528480..bb1cafa171 100644 --- a/src/Tgstation.Server.Host/Transfer/FileTransferService.cs +++ b/src/Tgstation.Server.Host/Transfer/FileTransferService.cs @@ -12,8 +12,6 @@ using Tgstation.Server.Host.IO; using Tgstation.Server.Host.Security; using Tgstation.Server.Host.Utils; -#nullable disable - namespace Tgstation.Server.Host.Transfer { /// @@ -71,6 +69,11 @@ namespace Tgstation.Server.Host.Transfer /// Task expireTask; + /// + /// If the is disposed. + /// + bool disposed; + /// /// Initializes a new instance of the class. /// @@ -103,12 +106,13 @@ namespace Tgstation.Server.Host.Transfer { Task toAwait; lock (synchronizationLock) - if (expireTask != null) + if (!disposed) { disposeCts.Cancel(); disposeCts.Dispose(); + disposed = true; toAwait = expireTask; - expireTask = null; + expireTask = Task.CompletedTask; } else toAwait = Task.CompletedTask; @@ -120,72 +124,86 @@ namespace Tgstation.Server.Host.Transfer public FileTicketResponse CreateDownload(FileDownloadProvider downloadProvider) { ArgumentNullException.ThrowIfNull(downloadProvider); + ObjectDisposedException.ThrowIf(disposed, this); logger.LogDebug("Creating download ticket for path {filePath}", downloadProvider.FilePath); - var ticketResult = CreateTicket(); + var ticket = cryptographySuite.GetSecureString(); lock (downloadTickets) - downloadTickets.Add(ticketResult.FileTicket, downloadProvider); + downloadTickets.Add(ticket, downloadProvider); QueueExpiry(() => { lock (downloadTickets) - if (downloadTickets.Remove(ticketResult.FileTicket)) - logger.LogTrace("Expired download ticket {ticket}...", ticketResult.FileTicket); + if (downloadTickets.Remove(ticket)) + logger.LogTrace("Expired download ticket {ticket}...", ticket); }); - logger.LogTrace("Created download ticket {ticket}", ticketResult.FileTicket); + logger.LogTrace("Created download ticket {ticket}", ticket); - return ticketResult; + return new FileTicketResponse + { + FileTicket = ticket, + }; } /// public IFileUploadTicket CreateUpload(FileUploadStreamKind streamKind) { + ObjectDisposedException.ThrowIf(disposed, this); + logger.LogDebug("Creating upload ticket..."); - var uploadTicket = new FileUploadProvider(CreateTicket(), streamKind); + var ticket = cryptographySuite.GetSecureString(); + var uploadTicket = new FileUploadProvider( + new FileTicketResponse + { + FileTicket = ticket, + }, + streamKind); lock (uploadTickets) - uploadTickets.Add(uploadTicket.Ticket.FileTicket, uploadTicket); + uploadTickets.Add(ticket, uploadTicket); QueueExpiry(() => { lock (uploadTickets) - if (uploadTickets.Remove(uploadTicket.Ticket.FileTicket)) - logger.LogTrace("Expired upload ticket {ticket}...", uploadTicket.Ticket.FileTicket); + if (uploadTickets.Remove(ticket)) + logger.LogTrace("Expired upload ticket {ticket}...", ticket); else return; uploadTicket.Expire(); }); - logger.LogTrace("Created upload ticket {ticket}", uploadTicket.Ticket.FileTicket); + logger.LogTrace("Created upload ticket {ticket}", ticket); return uploadTicket; } /// - public async ValueTask> RetrieveDownloadStream(FileTicketResponse ticket, CancellationToken cancellationToken) + public async ValueTask> RetrieveDownloadStream(FileTicketResponse ticketResponse, CancellationToken cancellationToken) { - ArgumentNullException.ThrowIfNull(ticket); + ArgumentNullException.ThrowIfNull(ticketResponse); + ObjectDisposedException.ThrowIf(disposed, this); - FileDownloadProvider downloadProvider; + var ticket = ticketResponse.FileTicket ?? throw new InvalidOperationException("ticketResponse must have FileTicket!"); + FileDownloadProvider? downloadProvider; lock (downloadTickets) { - if (!downloadTickets.TryGetValue(ticket.FileTicket, out downloadProvider)) + if (!downloadTickets.TryGetValue(ticket, out downloadProvider)) { - logger.LogTrace("Download ticket {ticket} not found!", ticket.FileTicket); - return Tuple.Create(null, null); + logger.LogTrace("Download ticket {ticket} not found!", ticket); + return Tuple.Create(null, null); } - downloadTickets.Remove(ticket.FileTicket); + downloadTickets.Remove(ticket); } var errorCode = downloadProvider.ActivationCallback(); if (errorCode.HasValue) { - logger.LogDebug("Download ticket {ticket} failed activation!", ticket.FileTicket); - return Tuple.Create(null, new ErrorMessageResponse(errorCode.Value)); + logger.LogDebug("Download ticket {ticket} failed activation!", ticket); + return Tuple.Create(null, new ErrorMessageResponse(errorCode.Value)); } Stream stream; @@ -198,7 +216,7 @@ namespace Tgstation.Server.Host.Transfer } catch (IOException ex) { - return Tuple.Create( + return Tuple.Create( null, new ErrorMessageResponse(ErrorCode.IOError) { @@ -208,8 +226,8 @@ namespace Tgstation.Server.Host.Transfer try { - logger.LogTrace("Ticket {ticket} downloading...", ticket.FileTicket); - return Tuple.Create(stream, null); + logger.LogTrace("Ticket {ticket} downloading...", ticket); + return Tuple.Create(stream, null); } catch { @@ -219,34 +237,27 @@ namespace Tgstation.Server.Host.Transfer } /// - public async ValueTask SetUploadStream(FileTicketResponse ticket, Stream stream, CancellationToken cancellationToken) + public async ValueTask SetUploadStream(FileTicketResponse ticketResponse, Stream stream, CancellationToken cancellationToken) { - ArgumentNullException.ThrowIfNull(ticket); + ArgumentNullException.ThrowIfNull(ticketResponse); + ObjectDisposedException.ThrowIf(disposed, this); - FileUploadProvider uploadProvider; + var ticket = ticketResponse.FileTicket ?? throw new InvalidOperationException("ticketResponse must have FileTicket!"); + FileUploadProvider? uploadProvider; lock (uploadTickets) { - if (!uploadTickets.TryGetValue(ticket.FileTicket, out uploadProvider)) + if (!uploadTickets.TryGetValue(ticket, out uploadProvider)) { - logger.LogTrace("Upload ticket {ticket} not found!", ticket.FileTicket); + logger.LogTrace("Upload ticket {ticket} not found!", ticket); return new ErrorMessageResponse(ErrorCode.ResourceNotPresent); } - uploadTickets.Remove(ticket.FileTicket); + uploadTickets.Remove(ticket); } return await uploadProvider.Completion(stream, cancellationToken); } - /// - /// Creates a new . - /// - /// A new . - FileTicketResponse CreateTicket() => new() - { - FileTicket = cryptographySuite.GetSecureString(), - }; - /// /// Queue an to run after . /// diff --git a/src/Tgstation.Server.Host/Transfer/IFileTransferStreamHandler.cs b/src/Tgstation.Server.Host/Transfer/IFileTransferStreamHandler.cs index f0ba7a396d..43ec027534 100644 --- a/src/Tgstation.Server.Host/Transfer/IFileTransferStreamHandler.cs +++ b/src/Tgstation.Server.Host/Transfer/IFileTransferStreamHandler.cs @@ -15,20 +15,20 @@ namespace Tgstation.Server.Host.Transfer public interface IFileTransferStreamHandler { /// - /// Sets the for a given associated with a pending upload. + /// Sets the for a given associated with a pending upload. /// - /// The . + /// The . /// The with uploaded data. /// The for the operation. /// A resulting in if the upload completed successfully, otherwise. - ValueTask SetUploadStream(FileTicketResponse ticket, Stream stream, CancellationToken cancellationToken); + ValueTask SetUploadStream(FileTicketResponse ticketResponse, Stream stream, CancellationToken cancellationToken); /// - /// Gets the the for a given associated with a pending download. + /// Gets the the for a given associated with a pending download. /// - /// The . + /// The . /// The for the operation. /// A resulting in a containing either a containing the data to download or an to return. - ValueTask> RetrieveDownloadStream(FileTicketResponse ticket, CancellationToken cancellationToken); + ValueTask> RetrieveDownloadStream(FileTicketResponse ticketResponse, CancellationToken cancellationToken); } }