Nullify FileTransferService

This commit is contained in:
Jordan Dominion
2023-11-24 10:02:07 -05:00
parent a310fd8ec9
commit dffe0d1d08
2 changed files with 59 additions and 48 deletions
@@ -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
{
/// <summary>
@@ -71,6 +69,11 @@ namespace Tgstation.Server.Host.Transfer
/// </summary>
Task expireTask;
/// <summary>
/// If the <see cref="FileTransferService"/> is disposed.
/// </summary>
bool disposed;
/// <summary>
/// Initializes a new instance of the <see cref="FileTransferService"/> class.
/// </summary>
@@ -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,
};
}
/// <inheritdoc />
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;
}
/// <inheritdoc />
public async ValueTask<Tuple<Stream, ErrorMessageResponse>> RetrieveDownloadStream(FileTicketResponse ticket, CancellationToken cancellationToken)
public async ValueTask<Tuple<Stream?, ErrorMessageResponse?>> 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<Stream, ErrorMessageResponse>(null, null);
logger.LogTrace("Download ticket {ticket} not found!", ticket);
return Tuple.Create<Stream?, ErrorMessageResponse?>(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<Stream, ErrorMessageResponse>(null, new ErrorMessageResponse(errorCode.Value));
logger.LogDebug("Download ticket {ticket} failed activation!", ticket);
return Tuple.Create<Stream?, ErrorMessageResponse?>(null, new ErrorMessageResponse(errorCode.Value));
}
Stream stream;
@@ -198,7 +216,7 @@ namespace Tgstation.Server.Host.Transfer
}
catch (IOException ex)
{
return Tuple.Create<Stream, ErrorMessageResponse>(
return Tuple.Create<Stream?, ErrorMessageResponse?>(
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, ErrorMessageResponse>(stream, null);
logger.LogTrace("Ticket {ticket} downloading...", ticket);
return Tuple.Create<Stream?, ErrorMessageResponse?>(stream, null);
}
catch
{
@@ -219,34 +237,27 @@ namespace Tgstation.Server.Host.Transfer
}
/// <inheritdoc />
public async ValueTask<ErrorMessageResponse> SetUploadStream(FileTicketResponse ticket, Stream stream, CancellationToken cancellationToken)
public async ValueTask<ErrorMessageResponse?> 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);
}
/// <summary>
/// Creates a new <see cref="FileTicketResponse"/>.
/// </summary>
/// <returns>A new <see cref="FileTicketResponse"/>.</returns>
FileTicketResponse CreateTicket() => new()
{
FileTicket = cryptographySuite.GetSecureString(),
};
/// <summary>
/// Queue an <paramref name="expireAction"/> to run after <see cref="TicketValidityMinutes"/>.
/// </summary>
@@ -15,20 +15,20 @@ namespace Tgstation.Server.Host.Transfer
public interface IFileTransferStreamHandler
{
/// <summary>
/// Sets the <see cref="Stream"/> for a given <paramref name="ticket"/> associated with a pending upload.
/// Sets the <see cref="Stream"/> for a given <paramref name="ticketResponse"/> associated with a pending upload.
/// </summary>
/// <param name="ticket">The <see cref="FileTicketResponse"/>.</param>
/// <param name="ticketResponse">The <see cref="FileTicketResponse"/>.</param>
/// <param name="stream">The <see cref="Stream"/> with uploaded data.</param>
/// <param name="cancellationToken">The <see cref="CancellationToken"/> for the operation.</param>
/// <returns>A <see cref="ValueTask{TResult}"/> resulting in <see langword="null"/> if the upload completed successfully, <see cref="ErrorMessageResponse"/> otherwise.</returns>
ValueTask<ErrorMessageResponse> SetUploadStream(FileTicketResponse ticket, Stream stream, CancellationToken cancellationToken);
ValueTask<ErrorMessageResponse> SetUploadStream(FileTicketResponse ticketResponse, Stream stream, CancellationToken cancellationToken);
/// <summary>
/// Gets the the <see cref="Stream"/> for a given <paramref name="ticket"/> associated with a pending download.
/// Gets the the <see cref="Stream"/> for a given <paramref name="ticketResponse"/> associated with a pending download.
/// </summary>
/// <param name="ticket">The <see cref="FileTicketResponse"/>.</param>
/// <param name="ticketResponse">The <see cref="FileTicketResponse"/>.</param>
/// <param name="cancellationToken">The <see cref="CancellationToken"/> for the operation.</param>
/// <returns>A <see cref="ValueTask{TResult}"/> resulting in a <see cref="Tuple{T1, T2}"/> containing either a <see cref="Stream"/> containing the data to download or an <see cref="ErrorMessageResponse"/> to return.</returns>
ValueTask<Tuple<Stream, ErrorMessageResponse>> RetrieveDownloadStream(FileTicketResponse ticket, CancellationToken cancellationToken);
ValueTask<Tuple<Stream, ErrorMessageResponse>> RetrieveDownloadStream(FileTicketResponse ticketResponse, CancellationToken cancellationToken);
}
}