feat: Implement message encryption and decryption support

- Added IMessageEncryptionService and its implementation MessageEncryptionService for handling message encryption.
- Updated ChannelsController and ChatService to encrypt messages before storing and sending.
- Introduced encryption key retrieval endpoint in ServerController.
- Modified EchoHubDbContext to accommodate increased message content and embed JSON lengths for encrypted data.
- Created migrations to support encryption-related database changes.
- Enhanced FirstRunSetup to ensure encryption key is generated if not present.
- Updated appsettings.example.json to include encryption configuration.
- Added comprehensive unit tests for encryption service and compatibility tests between client and server encryption.
This commit is contained in:
HueByte
2026-02-20 17:16:32 +01:00
parent dfa1b21d32
commit 2efe54e417
26 changed files with 1517 additions and 31 deletions
+35 -13
View File
@@ -16,6 +16,7 @@ public class ChatService : IChatService
private readonly PresenceTracker _presenceTracker;
private readonly IEnumerable<IChatBroadcaster> _broadcasters;
private readonly LinkEmbedService _embedService;
private readonly IMessageEncryptionService _encryption;
private readonly ILogger<ChatService> _logger;
public ChatService(
@@ -23,12 +24,14 @@ public class ChatService : IChatService
PresenceTracker presenceTracker,
IEnumerable<IChatBroadcaster> broadcasters,
LinkEmbedService embedService,
IMessageEncryptionService encryption,
ILogger<ChatService> logger)
{
_scopeFactory = scopeFactory;
_presenceTracker = presenceTracker;
_broadcasters = broadcasters;
_embedService = embedService;
_encryption = encryption;
_logger = logger;
}
@@ -142,14 +145,21 @@ public class ChatService : IChatService
if (!ValidationConstants.ChannelNameRegex().IsMatch(channelName))
return "Invalid channel name.";
if (string.IsNullOrWhiteSpace(content))
// Decrypt content (client sends encrypted; IRC sends plaintext — Decrypt handles both)
var plaintext = _encryption.Decrypt(content);
// Strip encryption prefix if a user typed it literally (prevents spoofing)
while (plaintext.StartsWith("$ENC$"))
plaintext = plaintext["$ENC$".Length..];
if (string.IsNullOrWhiteSpace(plaintext))
return "Message content cannot be empty.";
if (content.Length > HubConstants.MaxMessageLength)
if (plaintext.Length > HubConstants.MaxMessageLength)
return $"Message exceeds maximum length of {HubConstants.MaxMessageLength} characters.";
// Sanitize: collapse excessive newlines
content = SanitizeNewlines(content);
// Sanitize on plaintext: collapse excessive newlines
plaintext = SanitizeNewlines(plaintext);
using var scope = _scopeFactory.CreateScope();
var db = scope.ServiceProvider.GetRequiredService<EchoHubDbContext>();
@@ -175,35 +185,42 @@ public class ChatService : IChatService
}
}
// Attempt to fetch link embeds for URLs in the message
// Attempt to fetch link embeds for URLs in the plaintext message
List<EmbedDto>? embeds = null;
try
{
embeds = await _embedService.TryGetEmbedsAsync(content);
embeds = await _embedService.TryGetEmbedsAsync(plaintext);
}
catch (Exception ex)
{
_logger.LogWarning(ex, "Failed to fetch link embeds for message in '{Channel}'", channelName);
}
// Store in DB — encrypted at rest if enabled, plaintext otherwise
var embedJson = embeds is not null ? JsonSerializer.Serialize(embeds) : null;
var dbContent = _encryption.EncryptDatabaseEnabled ? _encryption.Encrypt(plaintext) : plaintext;
var dbEmbedJson = _encryption.EncryptDatabaseEnabled ? _encryption.EncryptNullable(embedJson) : embedJson;
var message = new Message
{
Id = Guid.NewGuid(),
Content = content,
Content = dbContent,
Type = MessageType.Text,
SentAt = DateTimeOffset.UtcNow,
ChannelId = channel.Id,
SenderUserId = userId,
SenderUsername = username,
EmbedJson = embeds is not null ? JsonSerializer.Serialize(embeds) : null,
EmbedJson = dbEmbedJson,
};
db.Messages.Add(message);
await db.SaveChangesAsync();
// Broadcast encrypted for SignalR clients; IRC broadcaster gets plaintext
var encryptedContent = _encryption.Encrypt(plaintext);
var messageDto = new MessageDto(
message.Id,
message.Content,
encryptedContent,
message.SenderUsername,
sender?.NicknameColor,
channelName,
@@ -398,7 +415,7 @@ public class ChatService : IChatService
return string.Join('\n', result);
}
private static async Task<List<MessageDto>> GetChannelHistoryInternalAsync(EchoHubDbContext db, string channelName, int count)
private async Task<List<MessageDto>> GetChannelHistoryInternalAsync(EchoHubDbContext db, string channelName, int count)
{
var channel = await db.Channels.FirstOrDefaultAsync(c => c.Name == channelName);
if (channel is null)
@@ -418,16 +435,21 @@ public class ChatService : IChatService
return raw.Select(x =>
{
// Decrypt DB content (handles both encrypted and plaintext via prefix detection)
var plaintext = _encryption.Decrypt(x.m.Content);
var embedJsonPlain = _encryption.DecryptNullable(x.m.EmbedJson);
List<EmbedDto>? embeds = null;
if (x.m.EmbedJson is not null)
if (embedJsonPlain is not null)
{
try { embeds = JsonSerializer.Deserialize<List<EmbedDto>>(x.m.EmbedJson); }
try { embeds = JsonSerializer.Deserialize<List<EmbedDto>>(embedJsonPlain); }
catch { /* ignore malformed JSON */ }
}
// Encrypt for transport — client decrypts
return new MessageDto(
x.m.Id,
x.m.Content,
_encryption.Encrypt(plaintext),
x.m.SenderUsername,
x.NicknameColor,
channelName,
@@ -0,0 +1,102 @@
using System.Security.Cryptography;
using System.Text;
using EchoHub.Core.Contracts;
using Microsoft.Extensions.Configuration;
using Microsoft.Extensions.Logging;
namespace EchoHub.Server.Services;
public class MessageEncryptionService : IMessageEncryptionService
{
private const string EncryptionPrefix = "$ENC$v1$";
private const int NonceSizeBytes = 12;
private const int TagSizeBytes = 16;
private readonly byte[] _key;
private readonly ILogger<MessageEncryptionService> _logger;
public bool EncryptDatabaseEnabled { get; }
public MessageEncryptionService(IConfiguration configuration, ILogger<MessageEncryptionService> logger)
{
_logger = logger;
var keyBase64 = configuration["Encryption:Key"]
?? throw new InvalidOperationException("Encryption:Key must be configured in appsettings.json.");
_key = Convert.FromBase64String(keyBase64);
if (_key.Length != 32)
throw new InvalidOperationException($"Encryption:Key must be exactly 32 bytes (256-bit). Got {_key.Length} bytes.");
EncryptDatabaseEnabled = configuration.GetValue<bool>("Encryption:EncryptDatabase");
}
public string Encrypt(string plaintext)
{
var plaintextBytes = Encoding.UTF8.GetBytes(plaintext);
var nonce = RandomNumberGenerator.GetBytes(NonceSizeBytes);
var ciphertext = new byte[plaintextBytes.Length];
var tag = new byte[TagSizeBytes];
using var aes = new AesGcm(_key, TagSizeBytes);
aes.Encrypt(nonce, plaintextBytes, ciphertext, tag);
// Combine ciphertext + tag for storage
var combined = new byte[ciphertext.Length + tag.Length];
Buffer.BlockCopy(ciphertext, 0, combined, 0, ciphertext.Length);
Buffer.BlockCopy(tag, 0, combined, ciphertext.Length, tag.Length);
return $"{EncryptionPrefix}{Convert.ToBase64String(nonce)}${Convert.ToBase64String(combined)}";
}
public string Decrypt(string content)
{
if (!content.StartsWith(EncryptionPrefix))
return content; // Legacy plaintext
try
{
var payload = content[EncryptionPrefix.Length..];
var separatorIndex = payload.IndexOf('$');
if (separatorIndex < 0)
{
_logger.LogWarning("Malformed encrypted content: missing separator");
return "[encrypted message — decryption failed]";
}
var nonceBase64 = payload[..separatorIndex];
var combinedBase64 = payload[(separatorIndex + 1)..];
var nonce = Convert.FromBase64String(nonceBase64);
var combined = Convert.FromBase64String(combinedBase64);
if (combined.Length < TagSizeBytes)
{
_logger.LogWarning("Malformed encrypted content: data too short");
return "[encrypted message — decryption failed]";
}
var ciphertextLength = combined.Length - TagSizeBytes;
var ciphertext = combined.AsSpan(0, ciphertextLength);
var tag = combined.AsSpan(ciphertextLength, TagSizeBytes);
var plaintext = new byte[ciphertextLength];
using var aes = new AesGcm(_key, TagSizeBytes);
aes.Decrypt(nonce, ciphertext, tag, plaintext);
return Encoding.UTF8.GetString(plaintext);
}
catch (Exception ex)
{
_logger.LogError(ex, "Failed to decrypt message content");
return "[encrypted message — decryption failed]";
}
}
public string? EncryptNullable(string? value)
=> value is null ? null : Encrypt(value);
public string? DecryptNullable(string? value)
=> value is null ? null : Decrypt(value);
}