Compare commits

..
13 changed files with 48 additions and 132 deletions
-13
View File
@@ -122,19 +122,6 @@ namespace Notesnook.API.Hubs
await base.OnConnectedAsync(); await base.OnConnectedAsync();
} }
public override async Task OnDisconnectedAsync(Exception? exception)
{
if (exception != null)
{
Logger.LogWarning(exception, "Connection {ConnectionId} disconnected with error (server-side drop)", Context.ConnectionId);
}
else
{
Logger.LogInformation("Connection {ConnectionId} disconnected cleanly (client-initiated)", Context.ConnectionId);
}
await base.OnDisconnectedAsync(exception);
}
public async Task<int> PushItems(string deviceId, SyncTransferItemV2 pushItem) public async Task<int> PushItems(string deviceId, SyncTransferItemV2 pushItem)
{ {
-1
View File
@@ -9,7 +9,6 @@
<ItemGroup> <ItemGroup>
<PackageReference Include="AngleSharp" Version="1.3.0" /> <PackageReference Include="AngleSharp" Version="1.3.0" />
<PackageReference Include="AspNetCore.HealthChecks.Aws.S3" Version="9.0.0" /> <PackageReference Include="AspNetCore.HealthChecks.Aws.S3" Version="9.0.0" />
<PackageReference Include="AspNetCore.HealthChecks.Redis" Version="9.0.0" />
<PackageReference Include="AWSSDK.Core" Version="3.7.304.31" /> <PackageReference Include="AWSSDK.Core" Version="3.7.304.31" />
<PackageReference Include="DotNetEnv" Version="2.3.0" /> <PackageReference Include="DotNetEnv" Version="2.3.0" />
<PackageReference Include="IdentityModel.AspNetCore.OAuth2Introspection" Version="6.2.0" /> <PackageReference Include="IdentityModel.AspNetCore.OAuth2Introspection" Version="6.2.0" />
+1 -15
View File
@@ -25,7 +25,6 @@ using System.Text;
using System.Text.Encodings.Web; using System.Text.Encodings.Web;
using System.Threading.Tasks; using System.Threading.Tasks;
using Amazon.Runtime; using Amazon.Runtime;
using StackExchange.Redis;
using IdentityModel.AspNetCore.OAuth2Introspection; using IdentityModel.AspNetCore.OAuth2Introspection;
using Microsoft.AspNetCore.Authentication; using Microsoft.AspNetCore.Authentication;
using Microsoft.AspNetCore.Authentication.JwtBearer; using Microsoft.AspNetCore.Authentication.JwtBearer;
@@ -220,20 +219,7 @@ namespace Notesnook.API
}).AddMessagePackProtocol().AddJsonProtocol(); }).AddMessagePackProtocol().AddJsonProtocol();
if (!string.IsNullOrEmpty(Constants.SIGNALR_REDIS_CONNECTION_STRING)) if (!string.IsNullOrEmpty(Constants.SIGNALR_REDIS_CONNECTION_STRING))
{ signalR.AddStackExchangeRedis(Constants.SIGNALR_REDIS_CONNECTION_STRING);
services.AddHealthChecks()
.AddRedis(Constants.SIGNALR_REDIS_CONNECTION_STRING, tags: ["ready"]);
signalR.AddStackExchangeRedis(options =>
{
options.Configuration = ConfigurationOptions.Parse(Constants.SIGNALR_REDIS_CONNECTION_STRING);
options.Configuration.AbortOnConnectFail = false;
options.Configuration.ConnectRetry = 5;
options.Configuration.ReconnectRetryPolicy = new ExponentialRetry(5000, 30000);
options.Configuration.KeepAlive = 60;
options.Configuration.ConnectTimeout = 5000;
options.Configuration.SyncTimeout = 5000;
});
}
services.AddResponseCompression(options => services.AddResponseCompression(options =>
{ {
@@ -190,8 +190,8 @@ namespace Streetwriters.Identity.Controllers
var client = Clients.FindClientById(form.ClientId); var client = Clients.FindClientById(form.ClientId);
if (client == null) return BadRequest("Invalid client_id."); if (client == null) return BadRequest("Invalid client_id.");
var user = await UserManager.FindByEmailAsync(form.Email); var user = await UserManager.FindByEmailAsync(form.Email) ?? throw new Exception("User not found.");
if (user == null || !await UserService.IsUserValidAsync(UserManager, user, form.ClientId)) return Ok(); if (!await UserService.IsUserValidAsync(UserManager, user, form.ClientId)) return Ok();
var code = await UserManager.GenerateUserTokenAsync(user, TokenOptions.DefaultProvider, "ResetPassword"); var code = await UserManager.GenerateUserTokenAsync(user, TokenOptions.DefaultProvider, "ResetPassword");
var callbackUrl = UrlExtensions.TokenLink(user.Id.ToString(), code, client.Id, TokenType.RESET_PASSWORD); var callbackUrl = UrlExtensions.TokenLink(user.Id.ToString(), code, client.Id, TokenType.RESET_PASSWORD);
@@ -79,9 +79,12 @@ namespace Streetwriters.Identity.Controllers
} }
[HttpDelete] [HttpDelete]
public IActionResult Disable2FA() public async Task<IActionResult> Disable2FA()
{ {
return BadRequest("2FA is mandatory and cannot be disabled."); var user = await UserManager.GetUserAsync(User) ?? throw new Exception("User not found.");
if (!await UserManager.GetTwoFactorEnabledAsync(user)) return Ok();
await MFAService.DisableMFAAsync(user);
return Ok();
} }
[HttpGet("codes")] [HttpGet("codes")]
@@ -24,7 +24,7 @@ namespace Streetwriters.Identity.Interfaces
{ {
public interface ISMSSender public interface ISMSSender
{ {
Task<string?> SendOTPAsync(string number, IClient client); Task<string> SendOTPAsync(string number, IClient client);
Task<bool> VerifyOTPAsync(string id, string code); Task<bool> VerifyOTPAsync(string id, string code);
} }
} }
@@ -186,8 +186,6 @@ namespace Streetwriters.Identity.Services
ArgumentNullException.ThrowIfNull(form.PhoneNumber); ArgumentNullException.ThrowIfNull(form.PhoneNumber);
await UserManager.SetPhoneNumberAsync(user, form.PhoneNumber); await UserManager.SetPhoneNumberAsync(user, form.PhoneNumber);
var id = await SMSSender.SendOTPAsync(form.PhoneNumber, client); var id = await SMSSender.SendOTPAsync(form.PhoneNumber, client);
if (string.IsNullOrEmpty(id)) throw new Exception("Failed to send SMS. Please try again.");
logger.LogInformation("SMS OTP sent for user: {UserId}, SMS ID: {SmsId}", user.Id, id); logger.LogInformation("SMS OTP sent for user: {UserId}, SMS ID: {SmsId}", user.Id, id);
await this.ReplaceClaimAsync(user, MFAService.SMS_ID_CLAIM, id); await this.ReplaceClaimAsync(user, MFAService.SMS_ID_CLAIM, id);
break; break;
+13 -33
View File
@@ -23,56 +23,36 @@ using Streetwriters.Common;
using Twilio.Rest.Verify.V2.Service; using Twilio.Rest.Verify.V2.Service;
using Twilio; using Twilio;
using System.Threading.Tasks; using System.Threading.Tasks;
using System;
using Microsoft.Extensions.Logging;
namespace Streetwriters.Identity.Services namespace Streetwriters.Identity.Services
{ {
public class SMSSender : ISMSSender public class SMSSender : ISMSSender
{ {
private readonly ILogger<SMSSender> Logger; public SMSSender()
public SMSSender(ILogger<SMSSender> logger)
{ {
Logger = logger;
if (!string.IsNullOrEmpty(Constants.TWILIO_ACCOUNT_SID) && !string.IsNullOrEmpty(Constants.TWILIO_AUTH_TOKEN)) if (!string.IsNullOrEmpty(Constants.TWILIO_ACCOUNT_SID) && !string.IsNullOrEmpty(Constants.TWILIO_AUTH_TOKEN))
{ {
TwilioClient.Init(Constants.TWILIO_ACCOUNT_SID, Constants.TWILIO_AUTH_TOKEN); TwilioClient.Init(Constants.TWILIO_ACCOUNT_SID, Constants.TWILIO_AUTH_TOKEN);
} }
} }
public async Task<string?> SendOTPAsync(string number, IClient app) public async Task<string> SendOTPAsync(string number, IClient app)
{ {
try var verification = await VerificationResource.CreateAsync(
{ to: number,
var verification = await VerificationResource.CreateAsync( channel: "sms",
to: number, pathServiceSid: Constants.TWILIO_SERVICE_SID
channel: "sms", );
pathServiceSid: Constants.TWILIO_SERVICE_SID return verification.Sid;
);
return verification.Sid;
}
catch (Exception ex)
{
Logger.LogError(ex, "Error sending OTP with Twilio");
return null;
}
} }
public async Task<bool> VerifyOTPAsync(string id, string code) public async Task<bool> VerifyOTPAsync(string id, string code)
{ {
try return (await VerificationCheckResource.CreateAsync(
{ verificationSid: id,
return (await VerificationCheckResource.CreateAsync( pathServiceSid: Constants.TWILIO_SERVICE_SID,
verificationSid: id, code: code
pathServiceSid: Constants.TWILIO_SERVICE_SID, )).Status == "approved";
code: code
)).Status == "approved";
}
catch (Exception ex)
{
Logger.LogError(ex, "Error verifying OTP with Twilio");
return false;
}
} }
} }
} }
@@ -34,12 +34,6 @@ namespace Streetwriters.Identity.Services
var claims = await userManager.GetClaimsAsync(user); var claims = await userManager.GetClaimsAsync(user);
var marketingConsentClaim = claims.FirstOrDefault((claim) => claim.Type == $"{clientId}:marketing_consent"); var marketingConsentClaim = claims.FirstOrDefault((claim) => claim.Type == $"{clientId}:marketing_consent");
if (await userManager.IsEmailConfirmedAsync(user) && !await userManager.GetTwoFactorEnabledAsync(user))
{
await mfaService.EnableMFAAsync(user, MFAMethods.Email);
user = await userManager.FindByIdAsync(userId);
ArgumentNullException.ThrowIfNull(user);
}
ArgumentNullException.ThrowIfNull(user.Email); ArgumentNullException.ThrowIfNull(user.Email);
return new UserModel return new UserModel
@@ -163,6 +157,7 @@ namespace Streetwriters.Identity.Services
} }
else else
{ {
await mfaService.EnableMFAAsync(user, MFAMethods.Email);
if (userAgent != null) await userManager.AddClaimAsync(user, new Claim("platform", PlatformFromUserAgent(userAgent))); if (userAgent != null) await userManager.AddClaimAsync(user, new Claim("platform", PlatformFromUserAgent(userAgent)));
var code = await userManager.GenerateEmailConfirmationTokenAsync(user); var code = await userManager.GenerateEmailConfirmationTokenAsync(user);
var callbackUrl = UrlExtensions.TokenLink(user.Id.ToString(), code, client.Id, TokenType.CONFRIM_EMAIL); var callbackUrl = UrlExtensions.TokenLink(user.Id.ToString(), code, client.Id, TokenType.CONFRIM_EMAIL);
@@ -59,7 +59,6 @@ namespace Streetwriters.Identity.Validation
public string GrantType => Config.EMAIL_GRANT_TYPE; public string GrantType => Config.EMAIL_GRANT_TYPE;
public async Task ValidateAsync(ExtensionGrantValidationContext context) public async Task ValidateAsync(ExtensionGrantValidationContext context)
{ {
var email = context.Request.Raw["email"]; var email = context.Request.Raw["email"];
@@ -76,8 +75,14 @@ namespace Streetwriters.Identity.Validation
}; };
var isMultiFactor = await UserManager.GetTwoFactorEnabledAsync(user); var isMultiFactor = await UserManager.GetTwoFactorEnabledAsync(user);
if (!isMultiFactor)
{
context.Result.IsError = false;
context.Result.Subject = await TokenGenerationService.TransformTokenRequestAsync(context.Request, user, GrantType, [Config.MFA_PASSWORD_GRANT_TYPE_SCOPE]);
return;
}
var primaryMethod = isMultiFactor ? MFAService.GetPrimaryMethod(user) : MFAMethods.Email; var primaryMethod = MFAService.GetPrimaryMethod(user);
var secondaryMethod = MFAService.GetSecondaryMethod(user); var secondaryMethod = MFAService.GetSecondaryMethod(user);
var sendPhoneNumber = primaryMethod == MFAMethods.SMS || secondaryMethod == MFAMethods.SMS; var sendPhoneNumber = primaryMethod == MFAMethods.SMS || secondaryMethod == MFAMethods.SMS;
+9 -30
View File
@@ -18,49 +18,28 @@ along with this program. If not, see <http://www.gnu.org/licenses/>.
*/ */
using System.Linq; using System.Linq;
using System;
using System.Threading;
using System.Threading.Tasks; using System.Threading.Tasks;
using Lib.AspNetCore.ServerSentEvents; using Lib.AspNetCore.ServerSentEvents;
using System.Security.Claims; using System.Security.Claims;
using System.Collections.Generic;
namespace Streetwriters.Messenger.Helpers namespace Streetwriters.Messenger.Helpers
{ {
public class SSEHelper public class SSEHelper
{ {
public static async Task SendEventToUserAsync(string data, IServerSentEventsService sseService, string userId, string? originTokenId = null, CancellationToken cancellationToken = default) public static async Task SendEventToUserAsync(string data, IServerSentEventsService sseService, string userId, string? originTokenId = null)
{
var clients = sseService.GetClients()
.Where(c => c.User?.FindFirstValue("sub") == userId)
.Where(c => originTokenId == null || c.User?.FindFirstValue("jti") != originTokenId);
await SendEventToClientsAsync(clients, data, cancellationToken);
}
public static async Task SendEventToAllUsersAsync(string data, IServerSentEventsService sseService, CancellationToken cancellationToken = default)
{
await SendEventToClientsAsync(sseService.GetClients(), data, cancellationToken);
}
private static async Task SendEventToClientsAsync(IEnumerable<IServerSentEventsClient> clients, string data, CancellationToken cancellationToken)
{ {
var clients = sseService.GetClients().Where(c => c.User.FindFirstValue("sub") == userId);
foreach (var client in clients) foreach (var client in clients)
{ {
if (originTokenId != null && client.User.FindFirstValue("jti") == originTokenId) continue;
if (!client.IsConnected) continue; if (!client.IsConnected) continue;
await client.SendEventAsync(data);
try
{
await client.SendEventAsync(data, cancellationToken);
}
catch (OperationCanceledException) when (cancellationToken.IsCancellationRequested)
{
break;
}
catch
{
}
} }
} }
public static async Task SendEventToAllUsersAsync(string data, IServerSentEventsService sseService)
{
await sseService.SendEventAsync(data);
}
} }
} }
@@ -21,7 +21,6 @@ using System;
using System.Threading; using System.Threading;
using System.Threading.Tasks; using System.Threading.Tasks;
using Microsoft.Extensions.Hosting; using Microsoft.Extensions.Hosting;
using Microsoft.Extensions.Logging;
using Lib.AspNetCore.ServerSentEvents; using Lib.AspNetCore.ServerSentEvents;
using Streetwriters.Messenger.Helpers; using Streetwriters.Messenger.Helpers;
using System.Text.Json; using System.Text.Json;
@@ -34,14 +33,12 @@ namespace Streetwriters.Messenger.Services
private const string HEARTBEAT_MESSAGE_FORMAT = "Streetwriters Heartbeat ({0} UTC)"; private const string HEARTBEAT_MESSAGE_FORMAT = "Streetwriters Heartbeat ({0} UTC)";
private readonly IServerSentEventsService _serverSentEventsService; private readonly IServerSentEventsService _serverSentEventsService;
private readonly ILogger<HeartbeatService> _logger;
#endregion #endregion
#region Constructor #region Constructor
public HeartbeatService(IServerSentEventsService serverSentEventsService, ILogger<HeartbeatService> logger) public HeartbeatService(IServerSentEventsService serverSentEventsService)
{ {
_serverSentEventsService = serverSentEventsService; _serverSentEventsService = serverSentEventsService;
_logger = logger;
} }
#endregion #endregion
@@ -50,28 +47,15 @@ namespace Streetwriters.Messenger.Services
{ {
while (!stoppingToken.IsCancellationRequested) while (!stoppingToken.IsCancellationRequested)
{ {
try var message = JsonSerializer.Serialize(new
{ {
var message = JsonSerializer.Serialize(new type = "heartbeat",
data = JsonSerializer.Serialize(new
{ {
type = "heartbeat", t = DateTimeOffset.UtcNow.ToUnixTimeMilliseconds()
data = JsonSerializer.Serialize(new })
{ });
t = DateTimeOffset.UtcNow.ToUnixTimeMilliseconds() await SSEHelper.SendEventToAllUsersAsync(message, _serverSentEventsService);
})
});
await SSEHelper.SendEventToAllUsersAsync(message, _serverSentEventsService, stoppingToken);
}
catch (OperationCanceledException) when (stoppingToken.IsCancellationRequested)
{
break;
}
catch (Exception ex)
{
_logger.LogWarning(ex, "Failed to send SSE heartbeat to one or more clients.");
}
await Task.Delay(TimeSpan.FromSeconds(5), stoppingToken); await Task.Delay(TimeSpan.FromSeconds(5), stoppingToken);
} }
} }
@@ -8,7 +8,7 @@
<ItemGroup> <ItemGroup>
<PackageReference Include="DotNetEnv" Version="2.3.0" /> <PackageReference Include="DotNetEnv" Version="2.3.0" />
<PackageReference Include="Lib.AspNetCore.ServerSentEvents" Version="9.1.0" /> <PackageReference Include="Lib.AspNetCore.ServerSentEvents" Version="6.0.0" />
<PackageReference Include="Microsoft.AspNetCore.Authentication.JwtBearer" Version="5.0.0" <PackageReference Include="Microsoft.AspNetCore.Authentication.JwtBearer" Version="5.0.0"
NoWarn="NU1605" /> NoWarn="NU1605" />
<PackageReference Include="Microsoft.AspNetCore.Authentication.OpenIdConnect" Version="5.0.0" <PackageReference Include="Microsoft.AspNetCore.Authentication.OpenIdConnect" Version="5.0.0"