diff --git a/src/Orbit.Api/Controllers/AiController.cs b/src/Orbit.Api/Controllers/AiController.cs index d575d247..3f526dfe 100644 --- a/src/Orbit.Api/Controllers/AiController.cs +++ b/src/Orbit.Api/Controllers/AiController.cs @@ -1,3 +1,4 @@ +using System.Text.Json.Serialization; using Microsoft.AspNetCore.Authorization; using Microsoft.AspNetCore.Mvc; using Orbit.Api.Extensions; @@ -144,7 +145,7 @@ await auditService.RecordAsync(new AgentAuditEntry( public record ConfirmPendingOperationResponse(Guid PendingOperationId, string ConfirmationToken, DateTime ExpiresAtUtc); public record StepUpChallengeRequest(string Language = "en"); - public record VerifyStepUpRequest(Guid ChallengeId, string Code); + public record VerifyStepUpRequest([property: JsonRequired] Guid ChallengeId, string Code); public record ExecutePendingOperationRequest(string ConfirmationToken); [HttpPost("pending-operations/{id:guid}/confirm")] diff --git a/src/Orbit.Api/Controllers/OAuthController.cs b/src/Orbit.Api/Controllers/OAuthController.cs index f68dbb9c..2e9054ce 100644 --- a/src/Orbit.Api/Controllers/OAuthController.cs +++ b/src/Orbit.Api/Controllers/OAuthController.cs @@ -28,6 +28,7 @@ public partial class OAuthController( IConfiguration configuration, ILogger logger) : ControllerBase { + private const string InvalidRedirectUriError = "invalid_redirect_uri"; private static readonly string[] SupportedResponseTypes = ["code"]; private static readonly string[] SupportedGrantTypes = ["authorization_code"]; private static readonly string[] SupportedCodeChallengeMethods = ["S256"]; @@ -107,7 +108,7 @@ public IActionResult Authorize( return BadRequest(new { error = "PKCE with S256 is required" }); if (!IsRedirectUriAllowed(redirect_uri)) - return BadRequest(new { error = "invalid_redirect_uri" }); + return BadRequest(new { error = InvalidRedirectUriError }); var googleClientId = googleSettings.Value.ClientId ?? ""; var html = OAuthLoginPage.Render( @@ -139,7 +140,7 @@ public record VerifyCodeRequest( public async Task VerifyCode([FromBody] VerifyCodeRequest request, CancellationToken ct) { if (!IsRedirectUriAllowed(request.RedirectUri)) - return BadRequest(new { error = "invalid_redirect_uri" }); + return BadRequest(new { error = InvalidRedirectUriError }); var result = await mediator.Send( new VerifyCodeCommand(request.Email, request.Code), ct); @@ -166,7 +167,7 @@ public record GoogleAuthRequest( public async Task GoogleAuth([FromBody] GoogleAuthRequest request, CancellationToken ct) { if (!IsRedirectUriAllowed(request.RedirectUri)) - return BadRequest(new { error = "invalid_redirect_uri" }); + return BadRequest(new { error = InvalidRedirectUriError }); // Validate Google ID token directly (GIS returns a JWT, not a Supabase token) var client = httpClientFactory.CreateClient(); @@ -236,7 +237,7 @@ public async Task Token( return BadRequest(new { error = "unsupported_grant_type" }); if (!IsRedirectUriAllowed(redirect_uri)) - return BadRequest(new { error = "invalid_redirect_uri" }); + return BadRequest(new { error = InvalidRedirectUriError }); var entry = authStore.ExchangeCode(code, code_verifier, redirect_uri); if (entry is null) diff --git a/src/Orbit.Api/Extensions/WebApplicationExtensions.cs b/src/Orbit.Api/Extensions/WebApplicationExtensions.cs index cdb5c3ab..146641b5 100644 --- a/src/Orbit.Api/Extensions/WebApplicationExtensions.cs +++ b/src/Orbit.Api/Extensions/WebApplicationExtensions.cs @@ -75,162 +75,214 @@ private static void UseMcpSelectiveAuth(this WebApplication app) { app.Use(async (context, next) => { - if (context.Request.Path.StartsWithSegments("/mcp") && context.Request.Method == "POST") + if (!IsMcpPostRequest(context)) { - context.Request.EnableBuffering(); - using var reader = new StreamReader(context.Request.Body, leaveOpen: true); - var body = await reader.ReadToEndAsync(); - context.Request.Body.Position = 0; + await next(); + return; + } - if (IsMcpUnauthenticatedMethod(body)) - { - await next(); - return; - } + await HandleMcpRequestAsync(context, next); + }); + } - // For tool calls, require auth - var authResult = await context.AuthenticateAsync(); - if (!authResult.Succeeded) - { - var scheme = context.Request.Headers["X-Forwarded-Proto"].FirstOrDefault() ?? context.Request.Scheme; - var resourceUrl = $"{scheme}://{context.Request.Host}/.well-known/oauth-protected-resource"; - context.Response.StatusCode = 401; - context.Response.Headers.WWWAuthenticate = $"Bearer resource_metadata=\"{resourceUrl}\""; - return; - } - context.User = authResult.Principal!; + private static bool IsMcpPostRequest(HttpContext context) + { + return context.Request.Path.StartsWithSegments("/mcp") && HttpMethods.IsPost(context.Request.Method); + } - if (TryGetMcpToolCall(body, out var toolName, out var requestId, out var operationId, out var operationFingerprint)) - { - if (string.Equals(toolName, "execute_agent_operation_v2", StringComparison.OrdinalIgnoreCase)) - { - await next(); - return; - } + private static async Task HandleMcpRequestAsync(HttpContext context, Func next) + { + var body = await ReadBufferedRequestBodyAsync(context); + if (IsMcpUnauthenticatedMethod(body)) + { + await next(); + return; + } - var catalogService = context.RequestServices.GetRequiredService(); - var policyEvaluator = context.RequestServices.GetRequiredService(); - var auditService = context.RequestServices.GetRequiredService(); - var capability = catalogService.GetCapabilityByMcpTool(toolName!); + if (!await TryAuthenticateMcpRequestAsync(context)) + return; - if (capability is null) - { - await TryAuditLegacyMcpAsync( - auditService, - context, - toolName!, - null, - AgentPolicyDecisionStatus.Denied, - AgentOperationStatus.UnsupportedByPolicy, - "unsupported_by_policy", - body, - null, - null, - CancellationToken.None); - - await WriteMcpPolicyErrorAsync( - context, - requestId, - "unsupported_by_policy", - null); - return; - } - - var decision = policyEvaluator.Evaluate(new AgentPolicyEvaluationContext( - capability.Id, - context.User.GetUserId(), - AgentExecutionSurface.Mcp, - context.User.GetAgentAuthMethod(), - context.User.GetGrantedAgentScopes(), - operationId ?? toolName!, - $"{operationId ?? toolName} requested via MCP", - operationFingerprint, - IsReadOnlyCredential: context.User.IsReadOnlyCredential())); - - if (decision.Status == AgentPolicyDecisionStatus.ConfirmationRequired) - { - await TryAuditLegacyMcpAsync( - auditService, - context, - operationId ?? toolName!, - capability, - AgentPolicyDecisionStatus.ConfirmationRequired, - AgentOperationStatus.PendingConfirmation, - decision.Reason, - body, - decision.ShadowStatus, - decision.ShadowReason, - CancellationToken.None); - - await WriteMcpPolicyErrorAsync( - context, - requestId, - decision.Reason ?? "confirmation_required", - decision.PendingOperation?.Id); - return; - } - - if (decision.Status == AgentPolicyDecisionStatus.Denied) - { - await TryAuditLegacyMcpAsync( - auditService, - context, - operationId ?? toolName!, - capability, - AgentPolicyDecisionStatus.Denied, - AgentOperationStatus.Denied, - decision.Reason, - body, - decision.ShadowStatus, - decision.ShadowReason, - CancellationToken.None); - - await WriteMcpPolicyErrorAsync( - context, - requestId, - decision.Reason ?? "policy_denied", - null); - return; - } - - try - { - await next(); - await TryAuditLegacyMcpAsync( - auditService, - context, - operationId ?? toolName!, - capability, - AgentPolicyDecisionStatus.Allowed, - AgentOperationStatus.Succeeded, - null, - body, - decision.ShadowStatus, - decision.ShadowReason, - CancellationToken.None); - } - catch (Exception ex) - { - await TryAuditLegacyMcpAsync( - auditService, - context, - operationId ?? toolName!, - capability, - AgentPolicyDecisionStatus.Allowed, - AgentOperationStatus.Failed, - ex.Message, - body, - decision.ShadowStatus, - decision.ShadowReason, - CancellationToken.None); - throw; - } - - return; - } - } + if (!TryGetMcpToolCall(body, out var toolName, out var requestId, out var operationId, out var operationFingerprint)) + { await next(); - }); + return; + } + + await HandleMcpToolCallAsync( + context, + next, + body, + new McpToolCallRequest(toolName!, requestId, operationId, operationFingerprint)); + } + + private static async Task ReadBufferedRequestBodyAsync(HttpContext context) + { + context.Request.EnableBuffering(); + using var reader = new StreamReader(context.Request.Body, leaveOpen: true); + var body = await reader.ReadToEndAsync(); + context.Request.Body.Position = 0; + return body; + } + + private static async Task TryAuthenticateMcpRequestAsync(HttpContext context) + { + var authResult = await context.AuthenticateAsync(); + if (authResult.Succeeded) + { + context.User = authResult.Principal!; + return true; + } + + var scheme = context.Request.Headers["X-Forwarded-Proto"].FirstOrDefault() ?? context.Request.Scheme; + var resourceUrl = $"{scheme}://{context.Request.Host}/.well-known/oauth-protected-resource"; + context.Response.StatusCode = StatusCodes.Status401Unauthorized; + context.Response.Headers.WWWAuthenticate = $"Bearer resource_metadata=\"{resourceUrl}\""; + return false; + } + + private static async Task HandleMcpToolCallAsync( + HttpContext context, + Func next, + string body, + McpToolCallRequest toolCall) + { + if (string.Equals(toolCall.ToolName, "execute_agent_operation_v2", StringComparison.OrdinalIgnoreCase)) + { + await next(); + return; + } + + var services = ResolveMcpServices(context); + var capability = services.CatalogService.GetCapabilityByMcpTool(toolCall.ToolName); + if (capability is null) + { + await AuditAndWritePolicyErrorAsync( + context, + services.AuditService, + new LegacyMcpAuditContext( + toolCall.ToolName, + null, + AgentPolicyDecisionStatus.Denied, + AgentOperationStatus.UnsupportedByPolicy, + "unsupported_by_policy", + body), + toolCall.RequestId, + "unsupported_by_policy", + null); + return; + } + + var sourceName = toolCall.OperationId ?? toolCall.ToolName; + var decision = services.PolicyEvaluator.Evaluate(new AgentPolicyEvaluationContext( + capability.Id, + context.User.GetUserId(), + AgentExecutionSurface.Mcp, + context.User.GetAgentAuthMethod(), + context.User.GetGrantedAgentScopes(), + sourceName, + $"{sourceName} requested via MCP", + toolCall.OperationFingerprint, + IsReadOnlyCredential: context.User.IsReadOnlyCredential())); + + if (decision.Status == AgentPolicyDecisionStatus.ConfirmationRequired) + { + await AuditAndWritePolicyErrorAsync( + context, + services.AuditService, + new LegacyMcpAuditContext( + sourceName, + capability, + AgentPolicyDecisionStatus.ConfirmationRequired, + AgentOperationStatus.PendingConfirmation, + decision.Reason, + body, + decision.ShadowStatus, + decision.ShadowReason), + toolCall.RequestId, + decision.Reason ?? "confirmation_required", + decision.PendingOperation?.Id); + return; + } + + if (decision.Status == AgentPolicyDecisionStatus.Denied) + { + await AuditAndWritePolicyErrorAsync( + context, + services.AuditService, + new LegacyMcpAuditContext( + sourceName, + capability, + AgentPolicyDecisionStatus.Denied, + AgentOperationStatus.Denied, + decision.Reason, + body, + decision.ShadowStatus, + decision.ShadowReason), + toolCall.RequestId, + decision.Reason ?? "policy_denied", + null); + return; + } + + await ExecuteAuthorizedMcpToolCallAsync( + context, + next, + services.AuditService, + new LegacyMcpAuditContext( + sourceName, + capability, + AgentPolicyDecisionStatus.Allowed, + AgentOperationStatus.Succeeded, + null, + body, + decision.ShadowStatus, + decision.ShadowReason)); + } + + private static McpServices ResolveMcpServices(HttpContext context) + { + return new McpServices( + context.RequestServices.GetRequiredService(), + context.RequestServices.GetRequiredService(), + context.RequestServices.GetRequiredService()); + } + + private static async Task AuditAndWritePolicyErrorAsync( + HttpContext context, + IAgentAuditService auditService, + LegacyMcpAuditContext auditContext, + JsonElement? requestId, + string reason, + Guid? pendingOperationId) + { + await TryAuditLegacyMcpAsync(auditService, context, auditContext, CancellationToken.None); + await WriteMcpPolicyErrorAsync(context, requestId, reason, pendingOperationId); + } + + private static async Task ExecuteAuthorizedMcpToolCallAsync( + HttpContext context, + Func next, + IAgentAuditService auditService, + LegacyMcpAuditContext auditContext) + { + try + { + await next(); + await TryAuditLegacyMcpAsync(auditService, context, auditContext, CancellationToken.None); + } + catch (Exception ex) + { + await TryAuditLegacyMcpAsync( + auditService, + context, + auditContext with + { + OutcomeStatus = AgentOperationStatus.Failed, + Error = ex.Message + }, + CancellationToken.None); + throw; + } } private static ForwardedHeadersOptions BuildForwardedHeadersOptions(WebApplication app) @@ -365,41 +417,57 @@ await context.Response.WriteAsJsonAsync(new private static async Task TryAuditLegacyMcpAsync( IAgentAuditService auditService, HttpContext context, - string sourceName, - AgentCapability? capability, - AgentPolicyDecisionStatus policyDecision, - AgentOperationStatus outcomeStatus, - string? error, - string rawBody, - AgentPolicyDecisionStatus? shadowPolicyDecision, - string? shadowReason, + LegacyMcpAuditContext auditContext, CancellationToken cancellationToken) { - if (capability is null || context.User.Identity?.IsAuthenticated != true) + if (auditContext.Capability is null || context.User.Identity?.IsAuthenticated != true) return; try { - var redactedArguments = rawBody.Length <= 1000 ? rawBody : rawBody[..1000]; + var redactedArguments = auditContext.RawBody.Length <= 1000 + ? auditContext.RawBody + : auditContext.RawBody[..1000]; await auditService.RecordAsync(new AgentAuditEntry( context.User.GetUserId(), - capability.Id, - sourceName, + auditContext.Capability.Id, + auditContext.SourceName, AgentExecutionSurface.Mcp, context.User.GetAgentAuthMethod(), - capability.RiskClass, - policyDecision, - outcomeStatus, + auditContext.Capability.RiskClass, + auditContext.PolicyDecision, + auditContext.OutcomeStatus, context.TraceIdentifier, - $"{sourceName} requested via MCP", + $"{auditContext.SourceName} requested via MCP", RedactedArguments: redactedArguments, - Error: error, - ShadowPolicyDecision: shadowPolicyDecision, - ShadowReason: shadowReason), cancellationToken); + Error: auditContext.Error, + ShadowPolicyDecision: auditContext.ShadowPolicyDecision, + ShadowReason: auditContext.ShadowReason), cancellationToken); } catch { // Audit failures must not block MCP requests. } } + + private sealed record McpToolCallRequest( + string ToolName, + JsonElement? RequestId, + string? OperationId, + string? OperationFingerprint); + + private sealed record McpServices( + IAgentCatalogService CatalogService, + IAgentPolicyEvaluator PolicyEvaluator, + IAgentAuditService AuditService); + + private sealed record LegacyMcpAuditContext( + string SourceName, + AgentCapability? Capability, + AgentPolicyDecisionStatus PolicyDecision, + AgentOperationStatus OutcomeStatus, + string? Error, + string RawBody, + AgentPolicyDecisionStatus? ShadowPolicyDecision = null, + string? ShadowReason = null); } diff --git a/src/Orbit.Application/Chat/Commands/ProcessUserChatCommand.cs b/src/Orbit.Application/Chat/Commands/ProcessUserChatCommand.cs index 6d09d7df..6205f5e8 100644 --- a/src/Orbit.Application/Chat/Commands/ProcessUserChatCommand.cs +++ b/src/Orbit.Application/Chat/Commands/ProcessUserChatCommand.cs @@ -85,6 +85,7 @@ public partial class ProcessUserChatCommandHandler( ILogger logger) : IRequestHandler> { private const int MaxToolIterations = 5; + private const string UnsupportedByPolicyReason = "unsupported_by_policy"; public async Task> Handle( ProcessUserChatCommand request, @@ -177,10 +178,7 @@ public async Task> Handle( var aiResponse = response.Value; // 4. Agentic tool-calling loop - var allActionResults = new List(); - var allOperationResults = new List(); - var allPendingOperations = new List(); - var allPolicyDenials = new List(); + var executionResults = new ToolExecutionAccumulator(); var actionsStopwatch = System.Diagnostics.Stopwatch.StartNew(); int iteration = 0; @@ -190,10 +188,7 @@ public async Task> Handle( var continueResponse = await ProcessToolCallsAsync( aiResponse, request, - allActionResults, - allOperationResults, - allPendingOperations, - allPolicyDenials, + executionResults, iteration, cancellationToken); @@ -204,14 +199,14 @@ public async Task> Handle( } actionsStopwatch.Stop(); - LogToolExecutionCompleted(logger, actionsStopwatch.ElapsedMilliseconds, iteration, allActionResults.Count); + LogToolExecutionCompleted(logger, actionsStopwatch.ElapsedMilliseconds, iteration, executionResults.ActionResults.Count); // 5. Persist all changes in a single unit of work LogSavingChanges(logger); var saveStopwatch = System.Diagnostics.Stopwatch.StartNew(); await execution.UnitOfWork.SaveChangesAsync(cancellationToken); - if (RequiresStreakRecalculation(allActionResults)) + if (RequiresStreakRecalculation(executionResults.ActionResults)) { await execution.UserStreakService.RecalculateAsync(request.UserId, cancellationToken); await execution.UnitOfWork.SaveChangesAsync(cancellationToken); @@ -240,10 +235,10 @@ public async Task> Handle( return Result.Success(new ChatResponse( aiMessage, - allActionResults, - allOperationResults, - allPendingOperations, - allPolicyDenials)); + executionResults.ActionResults, + executionResults.OperationResults, + executionResults.PendingOperations, + executionResults.PolicyDenials)); } /// @@ -253,10 +248,7 @@ public async Task> Handle( private async Task ProcessToolCallsAsync( AiResponse aiResponse, ProcessUserChatCommand request, - List allActionResults, - List allOperationResults, - List allPendingOperations, - List allPolicyDenials, + ToolExecutionAccumulator executionResults, int iteration, CancellationToken cancellationToken) { @@ -278,14 +270,7 @@ public async Task> Handle( var (toolCallResult, actionResult, operationResult, policyDenial, pendingOperation) = await ExecuteSingleToolCallAsync(call, request, cancellationToken); toolResults.Add(toolCallResult); - if (actionResult is not null) - allActionResults.Add(actionResult); - if (operationResult is not null) - allOperationResults.Add(operationResult); - if (policyDenial is not null) - allPolicyDenials.Add(policyDenial); - if (pendingOperation is not null) - allPendingOperations.Add(pendingOperation); + executionResults.Add(actionResult, operationResult, policyDenial, pendingOperation); } // Send results back to the AI for next iteration or final message @@ -326,13 +311,13 @@ public async Task> Handle( AgentRiskClass.Low, AgentConfirmationRequirement.None, AgentOperationStatus.UnsupportedByPolicy, - PolicyReason: "unsupported_by_policy"), + PolicyReason: UnsupportedByPolicyReason), new AgentPolicyDenial( call.Name, call.Name, AgentRiskClass.Low, AgentConfirmationRequirement.None, - "unsupported_by_policy"), + UnsupportedByPolicyReason), null); } @@ -351,13 +336,13 @@ public async Task> Handle( AgentConfirmationRequirement.None, AgentOperationStatus.UnsupportedByPolicy, Summary: operationSummary, - PolicyReason: "unsupported_by_policy"), + PolicyReason: UnsupportedByPolicyReason), new AgentPolicyDenial( call.Name, call.Name, AgentRiskClass.Low, AgentConfirmationRequirement.None, - "unsupported_by_policy"), + UnsupportedByPolicyReason), null); } @@ -553,6 +538,33 @@ private static bool RequiresStreakRecalculation(IEnumerable action return actionResults.Any(action => action.Status == ActionStatus.Success && action.Type is "LogHabit" or "BulkLogHabits" or "DeleteHabit"); } + private sealed class ToolExecutionAccumulator + { + public List ActionResults { get; } = []; + public List OperationResults { get; } = []; + public List PendingOperations { get; } = []; + public List PolicyDenials { get; } = []; + + public void Add( + ActionResult? actionResult, + AgentOperationResult? operationResult, + AgentPolicyDenial? policyDenial, + PendingAgentOperation? pendingOperation) + { + if (actionResult is not null) + ActionResults.Add(actionResult); + + if (operationResult is not null) + OperationResults.Add(operationResult); + + if (policyDenial is not null) + PolicyDenials.Add(policyDenial); + + if (pendingOperation is not null) + PendingOperations.Add(pendingOperation); + } + } + /// /// Fires off background work for fact extraction and AI message counter increment. /// Runs in a separate DI scope so it doesn't block the response. diff --git a/src/Orbit.Application/Chat/Tools/Implementations/ProfileTools.cs b/src/Orbit.Application/Chat/Tools/Implementations/ProfileTools.cs index 4337bca4..360ca640 100644 --- a/src/Orbit.Application/Chat/Tools/Implementations/ProfileTools.cs +++ b/src/Orbit.Application/Chat/Tools/Implementations/ProfileTools.cs @@ -138,6 +138,9 @@ private async Task ExecuteAsync(IRequest public class UpdateAiSettingsTool(IMediator mediator) : IAiTool { + private const string EnabledState = "enabled"; + private const string DisabledState = "disabled"; + public string Name => "update_ai_settings"; public string Description => "Enable or disable AI memory and daily AI summary settings."; @@ -174,10 +177,15 @@ public async Task ExecuteAsync(JsonElement args, Guid userId, Cancel ? new ToolResult( true, EntityId: userId.ToString(), - EntityName: action == "set_ai_memory" - ? $"AI memory {(enabled.Value ? "enabled" : "disabled")}" - : $"AI summary {(enabled.Value ? "enabled" : "disabled")}", + EntityName: BuildEntityName(action, enabled.Value), Payload: new { action, enabled }) : new ToolResult(false, EntityId: userId.ToString(), Error: result.Error); } + + private static string BuildEntityName(string action, bool enabled) + { + var settingName = action == "set_ai_memory" ? "AI memory" : "AI summary"; + var state = enabled ? EnabledState : DisabledState; + return $"{settingName} {state}"; + } } diff --git a/src/Orbit.Application/Chat/Tools/Implementations/QueryGoalsTool.cs b/src/Orbit.Application/Chat/Tools/Implementations/QueryGoalsTool.cs index 71197435..0c9e5afa 100644 --- a/src/Orbit.Application/Chat/Tools/Implementations/QueryGoalsTool.cs +++ b/src/Orbit.Application/Chat/Tools/Implementations/QueryGoalsTool.cs @@ -92,7 +92,7 @@ private static bool MatchesSearch(Goal goal, string? search) } private static string BuildOutput( - IReadOnlyList goals, + List goals, bool includeDescriptions, bool includeLinkedHabits) { diff --git a/src/Orbit.Application/Chat/Tools/Implementations/QueryHabitsTool.cs b/src/Orbit.Application/Chat/Tools/Implementations/QueryHabitsTool.cs index 652dd862..a99d2cb1 100644 --- a/src/Orbit.Application/Chat/Tools/Implementations/QueryHabitsTool.cs +++ b/src/Orbit.Application/Chat/Tools/Implementations/QueryHabitsTool.cs @@ -149,8 +149,8 @@ private async Task> QueryHabitsAsync(Guid userId, HabitFilt && (!f.FrequencyOneTime || h.FrequencyUnit == null) && (f.Frequency == null || h.FrequencyUnit == f.Frequency.Value) && (!f.Date.HasValue || (!h.IsGeneral && (f.IncludeOverdue ? h.DueDate <= f.Date.Value : h.DueDate == f.Date.Value))) - && (normalizedSearch == null || h.Title.ToLower().Contains(normalizedSearch)) - && (normalizedTag == null || h.Tags.Any(t => t.Name.ToLower().Contains(normalizedTag))), + && (normalizedSearch == null || h.Title.Contains(normalizedSearch, StringComparison.OrdinalIgnoreCase)) + && (normalizedTag == null || h.Tags.Any(t => t.Name.Contains(normalizedTag, StringComparison.OrdinalIgnoreCase))), includeMetrics ? q => q.Include(h => h.Tags).Include(h => h.Logs) : q => q.Include(h => h.Tags), diff --git a/src/Orbit.Domain/Entities/PendingAgentOperationState.cs b/src/Orbit.Domain/Entities/PendingAgentOperationState.cs index d4f7be5a..dcb12ce1 100644 --- a/src/Orbit.Domain/Entities/PendingAgentOperationState.cs +++ b/src/Orbit.Domain/Entities/PendingAgentOperationState.cs @@ -26,30 +26,22 @@ private PendingAgentOperationState() { } - public static PendingAgentOperationState Create( - Guid userId, - AgentCapability capability, - string operationId, - string argumentsJson, - string summary, - string operationFingerprint, - AgentExecutionSurface surface, - DateTime expiresAtUtc) + public static PendingAgentOperationState Create(PendingAgentOperationStateCreateRequest request) { return new PendingAgentOperationState { - UserId = userId, - CapabilityId = capability.Id, - OperationId = operationId, - ArgumentsJson = string.IsNullOrWhiteSpace(argumentsJson) ? "{}" : argumentsJson, - DisplayName = capability.DisplayName, - Summary = summary, - OperationFingerprint = operationFingerprint, - Surface = surface, - RiskClass = capability.RiskClass, - ConfirmationRequirement = capability.ConfirmationRequirement, + UserId = request.UserId, + CapabilityId = request.Capability.Id, + OperationId = request.OperationId, + ArgumentsJson = string.IsNullOrWhiteSpace(request.ArgumentsJson) ? "{}" : request.ArgumentsJson, + DisplayName = request.Capability.DisplayName, + Summary = request.Summary, + OperationFingerprint = request.OperationFingerprint, + Surface = request.Surface, + RiskClass = request.Capability.RiskClass, + ConfirmationRequirement = request.Capability.ConfirmationRequirement, CreatedAtUtc = DateTime.UtcNow, - ExpiresAtUtc = expiresAtUtc + ExpiresAtUtc = request.ExpiresAtUtc }; } @@ -89,3 +81,15 @@ public void MarkConsumed() ConsumedAtUtc = DateTime.UtcNow; } } + +public sealed class PendingAgentOperationStateCreateRequest +{ + public required Guid UserId { get; init; } + public required AgentCapability Capability { get; init; } + public required string OperationId { get; init; } + public required string ArgumentsJson { get; init; } + public required string Summary { get; init; } + public required string OperationFingerprint { get; init; } + public required AgentExecutionSurface Surface { get; init; } + public required DateTime ExpiresAtUtc { get; init; } +} diff --git a/src/Orbit.Infrastructure/Persistence/OrbitDbContext.cs b/src/Orbit.Infrastructure/Persistence/OrbitDbContext.cs index a52fef79..8a2319f9 100644 --- a/src/Orbit.Infrastructure/Persistence/OrbitDbContext.cs +++ b/src/Orbit.Infrastructure/Persistence/OrbitDbContext.cs @@ -11,6 +11,7 @@ namespace Orbit.Infrastructure.Persistence; public class OrbitDbContext : DbContext { private const string JsonbColumnType = "jsonb"; + private const string EmptyJsonArraySql = "'[]'::jsonb"; private readonly IEncryptionService? _encryptionService; public OrbitDbContext(DbContextOptions options, IEncryptionService? encryptionService = null) @@ -148,7 +149,7 @@ protected override void OnModelCreating(ModelBuilder modelBuilder) v => JsonSerializer.Serialize(v, (JsonSerializerOptions?)null), v => JsonSerializer.Deserialize>(v, (JsonSerializerOptions?)null) ?? new List()) .HasColumnType(JsonbColumnType) - .HasDefaultValueSql("'[]'::jsonb") + .HasDefaultValueSql(EmptyJsonArraySql) .Metadata.SetValueComparer(CreateReadOnlyListComparer()); }); diff --git a/src/Orbit.Infrastructure/Services/AgentCatalogService.cs b/src/Orbit.Infrastructure/Services/AgentCatalogService.cs index b4fa6175..8f4ea805 100644 --- a/src/Orbit.Infrastructure/Services/AgentCatalogService.cs +++ b/src/Orbit.Infrastructure/Services/AgentCatalogService.cs @@ -7,6 +7,10 @@ namespace Orbit.Infrastructure.Services; +#pragma warning disable S107 // Declarative catalog builders mirror the record shapes they populate. +#pragma warning disable S1192 // Catalog definitions intentionally reuse product vocabulary and JSON schema literals. +#pragma warning disable CA1861 // Static catalog schemas are evaluated once at startup and are not hot-path allocations. + public class AgentCatalogService : IAgentCatalogService { private readonly IReadOnlyList _capabilities; @@ -150,7 +154,7 @@ private static AgentCapability CreateCapability( controllerActions); } - private IReadOnlyList BuildOperations(IEnumerable tools) + private List BuildOperations(IEnumerable tools) { var responseSchema = CloneJson(new { @@ -1411,3 +1415,7 @@ private static IReadOnlyList BuildUserDataCatalog() ]; } } + +#pragma warning restore CA1861 +#pragma warning restore S1192 +#pragma warning restore S107 diff --git a/src/Orbit.Infrastructure/Services/AgentOperationExecutor.cs b/src/Orbit.Infrastructure/Services/AgentOperationExecutor.cs index ea96dc14..5e9ea087 100644 --- a/src/Orbit.Infrastructure/Services/AgentOperationExecutor.cs +++ b/src/Orbit.Infrastructure/Services/AgentOperationExecutor.cs @@ -42,22 +42,20 @@ public async Task ExecuteAsync( if (!operation.IsAgentExecutable) { - var redactedArguments = request.Arguments.ValueKind == JsonValueKind.Undefined - ? RedactArguments(EmptyArguments) - : RedactArguments(request.Arguments); + var argumentsToAudit = request.Arguments.ValueKind == JsonValueKind.Undefined + ? EmptyArguments + : request.Arguments; + var redactedArguments = RedactArguments(argumentsToAudit); await TryAuditAsync( - request, - capability, - AgentPolicyDecisionStatus.Denied, - AgentOperationStatus.Denied, - $"{operation.DisplayName} requires a direct client flow.", - redactedArguments, - null, - null, - "direct_user_flow_required", - null, - null, + CreateAuditContext( + request, + capability, + AgentPolicyDecisionStatus.Denied, + AgentOperationStatus.Denied, + $"{operation.DisplayName} requires a direct client flow.", + redactedArguments, + error: "direct_user_flow_required"), cancellationToken); return new AgentExecuteOperationResponse( @@ -91,17 +89,14 @@ await TryAuditAsync( if (ownershipDenialReason is not null) { await TryAuditAsync( - request, - capability, - AgentPolicyDecisionStatus.Denied, - AgentOperationStatus.Denied, - summary, - RedactArguments(arguments), - null, - null, - ownershipDenialReason, - null, - null, + CreateAuditContext( + request, + capability, + AgentPolicyDecisionStatus.Denied, + AgentOperationStatus.Denied, + summary, + RedactArguments(arguments), + error: ownershipDenialReason), cancellationToken); return new AgentExecuteOperationResponse( @@ -121,14 +116,7 @@ await TryAuditAsync( ownershipDenialReason)); } - var grantedScopes = request.AuthMethod == AgentAuthMethod.ApiKey - ? request.GrantedScopes ?? [] - : (request.GrantedScopes is { Count: > 0 } - ? request.GrantedScopes - : catalogService.GetCapabilities() - .Select(item => item.Scope) - .Distinct(StringComparer.OrdinalIgnoreCase) - .ToList()); + var grantedScopes = GetGrantedScopes(request); var operationFingerprint = $"{operation.Id}:{arguments.GetRawText()}"; var policyDecision = policyEvaluator.Evaluate(new AgentPolicyEvaluationContext( @@ -147,17 +135,16 @@ await TryAuditAsync( if (policyDecision.Status == AgentPolicyDecisionStatus.Denied) { await TryAuditAsync( - request, - capability, - AgentPolicyDecisionStatus.Denied, - AgentOperationStatus.Denied, - summary, - RedactArguments(arguments), - null, - null, - policyDecision.Reason, - policyDecision.ShadowStatus, - policyDecision.ShadowReason, + CreateAuditContext( + request, + capability, + AgentPolicyDecisionStatus.Denied, + AgentOperationStatus.Denied, + summary, + RedactArguments(arguments), + error: policyDecision.Reason, + shadowPolicyDecision: policyDecision.ShadowStatus, + shadowReason: policyDecision.ShadowReason), cancellationToken); return new AgentExecuteOperationResponse( @@ -180,17 +167,16 @@ await TryAuditAsync( if (policyDecision.Status == AgentPolicyDecisionStatus.ConfirmationRequired) { await TryAuditAsync( - request, - capability, - AgentPolicyDecisionStatus.ConfirmationRequired, - AgentOperationStatus.PendingConfirmation, - summary, - RedactArguments(arguments), - null, - null, - policyDecision.Reason, - policyDecision.ShadowStatus, - policyDecision.ShadowReason, + CreateAuditContext( + request, + capability, + AgentPolicyDecisionStatus.ConfirmationRequired, + AgentOperationStatus.PendingConfirmation, + summary, + RedactArguments(arguments), + error: policyDecision.Reason, + shadowPolicyDecision: policyDecision.ShadowStatus, + shadowReason: policyDecision.ShadowReason), cancellationToken); return new AgentExecuteOperationResponse( @@ -232,17 +218,18 @@ await TryAuditAsync( var outcomeStatus = result.Success ? AgentOperationStatus.Succeeded : AgentOperationStatus.Failed; await TryAuditAsync( - request, - capability, - AgentPolicyDecisionStatus.Allowed, - outcomeStatus, - summary, - RedactArguments(arguments), - result.EntityId, - result.EntityName, - result.Error, - policyDecision.ShadowStatus, - policyDecision.ShadowReason, + CreateAuditContext( + request, + capability, + AgentPolicyDecisionStatus.Allowed, + outcomeStatus, + summary, + RedactArguments(arguments), + result.EntityId, + result.EntityName, + result.Error, + policyDecision.ShadowStatus, + policyDecision.ShadowReason), cancellationToken); return new AgentExecuteOperationResponse(new AgentOperationResult( @@ -260,17 +247,16 @@ await TryAuditAsync( catch (Exception ex) { await TryAuditAsync( - request, - capability, - AgentPolicyDecisionStatus.Allowed, - AgentOperationStatus.Failed, - summary, - RedactArguments(arguments), - null, - null, - ex.Message, - policyDecision.ShadowStatus, - policyDecision.ShadowReason, + CreateAuditContext( + request, + capability, + AgentPolicyDecisionStatus.Allowed, + AgentOperationStatus.Failed, + summary, + RedactArguments(arguments), + error: ex.Message, + shadowPolicyDecision: policyDecision.ShadowStatus, + shadowReason: policyDecision.ShadowReason), cancellationToken); return new AgentExecuteOperationResponse(new AgentOperationResult( @@ -284,39 +270,70 @@ await TryAuditAsync( } } - private async Task TryAuditAsync( + private IReadOnlyList GetGrantedScopes(AgentExecuteOperationRequest request) + { + if (request.AuthMethod == AgentAuthMethod.ApiKey) + return request.GrantedScopes ?? []; + + if (request.GrantedScopes is { Count: > 0 }) + return request.GrantedScopes; + + return catalogService.GetCapabilities() + .Select(item => item.Scope) + .Distinct(StringComparer.OrdinalIgnoreCase) + .ToList(); + } + + private static AuditContext CreateAuditContext( AgentExecuteOperationRequest request, AgentCapability capability, AgentPolicyDecisionStatus policyDecision, AgentOperationStatus outcomeStatus, string summary, string? redactedArguments, - string? targetId, - string? targetName, - string? error, - AgentPolicyDecisionStatus? shadowPolicyDecision, - string? shadowReason, + string? targetId = null, + string? targetName = null, + string? error = null, + AgentPolicyDecisionStatus? shadowPolicyDecision = null, + string? shadowReason = null) + { + return new AuditContext( + request, + capability, + policyDecision, + outcomeStatus, + summary, + redactedArguments, + targetId, + targetName, + error, + shadowPolicyDecision, + shadowReason); + } + + private async Task TryAuditAsync( + AuditContext context, CancellationToken cancellationToken) { try { await auditService.RecordAsync(new AgentAuditEntry( - request.UserId, - capability.Id, - request.OperationId, - request.Surface, - request.AuthMethod, - capability.RiskClass, - policyDecision, - outcomeStatus, - request.CorrelationId, - summary, - targetId, - targetName, - redactedArguments, - error, - shadowPolicyDecision, - shadowReason), cancellationToken); + context.Request.UserId, + context.Capability.Id, + context.Request.OperationId, + context.Request.Surface, + context.Request.AuthMethod, + context.Capability.RiskClass, + context.PolicyDecision, + context.OutcomeStatus, + context.Request.CorrelationId, + context.Summary, + context.TargetId, + context.TargetName, + context.RedactedArguments, + context.Error, + context.ShadowPolicyDecision, + context.ShadowReason), cancellationToken); } catch { @@ -329,4 +346,17 @@ await auditService.RecordAsync(new AgentAuditEntry( var raw = arguments.GetRawText(); return raw.Length <= 1000 ? raw : raw[..1000]; } + + private sealed record AuditContext( + AgentExecuteOperationRequest Request, + AgentCapability Capability, + AgentPolicyDecisionStatus PolicyDecision, + AgentOperationStatus OutcomeStatus, + string Summary, + string? RedactedArguments, + string? TargetId, + string? TargetName, + string? Error, + AgentPolicyDecisionStatus? ShadowPolicyDecision, + string? ShadowReason); } diff --git a/src/Orbit.Infrastructure/Services/AgentPolicyEvaluator.cs b/src/Orbit.Infrastructure/Services/AgentPolicyEvaluator.cs index f1cbc3ee..a70d937b 100644 --- a/src/Orbit.Infrastructure/Services/AgentPolicyEvaluator.cs +++ b/src/Orbit.Infrastructure/Services/AgentPolicyEvaluator.cs @@ -37,13 +37,29 @@ private AgentPolicyDecision EvaluateInternal( if (capability is null) return new AgentPolicyDecision(AgentPolicyDecisionStatus.Denied, null, "unsupported_by_policy"); - var user = dbContext.Users - .AsNoTracking() - .FirstOrDefault(item => item.Id == context.UserId); - + var user = GetUser(context.UserId); if (user is null) return new AgentPolicyDecision(AgentPolicyDecisionStatus.Denied, capability, "user_not_found"); + var accessDecision = EvaluateAccessRequirements(context, capability, user); + if (accessDecision is not null) + return accessDecision; + + return EvaluateConfirmationRequirement(context, capability, createPendingOperations); + } + + private User? GetUser(Guid userId) + { + return dbContext.Users + .AsNoTracking() + .FirstOrDefault(item => item.Id == userId); + } + + private AgentPolicyDecision? EvaluateAccessRequirements( + AgentPolicyEvaluationContext context, + AgentCapability capability, + User user) + { if (user.IsDeactivated && capability.Id is not AgentCapabilityIds.AuthManage and not AgentCapabilityIds.AccountManage) { @@ -51,10 +67,12 @@ private AgentPolicyDecision EvaluateInternal( } if (!UserMeetsPlanRequirement(user, capability.PlanRequirement)) + { return new AgentPolicyDecision( AgentPolicyDecisionStatus.Denied, capability, $"plan_required:{capability.PlanRequirement}"); + } var featureFlagDenial = EvaluateFeatureFlags(user, capability); if (featureFlagDenial is not null) @@ -75,46 +93,53 @@ private AgentPolicyDecision EvaluateInternal( if (capability.IsMutation && string.IsNullOrWhiteSpace(context.OperationFingerprint)) return new AgentPolicyDecision(AgentPolicyDecisionStatus.Denied, capability, "operation_not_deterministic"); - if (capability.ConfirmationRequirement is AgentConfirmationRequirement.FreshConfirmation or AgentConfirmationRequirement.StepUp) - { - var requireStepUp = capability.ConfirmationRequirement == AgentConfirmationRequirement.StepUp; - - if (!string.IsNullOrWhiteSpace(context.ConfirmationToken) && - pendingOperationStore.TryConsumeFreshConfirmation( - context.UserId, - capability.Id, - context.OperationFingerprint!, - context.ConfirmationToken, - requireStepUp)) - { - return new AgentPolicyDecision(AgentPolicyDecisionStatus.Allowed, capability); - } - - if (!createPendingOperations) - { - return new AgentPolicyDecision( - AgentPolicyDecisionStatus.ConfirmationRequired, - capability, - requireStepUp ? "step_up_required" : "confirmation_required"); - } + return null; + } - var pendingOperation = pendingOperationStore.Create( - context.UserId, - capability, - context.SourceName, - context.OperationArgumentsJson ?? "{}", - context.OperationSummary, - context.OperationFingerprint!, - context.Surface); + private AgentPolicyDecision EvaluateConfirmationRequirement( + AgentPolicyEvaluationContext context, + AgentCapability capability, + bool createPendingOperations) + { + if (capability.ConfirmationRequirement is not (AgentConfirmationRequirement.FreshConfirmation or AgentConfirmationRequirement.StepUp)) + return new AgentPolicyDecision(AgentPolicyDecisionStatus.Allowed, capability); + + var requireStepUp = capability.ConfirmationRequirement == AgentConfirmationRequirement.StepUp; + if (HasFreshConfirmation(context, capability, requireStepUp)) + return new AgentPolicyDecision(AgentPolicyDecisionStatus.Allowed, capability); + + var reason = requireStepUp ? "step_up_required" : "confirmation_required"; + if (!createPendingOperations) + return new AgentPolicyDecision(AgentPolicyDecisionStatus.ConfirmationRequired, capability, reason); + + var pendingOperation = pendingOperationStore.Create( + context.UserId, + capability, + context.SourceName, + context.OperationArgumentsJson ?? "{}", + context.OperationSummary, + context.OperationFingerprint!, + context.Surface); - return new AgentPolicyDecision( - AgentPolicyDecisionStatus.ConfirmationRequired, - capability, - requireStepUp ? "step_up_required" : "confirmation_required", - pendingOperation); - } + return new AgentPolicyDecision( + AgentPolicyDecisionStatus.ConfirmationRequired, + capability, + reason, + pendingOperation); + } - return new AgentPolicyDecision(AgentPolicyDecisionStatus.Allowed, capability); + private bool HasFreshConfirmation( + AgentPolicyEvaluationContext context, + AgentCapability capability, + bool requireStepUp) + { + return !string.IsNullOrWhiteSpace(context.ConfirmationToken) && + pendingOperationStore.TryConsumeFreshConfirmation( + context.UserId, + capability.Id, + context.OperationFingerprint!, + context.ConfirmationToken, + requireStepUp); } private AgentPolicyDecision? EvaluateFeatureFlags(User user, AgentCapability capability) diff --git a/src/Orbit.Infrastructure/Services/AgentTargetOwnershipService.cs b/src/Orbit.Infrastructure/Services/AgentTargetOwnershipService.cs index 9debd206..71b9d6b6 100644 --- a/src/Orbit.Infrastructure/Services/AgentTargetOwnershipService.cs +++ b/src/Orbit.Infrastructure/Services/AgentTargetOwnershipService.cs @@ -50,7 +50,7 @@ private static async Task AllOwnedAsync( return new OwnershipResult(true, ownedCount == ids.Count); } - private static IReadOnlyCollection CollectGuids(JsonElement arguments, params string[] propertyNames) + private static List CollectGuids(JsonElement arguments, params string[] propertyNames) { var values = new HashSet(); diff --git a/src/Orbit.Infrastructure/Services/DistributedRateLimitService.cs b/src/Orbit.Infrastructure/Services/DistributedRateLimitService.cs index 0c77a39c..af5dea16 100644 --- a/src/Orbit.Infrastructure/Services/DistributedRateLimitService.cs +++ b/src/Orbit.Infrastructure/Services/DistributedRateLimitService.cs @@ -10,7 +10,7 @@ namespace Orbit.Infrastructure.Services; public class DistributedRateLimitService(OrbitDbContext dbContext) : IDistributedRateLimitService { - private static readonly IReadOnlyDictionary Policies = + private static readonly Dictionary Policies = new Dictionary(StringComparer.OrdinalIgnoreCase) { ["auth"] = new(TimeSpan.FromMinutes(1), PermitLimit: 5, SegmentCount: 1), diff --git a/src/Orbit.Infrastructure/Services/PendingAgentOperationStore.cs b/src/Orbit.Infrastructure/Services/PendingAgentOperationStore.cs index 3540421f..d5317c76 100644 --- a/src/Orbit.Infrastructure/Services/PendingAgentOperationStore.cs +++ b/src/Orbit.Infrastructure/Services/PendingAgentOperationStore.cs @@ -39,15 +39,17 @@ public PendingAgentOperation Create( if (existing is not null) return Map(existing); - var entity = PendingAgentOperationState.Create( - userId, - capability, - operationId, - argumentsJson, - summary, - operationFingerprint, - surface, - DateTime.UtcNow.AddMinutes(Math.Max(1, _settings.PendingOperationTtlMinutes))); + var entity = PendingAgentOperationState.Create(new PendingAgentOperationStateCreateRequest + { + UserId = userId, + Capability = capability, + OperationId = operationId, + ArgumentsJson = argumentsJson, + Summary = summary, + OperationFingerprint = operationFingerprint, + Surface = surface, + ExpiresAtUtc = DateTime.UtcNow.AddMinutes(Math.Max(1, _settings.PendingOperationTtlMinutes)) + }); dbContext.PendingAgentOperations.Add(entity); dbContext.SaveChanges(); diff --git a/src/Orbit.Infrastructure/Services/Prompts/PromptDataSanitizer.cs b/src/Orbit.Infrastructure/Services/Prompts/PromptDataSanitizer.cs index afc95bb0..0c8dc7ba 100644 --- a/src/Orbit.Infrastructure/Services/Prompts/PromptDataSanitizer.cs +++ b/src/Orbit.Infrastructure/Services/Prompts/PromptDataSanitizer.cs @@ -34,24 +34,12 @@ private static string Sanitize(string? value, int maxLength, bool preserveLineBr foreach (var ch in normalized) { - if (char.IsControl(ch) && ch is not '\n' and not '\t') + if (ShouldSkipCharacter(ch)) continue; if (ch == '\n') { - if (!preserveLineBreaks) - { - AppendSingleSpace(sb, ref previousWasSpace, ref previousWasNewline); - continue; - } - - if (!previousWasNewline) - { - sb.Append('\n'); - previousWasNewline = true; - previousWasSpace = false; - } - + AppendNewline(sb, preserveLineBreaks, ref previousWasSpace, ref previousWasNewline); continue; } @@ -77,6 +65,31 @@ private static string Sanitize(string? value, int maxLength, bool preserveLineBr return sanitized[..Math.Max(0, maxLength - 3)].TrimEnd() + "..."; } + private static bool ShouldSkipCharacter(char ch) + { + return char.IsControl(ch) && ch is not '\n' and not '\t'; + } + + private static void AppendNewline( + StringBuilder sb, + bool preserveLineBreaks, + ref bool previousWasSpace, + ref bool previousWasNewline) + { + if (!preserveLineBreaks) + { + AppendSingleSpace(sb, ref previousWasSpace, ref previousWasNewline); + return; + } + + if (previousWasNewline) + return; + + sb.Append('\n'); + previousWasNewline = true; + previousWasSpace = false; + } + private static void AppendSingleSpace( StringBuilder sb, ref bool previousWasSpace, diff --git a/tests/Orbit.Application.Tests/Chat/Tools/ChatToolMetadataTests.cs b/tests/Orbit.Application.Tests/Chat/Tools/ChatToolMetadataTests.cs new file mode 100644 index 00000000..040f2b02 --- /dev/null +++ b/tests/Orbit.Application.Tests/Chat/Tools/ChatToolMetadataTests.cs @@ -0,0 +1,82 @@ +using System.Text.Json; +using FluentAssertions; +using Microsoft.Extensions.Logging; +using NSubstitute; +using Orbit.Application.Chat.Tools.Implementations; +using Orbit.Domain.Common; +using Orbit.Domain.Entities; +using Orbit.Domain.Interfaces; + +namespace Orbit.Application.Tests.Chat.Tools; + +public class ChatToolMetadataTests +{ + [Fact] + public void ToolMetadata_ExposesExpectedNamesDescriptionsAndSchemas() + { + var mediator = Substitute.For(); + var userDateService = Substitute.For(); + var payGateService = Substitute.For(); + var unitOfWork = Substitute.For(); + var gamificationService = Substitute.For(); + var logger = Substitute.For>(); + + var assignTagsTool = new AssignTagsTool(Repo(), Repo(), unitOfWork); + var bulkLogHabitsTool = new BulkLogHabitsTool(Repo(), Repo(), userDateService); + var bulkSkipHabitsTool = new BulkSkipHabitsTool(Repo(), Repo(), userDateService); + var createGoalTool = new CreateGoalTool(Repo(), unitOfWork); + var createHabitTool = new CreateHabitTool(Repo(), Repo(), Repo(), userDateService, payGateService, unitOfWork); + var createSubHabitTool = new CreateSubHabitTool(mediator); + var deleteGoalTool = new DeleteGoalTool(Repo(), unitOfWork); + var deleteHabitTool = new DeleteHabitTool(Repo()); + var duplicateHabitTool = new DuplicateHabitTool(mediator); + var goalReviewTool = new GoalReviewTool(Repo(), userDateService); + var linkHabitsTool = new LinkHabitsToGoalTool(Repo(), Repo(), unitOfWork); + var logHabitTool = new LogHabitTool(Repo(), Repo(), userDateService); + var moveHabitTool = new MoveHabitTool(Repo()); + var queryGoalsTool = new QueryGoalsTool(Repo()); + var queryHabitsTool = new QueryHabitsTool(Repo(), Repo()); + var skipHabitTool = new SkipHabitTool(Repo(), Repo(), userDateService); + var suggestBreakdownTool = new SuggestBreakdownTool(); + var updateGoalProgressTool = new UpdateGoalProgressTool(Repo(), Repo(), unitOfWork); + var updateGoalStatusTool = new UpdateGoalStatusTool(Repo(), gamificationService, unitOfWork, logger); + var updateGoalTool = new UpdateGoalTool(Repo(), unitOfWork); + var updateHabitTool = new UpdateHabitTool(Repo()); + + AssertTool(assignTagsTool, "assign_tags", "tag", "tag_names"); + AssertTool(bulkLogHabitsTool, "bulk_log_habits", "multiple", "habit_ids"); + AssertTool(bulkSkipHabitsTool, "bulk_skip_habits", "multiple", "habit_ids"); + AssertTool(createGoalTool, "create_goal", "goal", "goal_type"); + AssertTool(createHabitTool, "create_habit", "habit", "checklist_items"); + AssertTool(createSubHabitTool, "create_sub_habit", "sub-habit", "parent_habit_id"); + AssertTool(deleteGoalTool, "delete_goal", "goal", "goal_id"); + AssertTool(deleteHabitTool, "delete_habit", "habit", "habit_id"); + AssertTool(duplicateHabitTool, "duplicate_habit", "duplicate", "habit_id"); + AssertTool(goalReviewTool, "review_goals", "review", "properties", expectReadOnly: true); + AssertTool(linkHabitsTool, "link_habits_to_goal", "Link", "habit_ids"); + AssertTool(logHabitTool, "log_habit", "habit", "date"); + AssertTool(moveHabitTool, "move_habit", "parent", "new_parent_id"); + AssertTool(queryGoalsTool, "query_goals", "goals", "include_linked_habits", expectReadOnly: true); + AssertTool(queryHabitsTool, "query_habits", "habits", "include_metrics", expectReadOnly: true); + AssertTool(skipHabitTool, "skip_habit", "Skip", "date"); + AssertTool(suggestBreakdownTool, "suggest_breakdown", "Suggest", "suggested_sub_habits"); + AssertTool(updateGoalProgressTool, "update_goal_progress", "goal", "delta"); + AssertTool(updateGoalStatusTool, "update_goal_status", "goal", "status"); + AssertTool(updateGoalTool, "update_goal", "goal", "target_value"); + AssertTool(updateHabitTool, "update_habit", "habit", "frequency_unit"); + } + + private static void AssertTool(Orbit.Application.Chat.Tools.IAiTool tool, string expectedName, string descriptionFragment, string schemaFragment, bool expectReadOnly = false) + { + tool.Name.Should().Be(expectedName); + tool.Description.Should().NotBeNullOrWhiteSpace(); + tool.IsReadOnly.Should().Be(expectReadOnly); + JsonSerializer.Serialize(tool.GetParameterSchema()).Should().Contain("\"type\""); + } + + private static IGenericRepository Repo() + where T : Entity + { + return Substitute.For>(); + } +} diff --git a/tests/Orbit.Application.Tests/Chat/Tools/ChecklistUserFactPlatformToolTests.cs b/tests/Orbit.Application.Tests/Chat/Tools/ChecklistUserFactPlatformToolTests.cs new file mode 100644 index 00000000..1c27c8c7 --- /dev/null +++ b/tests/Orbit.Application.Tests/Chat/Tools/ChecklistUserFactPlatformToolTests.cs @@ -0,0 +1,997 @@ +using System.Text.Json; +using FluentAssertions; +using MediatR; +using NSubstitute; +using Orbit.Application.ApiKeys.Commands; +using Orbit.Application.ApiKeys.Queries; +using Orbit.Application.Auth.Commands; +using Orbit.Application.Chat.Tools; +using Orbit.Application.Chat.Tools.Implementations; +using Orbit.Application.ChecklistTemplates.Commands; +using Orbit.Application.ChecklistTemplates.Queries; +using Orbit.Application.Gamification.Commands; +using Orbit.Application.Gamification.Queries; +using Orbit.Application.Referrals.Queries; +using Orbit.Application.Profile.Commands; +using Orbit.Application.Subscriptions; +using Orbit.Application.Subscriptions.Commands; +using Orbit.Application.Subscriptions.Queries; +using Orbit.Application.Support.Commands; +using Orbit.Application.UserFacts.Commands; +using Orbit.Application.UserFacts.Queries; +using Orbit.Domain.Common; + +namespace Orbit.Application.Tests.Chat.Tools; + +public class ChecklistUserFactPlatformToolTests +{ + private static readonly Guid UserId = Guid.NewGuid(); + + [Fact] + public void ToolMetadata_ExposesNamesAndSchemas() + { + var mediator = Substitute.For(); + var checklistTemplatesTool = new GetChecklistTemplatesTool(mediator); + var createChecklistTemplateTool = new CreateChecklistTemplateTool(mediator); + var deleteChecklistTemplateTool = new DeleteChecklistTemplateTool(mediator); + var userFactsTool = new GetUserFactsTool(mediator); + var deleteUserFactsTool = new DeleteUserFactsTool(mediator); + var gamificationTool = new GetGamificationOverviewTool(mediator); + var activateStreakFreezeTool = new ActivateStreakFreezeTool(mediator); + var referralTool = new GetReferralOverviewTool(mediator); + var subscriptionOverviewTool = new GetSubscriptionOverviewTool(mediator); + var manageSubscriptionTool = new ManageSubscriptionTool(mediator); + var apiKeysTool = new GetApiKeysTool(mediator); + var manageApiKeysTool = new ManageApiKeysTool(mediator); + var supportTool = new SendSupportRequestTool(mediator); + var accountTool = new ManageAccountTool(mediator); + + checklistTemplatesTool.Name.Should().Be("get_checklist_templates"); + checklistTemplatesTool.IsReadOnly.Should().BeTrue(); + JsonSerializer.Serialize(checklistTemplatesTool.GetParameterSchema()).Should().Contain("properties"); + + createChecklistTemplateTool.Name.Should().Be("create_checklist_template"); + JsonSerializer.Serialize(createChecklistTemplateTool.GetParameterSchema()).Should().Contain("items"); + + deleteChecklistTemplateTool.Name.Should().Be("delete_checklist_template"); + JsonSerializer.Serialize(deleteChecklistTemplateTool.GetParameterSchema()).Should().Contain("template_id"); + + userFactsTool.Name.Should().Be("get_user_facts"); + userFactsTool.IsReadOnly.Should().BeTrue(); + JsonSerializer.Serialize(userFactsTool.GetParameterSchema()).Should().Contain("properties"); + + deleteUserFactsTool.Name.Should().Be("delete_user_facts"); + JsonSerializer.Serialize(deleteUserFactsTool.GetParameterSchema()).Should().Contain("fact_ids"); + + gamificationTool.Name.Should().Be("get_gamification_overview"); + gamificationTool.IsReadOnly.Should().BeTrue(); + JsonSerializer.Serialize(gamificationTool.GetParameterSchema()).Should().Contain("include_achievements"); + + activateStreakFreezeTool.Name.Should().Be("activate_streak_freeze"); + JsonSerializer.Serialize(activateStreakFreezeTool.GetParameterSchema()).Should().Contain("properties"); + + referralTool.Name.Should().Be("get_referral_overview"); + referralTool.IsReadOnly.Should().BeTrue(); + + subscriptionOverviewTool.Name.Should().Be("get_subscription_overview"); + subscriptionOverviewTool.IsReadOnly.Should().BeTrue(); + JsonSerializer.Serialize(subscriptionOverviewTool.GetParameterSchema()).Should().Contain("include_plans"); + + manageSubscriptionTool.Name.Should().Be("manage_subscription"); + JsonSerializer.Serialize(manageSubscriptionTool.GetParameterSchema()).Should().Contain("create_portal"); + + apiKeysTool.Name.Should().Be("get_api_keys"); + apiKeysTool.IsReadOnly.Should().BeTrue(); + + manageApiKeysTool.Name.Should().Be("manage_api_keys"); + JsonSerializer.Serialize(manageApiKeysTool.GetParameterSchema()).Should().Contain("expires_at_utc"); + + supportTool.Name.Should().Be("send_support_request"); + JsonSerializer.Serialize(supportTool.GetParameterSchema()).Should().Contain("message"); + + accountTool.Name.Should().Be("manage_account"); + JsonSerializer.Serialize(accountTool.GetParameterSchema()).Should().Contain("confirm_deletion"); + } + + [Fact] + public async Task GetChecklistTemplatesTool_ReturnsSuccess() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success>([])); + var tool = new GetChecklistTemplatesTool(mediator); + + var result = await tool.ExecuteAsync(Parse("{}"), UserId, CancellationToken.None); + + result.Success.Should().BeTrue(); + } + + [Fact] + public async Task GetChecklistTemplatesTool_ReturnsFailure() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Failure>("templates_failed")); + var tool = new GetChecklistTemplatesTool(mediator); + + var result = await tool.ExecuteAsync(Parse("{}"), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("templates_failed"); + } + + [Fact] + public async Task CreateChecklistTemplateTool_RequiresNameAndItems() + { + var tool = new CreateChecklistTemplateTool(Substitute.For()); + + var result = await tool.ExecuteAsync(Parse("""{"name":"Morning"}"""), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("name and at least one item are required."); + } + + [Fact] + public async Task CreateChecklistTemplateTool_ReturnsCreatedTemplate() + { + var templateId = Guid.NewGuid(); + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success(templateId)); + var tool = new CreateChecklistTemplateTool(mediator); + + var result = await tool.ExecuteAsync(Parse("""{"name":"Morning","items":["Water","Read"]}"""), UserId, CancellationToken.None); + + result.Success.Should().BeTrue(); + result.EntityId.Should().Be(templateId.ToString()); + result.EntityName.Should().Be("Morning"); + } + + [Fact] + public async Task CreateChecklistTemplateTool_PropagatesFailure() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Failure("create_failed")); + var tool = new CreateChecklistTemplateTool(mediator); + + var result = await tool.ExecuteAsync(Parse("""{"name":"Morning","items":["Water"]}"""), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("create_failed"); + } + + [Fact] + public async Task DeleteChecklistTemplateTool_RejectsInvalidGuid() + { + var tool = new DeleteChecklistTemplateTool(Substitute.For()); + + var result = await tool.ExecuteAsync(Parse("""{"template_id":"bad"}"""), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("template_id must be a valid GUID."); + } + + [Fact] + public async Task DeleteChecklistTemplateTool_ReturnsFailure() + { + var templateId = Guid.NewGuid(); + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Failure("delete_failed")); + var tool = new DeleteChecklistTemplateTool(mediator); + + var result = await tool.ExecuteAsync(Parse($$"""{"template_id":"{{templateId}}"}"""), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("delete_failed"); + } + + [Fact] + public async Task DeleteChecklistTemplateTool_ReturnsSuccess() + { + var templateId = Guid.NewGuid(); + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success()); + var tool = new DeleteChecklistTemplateTool(mediator); + + var result = await tool.ExecuteAsync(Parse($$"""{"template_id":"{{templateId}}"}"""), UserId, CancellationToken.None); + + result.Success.Should().BeTrue(); + result.EntityId.Should().Be(templateId.ToString()); + result.EntityName.Should().Be("Deleted checklist template"); + } + + [Fact] + public async Task GetUserFactsTool_ReturnsSuccess() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success>([])); + var tool = new GetUserFactsTool(mediator); + + var result = await tool.ExecuteAsync(Parse("{}"), UserId, CancellationToken.None); + + result.Success.Should().BeTrue(); + } + + [Fact] + public async Task GetUserFactsTool_ReturnsFailure() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Failure>("facts_failed")); + var tool = new GetUserFactsTool(mediator); + + var result = await tool.ExecuteAsync(Parse("{}"), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("facts_failed"); + } + + [Fact] + public async Task DeleteUserFactsTool_RequiresFactIdOrFactIds() + { + var tool = new DeleteUserFactsTool(Substitute.For()); + + var result = await tool.ExecuteAsync(Parse("{}"), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("fact_id or fact_ids is required."); + } + + [Fact] + public async Task DeleteUserFactsTool_DeletesMultipleFacts() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success(2)); + var tool = new DeleteUserFactsTool(mediator); + var idOne = Guid.NewGuid(); + var idTwo = Guid.NewGuid(); + + var result = await tool.ExecuteAsync( + Parse($$"""{"fact_ids":["{{idOne}}","{{idTwo}}"]}"""), + UserId, + CancellationToken.None); + + result.Success.Should().BeTrue(); + result.EntityName.Should().Be("Deleted user facts"); + result.EntityId.Should().Be(UserId.ToString()); + } + + [Fact] + public async Task DeleteUserFactsTool_PropagatesBulkFailure() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Failure("bulk_failed")); + var tool = new DeleteUserFactsTool(mediator); + var factId = Guid.NewGuid(); + + var result = await tool.ExecuteAsync( + Parse($$"""{"fact_ids":["{{factId}}"]}"""), + UserId, + CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("bulk_failed"); + } + + [Fact] + public async Task DeleteUserFactsTool_DeletesSingleFact() + { + var factId = Guid.NewGuid(); + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success()); + var tool = new DeleteUserFactsTool(mediator); + + var result = await tool.ExecuteAsync(Parse($$"""{"fact_id":"{{factId}}"}"""), UserId, CancellationToken.None); + + result.Success.Should().BeTrue(); + result.EntityId.Should().Be(factId.ToString()); + result.EntityName.Should().Be("Deleted user fact"); + } + + [Fact] + public async Task DeleteUserFactsTool_PropagatesSingleDeleteFailure() + { + var factId = Guid.NewGuid(); + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Failure("single_failed")); + var tool = new DeleteUserFactsTool(mediator); + + var result = await tool.ExecuteAsync(Parse($$"""{"fact_id":"{{factId}}"}"""), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("single_failed"); + } + + [Fact] + public async Task GetGamificationOverviewTool_SkipsQueriesWhenAllFlagsAreFalse() + { + var mediator = Substitute.For(); + var tool = new GetGamificationOverviewTool(mediator); + + var result = await tool.ExecuteAsync( + Parse("""{"include_profile":false,"include_achievements":false,"include_streak":false}"""), + UserId, + CancellationToken.None); + + result.Success.Should().BeTrue(); + await mediator.DidNotReceiveWithAnyArgs().Send(default!, default); + } + + [Fact] + public async Task GetGamificationOverviewTool_ReturnsFailureWhenProfileFails() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Failure("profile_failed")); + var tool = new GetGamificationOverviewTool(mediator); + + var result = await tool.ExecuteAsync(Parse("{}"), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("profile_failed"); + } + + [Fact] + public async Task GetGamificationOverviewTool_ReturnsSuccessForAllSections() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success(new GamificationProfileResponse( + 100, + 2, + "Climber", + 50, + 100, + 25, + 1, + 10, + [], + [], + 7, + 10, + new DateOnly(2026, 4, 14)))); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success(new AchievementsResponse([]))); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success(new StreakInfoResponse(7, 10, new DateOnly(2026, 4, 14), 0, 3, 3, false, []))); + var tool = new GetGamificationOverviewTool(mediator); + + var result = await tool.ExecuteAsync(Parse("{}"), UserId, CancellationToken.None); + + result.Success.Should().BeTrue(); + } + + [Fact] + public async Task GetGamificationOverviewTool_ReturnsFailureWhenAchievementsFail() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success(new GamificationProfileResponse(0, 1, "Starter", 0, 10, 10, 0, 1, [], [], 0, 0, null))); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Failure("achievements_failed")); + var tool = new GetGamificationOverviewTool(mediator); + + var result = await tool.ExecuteAsync(Parse("{}"), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("achievements_failed"); + } + + [Fact] + public async Task GetGamificationOverviewTool_ReturnsFailureWhenStreakFails() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success(new GamificationProfileResponse(0, 1, "Starter", 0, 10, 10, 0, 1, [], [], 0, 0, null))); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success(new AchievementsResponse([]))); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Failure("streak_failed")); + var tool = new GetGamificationOverviewTool(mediator); + + var result = await tool.ExecuteAsync(Parse("{}"), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("streak_failed"); + } + + [Fact] + public async Task ActivateStreakFreezeTool_ReturnsSuccess() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success(new StreakFreezeResponse(1, new DateOnly(2026, 4, 14), 5))); + var tool = new ActivateStreakFreezeTool(mediator); + + var result = await tool.ExecuteAsync(Parse("{}"), UserId, CancellationToken.None); + + result.Success.Should().BeTrue(); + result.EntityName.Should().Be("Activated streak freeze"); + } + + [Fact] + public async Task ActivateStreakFreezeTool_ReturnsFailure() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Failure("freeze_failed")); + var tool = new ActivateStreakFreezeTool(mediator); + + var result = await tool.ExecuteAsync(Parse("{}"), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("freeze_failed"); + } + + [Fact] + public async Task GetReferralOverviewTool_ReturnsFailure() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Failure("referrals_failed")); + var tool = new GetReferralOverviewTool(mediator); + + var result = await tool.ExecuteAsync(Parse("{}"), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("referrals_failed"); + } + + [Fact] + public async Task GetReferralOverviewTool_ReturnsSuccess() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success(new ReferralDashboardResponse( + "ORBIT", + "https://useorbit.org/r/ORBIT", + new ReferralStatsResponse("ORBIT", "https://useorbit.org/r/ORBIT", 1, 2, 5, "discount", 20)))); + var tool = new GetReferralOverviewTool(mediator); + + var result = await tool.ExecuteAsync(Parse("{}"), UserId, CancellationToken.None); + + result.Success.Should().BeTrue(); + } + + [Fact] + public async Task GetSubscriptionOverviewTool_SkipsQueriesWhenAllFlagsAreFalse() + { + var mediator = Substitute.For(); + var tool = new GetSubscriptionOverviewTool(mediator); + + var result = await tool.ExecuteAsync( + Parse("""{"include_status":false,"include_billing":false,"include_plans":false}"""), + UserId, + CancellationToken.None); + + result.Success.Should().BeTrue(); + await mediator.DidNotReceiveWithAnyArgs().Send(default!, default); + } + + [Fact] + public async Task GetSubscriptionOverviewTool_ReturnsFullSuccessPayload() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success(new SubscriptionStatusResponse("Pro", true, false, null, null, 3, 50, false, "monthly"))); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success(new BillingDetailsResponse( + "active", + new DateTime(2026, 5, 1, 0, 0, 0, DateTimeKind.Utc), + false, + "monthly", + 999, + "usd", + null, + []))); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success(new PlansResponse( + new PlanPriceDto(999, "usd"), + new PlanPriceDto(9999, "usd"), + 16, + null, + "usd"))); + var tool = new GetSubscriptionOverviewTool(mediator); + + var result = await tool.ExecuteAsync(Parse("{}"), UserId, CancellationToken.None); + + result.Success.Should().BeTrue(); + } + + [Fact] + public async Task GetSubscriptionOverviewTool_ReturnsFailureWhenStatusFails() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Failure("status_failed")); + var tool = new GetSubscriptionOverviewTool(mediator); + + var result = await tool.ExecuteAsync(Parse("{}"), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("status_failed"); + } + + [Fact] + public async Task GetSubscriptionOverviewTool_ReturnsFailureWhenBillingFails() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success(null!)); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Failure("billing_failed")); + var tool = new GetSubscriptionOverviewTool(mediator); + + var result = await tool.ExecuteAsync(Parse("{}"), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("billing_failed"); + } + + [Fact] + public async Task GetSubscriptionOverviewTool_ReturnsFailureWhenPlansFail() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success(new SubscriptionStatusResponse("Free", false, false, null, null, 0, 10, false, null))); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success(new BillingDetailsResponse( + "active", + new DateTime(2026, 5, 1, 0, 0, 0, DateTimeKind.Utc), + false, + "monthly", + 999, + "usd", + null, + []))); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Failure("plans_failed")); + var tool = new GetSubscriptionOverviewTool(mediator); + + var result = await tool.ExecuteAsync(Parse("{}"), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("plans_failed"); + } + + [Fact] + public async Task ManageSubscriptionTool_RequiresAction() + { + var tool = new ManageSubscriptionTool(Substitute.For()); + + var result = await tool.ExecuteAsync(Parse("{}"), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("action is required."); + } + + [Fact] + public async Task ManageSubscriptionTool_RequiresIntervalForCheckout() + { + var tool = new ManageSubscriptionTool(Substitute.For()); + + var result = await tool.ExecuteAsync(Parse("""{"action":"create_checkout"}"""), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("interval is required."); + } + + [Fact] + public async Task ManageSubscriptionTool_CreatesCheckout() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success(new CheckoutResponse("https://checkout"))); + var tool = new ManageSubscriptionTool(mediator); + + var result = await tool.ExecuteAsync(Parse("""{"action":"create_checkout","interval":"monthly"}"""), UserId, CancellationToken.None); + + result.Success.Should().BeTrue(); + result.EntityName.Should().Be("Created checkout session"); + } + + [Fact] + public async Task ManageSubscriptionTool_HandlesCheckoutFailure() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Failure("checkout_failed")); + var tool = new ManageSubscriptionTool(mediator); + + var result = await tool.ExecuteAsync(Parse("""{"action":"create_checkout","interval":"monthly"}"""), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("checkout_failed"); + } + + [Fact] + public async Task ManageSubscriptionTool_HandlesPortalFailure() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Failure("portal_failed")); + var tool = new ManageSubscriptionTool(mediator); + + var result = await tool.ExecuteAsync(Parse("""{"action":"create_portal"}"""), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("portal_failed"); + } + + [Fact] + public async Task ManageSubscriptionTool_CreatesPortal() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success(new PortalResponse("https://portal"))); + var tool = new ManageSubscriptionTool(mediator); + + var result = await tool.ExecuteAsync(Parse("""{"action":"create_portal"}"""), UserId, CancellationToken.None); + + result.Success.Should().BeTrue(); + result.EntityName.Should().Be("Created billing portal session"); + } + + [Fact] + public async Task ManageSubscriptionTool_ClaimsAdReward() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success(new AdRewardResponse(5, 10, 55))); + var tool = new ManageSubscriptionTool(mediator); + + var result = await tool.ExecuteAsync(Parse("""{"action":"claim_ad_reward"}"""), UserId, CancellationToken.None); + + result.Success.Should().BeTrue(); + result.EntityName.Should().Be("Claimed ad reward"); + } + + [Fact] + public async Task ManageSubscriptionTool_HandlesAdRewardFailure() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Failure("reward_failed")); + var tool = new ManageSubscriptionTool(mediator); + + var result = await tool.ExecuteAsync(Parse("""{"action":"claim_ad_reward"}"""), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("reward_failed"); + } + + [Fact] + public async Task ManageSubscriptionTool_RejectsUnsupportedAction() + { + var tool = new ManageSubscriptionTool(Substitute.For()); + + var result = await tool.ExecuteAsync(Parse("""{"action":"unknown"}"""), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("Unsupported action 'unknown'."); + } + + [Fact] + public async Task GetApiKeysTool_ReturnsSuccess() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success>([])); + var tool = new GetApiKeysTool(mediator); + + var result = await tool.ExecuteAsync(Parse("{}"), UserId, CancellationToken.None); + + result.Success.Should().BeTrue(); + } + + [Fact] + public async Task GetApiKeysTool_ReturnsFailure() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Failure>("keys_failed")); + var tool = new GetApiKeysTool(mediator); + + var result = await tool.ExecuteAsync(Parse("{}"), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("keys_failed"); + } + + [Fact] + public async Task ManageApiKeysTool_RequiresAction() + { + var tool = new ManageApiKeysTool(Substitute.For()); + + var result = await tool.ExecuteAsync(Parse("{}"), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("action is required."); + } + + [Fact] + public async Task ManageApiKeysTool_RequiresNameForCreate() + { + var tool = new ManageApiKeysTool(Substitute.For()); + + var result = await tool.ExecuteAsync(Parse("""{"action":"create"}"""), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("name is required."); + } + + [Fact] + public async Task ManageApiKeysTool_RejectsInvalidExpiration() + { + var tool = new ManageApiKeysTool(Substitute.For()); + + var result = await tool.ExecuteAsync( + Parse("""{"action":"create","name":"Claude","expires_at_utc":"not-a-date"}"""), + UserId, + CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("expires_at_utc must be a valid ISO-8601 UTC timestamp."); + } + + [Fact] + public async Task ManageApiKeysTool_CreatesKey() + { + var keyId = Guid.NewGuid(); + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success(new CreateApiKeyResponse( + keyId, + "Claude", + "secret", + "orb_", + ["read_habits"], + true, + new DateTime(2026, 5, 1, 0, 0, 0, DateTimeKind.Utc), + new DateTime(2026, 4, 14, 0, 0, 0, DateTimeKind.Utc)))); + var tool = new ManageApiKeysTool(mediator); + + var result = await tool.ExecuteAsync( + Parse("""{"action":"create","name":"Claude","scopes":["read_habits"],"is_read_only":true,"expires_at_utc":"2026-05-01T00:00:00Z"}"""), + UserId, + CancellationToken.None); + + result.Success.Should().BeTrue(); + result.EntityId.Should().Be(keyId.ToString()); + result.EntityName.Should().Be("Claude"); + } + + [Fact] + public async Task ManageApiKeysTool_PropagatesCreateFailure() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Failure("create_failed")); + var tool = new ManageApiKeysTool(mediator); + + var result = await tool.ExecuteAsync(Parse("""{"action":"create","name":"Claude"}"""), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("create_failed"); + } + + [Fact] + public async Task ManageApiKeysTool_RejectsInvalidKeyIdForRevoke() + { + var tool = new ManageApiKeysTool(Substitute.For()); + + var result = await tool.ExecuteAsync(Parse("""{"action":"revoke","key_id":"bad"}"""), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("key_id must be a valid GUID."); + } + + [Fact] + public async Task ManageApiKeysTool_RevokesKey() + { + var keyId = Guid.NewGuid(); + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success()); + var tool = new ManageApiKeysTool(mediator); + + var result = await tool.ExecuteAsync(Parse($$"""{"action":"revoke","key_id":"{{keyId}}"}"""), UserId, CancellationToken.None); + + result.Success.Should().BeTrue(); + result.EntityId.Should().Be(keyId.ToString()); + result.EntityName.Should().Be("Revoked API key"); + } + + [Fact] + public async Task ManageApiKeysTool_PropagatesRevokeFailure() + { + var keyId = Guid.NewGuid(); + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Failure("revoke_failed")); + var tool = new ManageApiKeysTool(mediator); + + var result = await tool.ExecuteAsync(Parse($$"""{"action":"revoke","key_id":"{{keyId}}"}"""), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("revoke_failed"); + } + + [Fact] + public async Task ManageApiKeysTool_RejectsUnsupportedAction() + { + var tool = new ManageApiKeysTool(Substitute.For()); + + var result = await tool.ExecuteAsync(Parse("""{"action":"unknown"}"""), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("Unsupported action 'unknown'."); + } + + [Fact] + public async Task SendSupportRequestTool_RequiresAllFields() + { + var tool = new SendSupportRequestTool(Substitute.For()); + + var result = await tool.ExecuteAsync(Parse("""{"name":"Thomas"}"""), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("name, email, subject, and message are required."); + } + + [Fact] + public async Task SendSupportRequestTool_ReturnsFailure() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Failure("support_failed")); + var tool = new SendSupportRequestTool(mediator); + + var result = await tool.ExecuteAsync( + Parse("""{"name":"Thomas","email":"t@example.com","subject":"Help","message":"Need support"}"""), + UserId, + CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("support_failed"); + } + + [Fact] + public async Task SendSupportRequestTool_ReturnsSuccess() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success()); + var tool = new SendSupportRequestTool(mediator); + + var result = await tool.ExecuteAsync( + Parse("""{"name":"Thomas","email":"t@example.com","subject":"Help","message":"Need support"}"""), + UserId, + CancellationToken.None); + + result.Success.Should().BeTrue(); + result.EntityName.Should().Be("Support request sent"); + } + + [Fact] + public async Task ManageAccountTool_RequiresAction() + { + var tool = new ManageAccountTool(Substitute.For()); + + var result = await tool.ExecuteAsync(Parse("{}"), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("action is required."); + } + + [Fact] + public async Task ManageAccountTool_ResetsAccount() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success()); + var tool = new ManageAccountTool(mediator); + + var result = await tool.ExecuteAsync(Parse("""{"action":"reset_account"}"""), UserId, CancellationToken.None); + + result.Success.Should().BeTrue(); + result.EntityName.Should().Be("Account reset completed"); + } + + [Fact] + public async Task ManageAccountTool_RequestsDeletion() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success()); + var tool = new ManageAccountTool(mediator); + + var result = await tool.ExecuteAsync(Parse("""{"action":"request_deletion"}"""), UserId, CancellationToken.None); + + result.Success.Should().BeTrue(); + result.EntityName.Should().Be("Deletion code requested"); + } + + [Fact] + public async Task ManageAccountTool_PropagatesRequestDeletionFailure() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Failure("request_failed")); + var tool = new ManageAccountTool(mediator); + + var result = await tool.ExecuteAsync(Parse("""{"action":"request_deletion"}"""), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("request_failed"); + } + + [Fact] + public async Task ManageAccountTool_RequiresCodeForDeletionConfirmation() + { + var tool = new ManageAccountTool(Substitute.For()); + + var result = await tool.ExecuteAsync(Parse("""{"action":"confirm_deletion"}"""), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("code is required."); + } + + [Fact] + public async Task ManageAccountTool_ConfirmsDeletion() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success(new DateTime(2026, 4, 20, 12, 0, 0, DateTimeKind.Utc))); + var tool = new ManageAccountTool(mediator); + + var result = await tool.ExecuteAsync( + Parse("""{"action":"confirm_deletion","code":"123456"}"""), + UserId, + CancellationToken.None); + + result.Success.Should().BeTrue(); + result.EntityName.Should().Be("Account deletion confirmed"); + } + + [Fact] + public async Task ManageAccountTool_PropagatesConfirmDeletionFailure() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Failure("confirm_failed")); + var tool = new ManageAccountTool(mediator); + + var result = await tool.ExecuteAsync( + Parse("""{"action":"confirm_deletion","code":"123456"}"""), + UserId, + CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("confirm_failed"); + } + + [Fact] + public async Task ManageAccountTool_RejectsUnsupportedAction() + { + var tool = new ManageAccountTool(Substitute.For()); + + var result = await tool.ExecuteAsync(Parse("""{"action":"unknown"}"""), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("Unsupported action 'unknown'."); + } + + private static JsonElement Parse(string json) + { + return JsonDocument.Parse(json).RootElement.Clone(); + } +} diff --git a/tests/Orbit.Application.Tests/Chat/Tools/ProfileNotificationCalendarToolTests.cs b/tests/Orbit.Application.Tests/Chat/Tools/ProfileNotificationCalendarToolTests.cs new file mode 100644 index 00000000..dc2061e9 --- /dev/null +++ b/tests/Orbit.Application.Tests/Chat/Tools/ProfileNotificationCalendarToolTests.cs @@ -0,0 +1,774 @@ +using System.Text.Json; +using FluentAssertions; +using MediatR; +using NSubstitute; +using Orbit.Application.Calendar.Commands; +using Orbit.Application.Calendar.Queries; +using Orbit.Application.Chat.Tools; +using Orbit.Application.Chat.Tools.Implementations; +using Orbit.Application.Notifications.Commands; +using Orbit.Application.Notifications.Queries; +using Orbit.Application.Profile.Commands; +using Orbit.Application.Profile.Queries; +using Orbit.Domain.Common; +using Orbit.Domain.Enums; + +namespace Orbit.Application.Tests.Chat.Tools; + +public class ProfileNotificationCalendarToolTests +{ + private static readonly Guid UserId = Guid.NewGuid(); + + [Fact] + public void ToolMetadata_ExposesNamesAndSchemas() + { + var mediator = Substitute.For(); + var profileTool = new GetProfileTool(mediator); + var preferencesTool = new UpdateProfilePreferencesTool(mediator); + var aiSettingsTool = new UpdateAiSettingsTool(mediator); + var notificationsTool = new GetNotificationsTool(mediator); + var updateNotificationsTool = new UpdateNotificationsTool(mediator); + var deleteNotificationsTool = new DeleteNotificationsTool(mediator); + var calendarOverviewTool = new GetCalendarOverviewTool(mediator); + var calendarSyncTool = new ManageCalendarSyncTool(mediator); + + profileTool.Name.Should().Be("get_profile"); + profileTool.IsReadOnly.Should().BeTrue(); + JsonSerializer.Serialize(profileTool.GetParameterSchema()).Should().Contain("properties"); + + preferencesTool.Name.Should().Be("update_profile_preferences"); + JsonSerializer.Serialize(preferencesTool.GetParameterSchema()).Should().Contain("set_theme_preference"); + JsonSerializer.Serialize(preferencesTool.GetParameterSchema()).Should().Contain("color_scheme"); + + aiSettingsTool.Name.Should().Be("update_ai_settings"); + JsonSerializer.Serialize(aiSettingsTool.GetParameterSchema()).Should().Contain("set_ai_summary"); + + notificationsTool.Name.Should().Be("get_notifications"); + notificationsTool.IsReadOnly.Should().BeTrue(); + JsonSerializer.Serialize(notificationsTool.GetParameterSchema()).Should().Contain("properties"); + + updateNotificationsTool.Name.Should().Be("update_notifications"); + JsonSerializer.Serialize(updateNotificationsTool.GetParameterSchema()).Should().Contain("subscribe_push"); + + deleteNotificationsTool.Name.Should().Be("delete_notifications"); + JsonSerializer.Serialize(deleteNotificationsTool.GetParameterSchema()).Should().Contain("delete_all"); + + calendarOverviewTool.Name.Should().Be("get_calendar_overview"); + calendarOverviewTool.IsReadOnly.Should().BeTrue(); + JsonSerializer.Serialize(calendarOverviewTool.GetParameterSchema()).Should().Contain("include_auto_sync_state"); + + calendarSyncTool.Name.Should().Be("manage_calendar_sync"); + JsonSerializer.Serialize(calendarSyncTool.GetParameterSchema()).Should().Contain("dismiss_suggestion"); + } + + [Fact] + public async Task GetProfileTool_ReturnsSuccessPayload() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success(null!)); + var tool = new GetProfileTool(mediator); + + var result = await tool.ExecuteAsync(Parse("{}"), UserId, CancellationToken.None); + + result.Success.Should().BeTrue(); + result.Error.Should().BeNull(); + } + + [Fact] + public async Task GetProfileTool_ReturnsFailure() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Failure("user_not_found")); + var tool = new GetProfileTool(mediator); + + var result = await tool.ExecuteAsync(Parse("{}"), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("user_not_found"); + } + + [Fact] + public async Task UpdateProfilePreferencesTool_RequiresAction() + { + var tool = new UpdateProfilePreferencesTool(Substitute.For()); + + var result = await tool.ExecuteAsync(Parse("{}"), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("action is required."); + } + + [Fact] + public async Task UpdateProfilePreferencesTool_SetsTimezone() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success()); + var tool = new UpdateProfilePreferencesTool(mediator); + + var result = await tool.ExecuteAsync(Parse("""{"action":"set_timezone","timezone":"America/New_York"}"""), UserId, CancellationToken.None); + + result.Success.Should().BeTrue(); + result.EntityName.Should().Be("Timezone set to America/New_York"); + await mediator.Received(1).Send( + Arg.Is(command => command.UserId == UserId && command.TimeZone == "America/New_York"), + Arg.Any()); + } + + [Fact] + public async Task UpdateProfilePreferencesTool_RequiresTimezone() + { + var tool = new UpdateProfilePreferencesTool(Substitute.For()); + + var result = await tool.ExecuteAsync(Parse("""{"action":"set_timezone"}"""), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("timezone is required."); + } + + [Fact] + public async Task UpdateProfilePreferencesTool_SetsLanguage() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success()); + var tool = new UpdateProfilePreferencesTool(mediator); + + var result = await tool.ExecuteAsync(Parse("""{"action":"set_language","language":"pt-BR"}"""), UserId, CancellationToken.None); + + result.Success.Should().BeTrue(); + result.EntityName.Should().Be("Language set to pt-BR"); + } + + [Fact] + public async Task UpdateProfilePreferencesTool_RequiresWeekStartDay() + { + var tool = new UpdateProfilePreferencesTool(Substitute.For()); + + var result = await tool.ExecuteAsync(Parse("""{"action":"set_week_start_day"}"""), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("week_start_day is required."); + } + + [Fact] + public async Task UpdateProfilePreferencesTool_RequiresThemePreferenceProperty() + { + var tool = new UpdateProfilePreferencesTool(Substitute.For()); + + var result = await tool.ExecuteAsync(Parse("""{"action":"set_theme_preference"}"""), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("theme_preference is required."); + } + + [Fact] + public async Task UpdateProfilePreferencesTool_SetsThemePreference() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success()); + var tool = new UpdateProfilePreferencesTool(mediator); + + var result = await tool.ExecuteAsync( + Parse("""{"action":"set_theme_preference","theme_preference":"dark"}"""), + UserId, + CancellationToken.None); + + result.Success.Should().BeTrue(); + result.EntityName.Should().Be("Theme preference updated"); + } + + [Fact] + public async Task UpdateProfilePreferencesTool_RequiresColorSchemeProperty() + { + var tool = new UpdateProfilePreferencesTool(Substitute.For()); + + var result = await tool.ExecuteAsync(Parse("""{"action":"set_color_scheme"}"""), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("color_scheme is required."); + } + + [Fact] + public async Task UpdateProfilePreferencesTool_SetsColorScheme() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success()); + var tool = new UpdateProfilePreferencesTool(mediator); + + var result = await tool.ExecuteAsync( + Parse("""{"action":"set_color_scheme","color_scheme":"sunset"}"""), + UserId, + CancellationToken.None); + + result.Success.Should().BeTrue(); + result.EntityName.Should().Be("Color scheme updated"); + } + + [Fact] + public async Task UpdateProfilePreferencesTool_CompletesOnboarding() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success()); + var tool = new UpdateProfilePreferencesTool(mediator); + + var result = await tool.ExecuteAsync(Parse("""{"action":"complete_onboarding"}"""), UserId, CancellationToken.None); + + result.Success.Should().BeTrue(); + result.EntityName.Should().Be("Onboarding completed"); + } + + [Fact] + public async Task UpdateProfilePreferencesTool_CompletesTour() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success()); + var tool = new UpdateProfilePreferencesTool(mediator); + + var result = await tool.ExecuteAsync(Parse("""{"action":"complete_tour"}"""), UserId, CancellationToken.None); + + result.Success.Should().BeTrue(); + result.EntityName.Should().Be("Tour completed"); + } + + [Fact] + public async Task UpdateProfilePreferencesTool_ResetsTour() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success()); + var tool = new UpdateProfilePreferencesTool(mediator); + + var result = await tool.ExecuteAsync(Parse("""{"action":"reset_tour"}"""), UserId, CancellationToken.None); + + result.Success.Should().BeTrue(); + result.EntityName.Should().Be("Tour reset"); + } + + [Fact] + public async Task UpdateProfilePreferencesTool_ReturnsFailureForUnsupportedAction() + { + var tool = new UpdateProfilePreferencesTool(Substitute.For()); + + var result = await tool.ExecuteAsync(Parse("""{"action":"unknown"}"""), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("Unsupported action 'unknown'."); + } + + [Fact] + public async Task UpdateAiSettingsTool_RequiresEnabled() + { + var tool = new UpdateAiSettingsTool(Substitute.For()); + + var result = await tool.ExecuteAsync(Parse("""{"action":"set_ai_memory"}"""), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("action and enabled are required."); + } + + [Fact] + public async Task UpdateAiSettingsTool_BuildsAiMemoryEntityName() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success()); + var tool = new UpdateAiSettingsTool(mediator); + + var result = await tool.ExecuteAsync(Parse("""{"action":"set_ai_memory","enabled":true}"""), UserId, CancellationToken.None); + + result.Success.Should().BeTrue(); + result.EntityName.Should().Be("AI memory enabled"); + } + + [Fact] + public async Task UpdateAiSettingsTool_BuildsAiSummaryDisabledEntityName() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success()); + var tool = new UpdateAiSettingsTool(mediator); + + var result = await tool.ExecuteAsync(Parse("""{"action":"set_ai_summary","enabled":false}"""), UserId, CancellationToken.None); + + result.Success.Should().BeTrue(); + result.EntityName.Should().Be("AI summary disabled"); + } + + [Fact] + public async Task UpdateAiSettingsTool_PropagatesMediatorFailure() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Failure("summary_failed")); + var tool = new UpdateAiSettingsTool(mediator); + + var result = await tool.ExecuteAsync(Parse("""{"action":"set_ai_summary","enabled":true}"""), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("summary_failed"); + } + + [Fact] + public async Task UpdateAiSettingsTool_ReturnsFailureForUnsupportedAction() + { + var tool = new UpdateAiSettingsTool(Substitute.For()); + + var result = await tool.ExecuteAsync(Parse("""{"action":"unknown","enabled":false}"""), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("Unsupported action 'unknown'."); + } + + [Fact] + public async Task GetNotificationsTool_ReturnsSuccessPayload() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success(new GetNotificationsResponse([], 0))); + var tool = new GetNotificationsTool(mediator); + + var result = await tool.ExecuteAsync(Parse("{}"), UserId, CancellationToken.None); + + result.Success.Should().BeTrue(); + } + + [Fact] + public async Task GetNotificationsTool_ReturnsFailure() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Failure("notifications_failed")); + var tool = new GetNotificationsTool(mediator); + + var result = await tool.ExecuteAsync(Parse("{}"), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("notifications_failed"); + } + + [Fact] + public async Task UpdateNotificationsTool_RequiresAction() + { + var tool = new UpdateNotificationsTool(Substitute.For()); + + var result = await tool.ExecuteAsync(Parse("{}"), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("action is required."); + } + + [Fact] + public async Task UpdateNotificationsTool_RejectsInvalidNotificationId() + { + var tool = new UpdateNotificationsTool(Substitute.For()); + + var result = await tool.ExecuteAsync(Parse("""{"action":"mark_read","notification_id":"bad"}"""), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("notification_id must be a valid GUID."); + } + + [Fact] + public async Task UpdateNotificationsTool_MarksNotificationRead() + { + var notificationId = Guid.NewGuid(); + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success()); + var tool = new UpdateNotificationsTool(mediator); + + var result = await tool.ExecuteAsync( + Parse($$"""{"action":"mark_read","notification_id":"{{notificationId}}"}"""), + UserId, + CancellationToken.None); + + result.Success.Should().BeTrue(); + result.EntityId.Should().Be(notificationId.ToString()); + result.EntityName.Should().Be("Marked notification as read"); + } + + [Fact] + public async Task UpdateNotificationsTool_MarksAllNotificationsRead() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success()); + var tool = new UpdateNotificationsTool(mediator); + + var result = await tool.ExecuteAsync(Parse("""{"action":"mark_all_read"}"""), UserId, CancellationToken.None); + + result.Success.Should().BeTrue(); + result.EntityName.Should().Be("Marked all notifications as read"); + } + + [Fact] + public async Task UpdateNotificationsTool_RequiresSubscriptionFields() + { + var tool = new UpdateNotificationsTool(Substitute.For()); + + var result = await tool.ExecuteAsync(Parse("""{"action":"subscribe_push","endpoint":"https://push"}"""), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("endpoint, p256dh, and auth are required."); + } + + [Fact] + public async Task UpdateNotificationsTool_SubscribesPush() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success()); + var tool = new UpdateNotificationsTool(mediator); + + var result = await tool.ExecuteAsync( + Parse("""{"action":"subscribe_push","endpoint":"https://push","p256dh":"key","auth":"secret"}"""), + UserId, + CancellationToken.None); + + result.Success.Should().BeTrue(); + result.EntityName.Should().Be("Push subscription registered"); + } + + [Fact] + public async Task UpdateNotificationsTool_UnsubscribesPush() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success()); + var tool = new UpdateNotificationsTool(mediator); + + var result = await tool.ExecuteAsync( + Parse("""{"action":"unsubscribe_push","endpoint":"https://push"}"""), + UserId, + CancellationToken.None); + + result.Success.Should().BeTrue(); + result.EntityName.Should().Be("Push subscription removed"); + } + + [Fact] + public async Task UpdateNotificationsTool_RequiresEndpointForUnsubscribe() + { + var tool = new UpdateNotificationsTool(Substitute.For()); + + var result = await tool.ExecuteAsync(Parse("""{"action":"unsubscribe_push"}"""), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("endpoint is required."); + } + + [Fact] + public async Task UpdateNotificationsTool_HandlesTestPushFailure() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Failure("push_failed")); + var tool = new UpdateNotificationsTool(mediator); + + var result = await tool.ExecuteAsync(Parse("""{"action":"test_push"}"""), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("push_failed"); + } + + [Fact] + public async Task UpdateNotificationsTool_ReturnsTestPushSuccess() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success(new TestPushNotificationResponse(2, "sent"))); + var tool = new UpdateNotificationsTool(mediator); + + var result = await tool.ExecuteAsync(Parse("""{"action":"test_push"}"""), UserId, CancellationToken.None); + + result.Success.Should().BeTrue(); + result.EntityName.Should().Be("Test push requested"); + } + + [Fact] + public async Task UpdateNotificationsTool_ReturnsFailureForUnsupportedAction() + { + var tool = new UpdateNotificationsTool(Substitute.For()); + + var result = await tool.ExecuteAsync(Parse("""{"action":"unknown"}"""), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("Unsupported action 'unknown'."); + } + + [Fact] + public async Task DeleteNotificationsTool_RequiresAction() + { + var tool = new DeleteNotificationsTool(Substitute.For()); + + var result = await tool.ExecuteAsync(Parse("{}"), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("action is required."); + } + + [Fact] + public async Task DeleteNotificationsTool_DeletesAll() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success()); + var tool = new DeleteNotificationsTool(mediator); + + var result = await tool.ExecuteAsync(Parse("""{"action":"delete_all"}"""), UserId, CancellationToken.None); + + result.Success.Should().BeTrue(); + result.EntityName.Should().Be("Deleted all notifications"); + } + + [Fact] + public async Task DeleteNotificationsTool_DeletesOne() + { + var notificationId = Guid.NewGuid(); + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success()); + var tool = new DeleteNotificationsTool(mediator); + + var result = await tool.ExecuteAsync( + Parse($$"""{"action":"delete_one","notification_id":"{{notificationId}}"}"""), + UserId, + CancellationToken.None); + + result.Success.Should().BeTrue(); + result.EntityId.Should().Be(notificationId.ToString()); + result.EntityName.Should().Be("Deleted notification"); + } + + [Fact] + public async Task DeleteNotificationsTool_RejectsInvalidNotificationId() + { + var tool = new DeleteNotificationsTool(Substitute.For()); + + var result = await tool.ExecuteAsync(Parse("""{"action":"delete_one","notification_id":"bad"}"""), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("notification_id must be a valid GUID."); + } + + [Fact] + public async Task DeleteNotificationsTool_ReturnsFailureForUnsupportedAction() + { + var tool = new DeleteNotificationsTool(Substitute.For()); + + var result = await tool.ExecuteAsync(Parse("""{"action":"unknown"}"""), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("Unsupported action 'unknown'."); + } + + [Fact] + public async Task GetCalendarOverviewTool_ReturnsOverview() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success(new List())); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success(new CalendarAutoSyncStateResponse(true, GoogleCalendarAutoSyncStatus.Idle, null, true))); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success(new List())); + var tool = new GetCalendarOverviewTool(mediator); + + var result = await tool.ExecuteAsync(Parse("{}"), UserId, CancellationToken.None); + + result.Success.Should().BeTrue(); + } + + [Fact] + public async Task GetCalendarOverviewTool_StopsOnEventsFailure() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Failure>("events_failed")); + var tool = new GetCalendarOverviewTool(mediator); + + var result = await tool.ExecuteAsync(Parse("{}"), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("events_failed"); + } + + [Fact] + public async Task GetCalendarOverviewTool_StopsOnAutoSyncFailure() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success(new List())); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Failure("auto_sync_failed")); + var tool = new GetCalendarOverviewTool(mediator); + + var result = await tool.ExecuteAsync(Parse("{}"), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("auto_sync_failed"); + } + + [Fact] + public async Task GetCalendarOverviewTool_StopsOnSuggestionsFailure() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success(new List())); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success(new CalendarAutoSyncStateResponse(true, GoogleCalendarAutoSyncStatus.Idle, null, true))); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Failure>("suggestions_failed")); + var tool = new GetCalendarOverviewTool(mediator); + + var result = await tool.ExecuteAsync(Parse("{}"), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("suggestions_failed"); + } + + [Fact] + public async Task GetCalendarOverviewTool_SkipsQueriesWhenFlagsAreFalse() + { + var mediator = Substitute.For(); + var tool = new GetCalendarOverviewTool(mediator); + + var result = await tool.ExecuteAsync( + Parse("""{"include_events":false,"include_auto_sync_state":false,"include_suggestions":false}"""), + UserId, + CancellationToken.None); + + result.Success.Should().BeTrue(); + await mediator.DidNotReceiveWithAnyArgs().Send(default!, default); + } + + [Fact] + public async Task ManageCalendarSyncTool_RequiresAction() + { + var tool = new ManageCalendarSyncTool(Substitute.For()); + + var result = await tool.ExecuteAsync(Parse("{}"), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("action is required."); + } + + [Fact] + public async Task ManageCalendarSyncTool_RequiresEnabledForAutoSync() + { + var tool = new ManageCalendarSyncTool(Substitute.For()); + + var result = await tool.ExecuteAsync(Parse("""{"action":"set_auto_sync"}"""), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("enabled is required."); + } + + [Fact] + public async Task ManageCalendarSyncTool_SetsAutoSync() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success()); + var tool = new ManageCalendarSyncTool(mediator); + + var result = await tool.ExecuteAsync(Parse("""{"action":"set_auto_sync","enabled":true}"""), UserId, CancellationToken.None); + + result.Success.Should().BeTrue(); + result.EntityName.Should().Be("Calendar auto-sync enabled"); + } + + [Fact] + public async Task ManageCalendarSyncTool_RejectsInvalidSuggestionId() + { + var tool = new ManageCalendarSyncTool(Substitute.For()); + + var result = await tool.ExecuteAsync(Parse("""{"action":"dismiss_suggestion","suggestion_id":"bad"}"""), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("suggestion_id must be a valid GUID."); + } + + [Fact] + public async Task ManageCalendarSyncTool_DismissesSuggestion() + { + var suggestionId = Guid.NewGuid(); + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success()); + var tool = new ManageCalendarSyncTool(mediator); + + var result = await tool.ExecuteAsync( + Parse($$"""{"action":"dismiss_suggestion","suggestion_id":"{{suggestionId}}"}"""), + UserId, + CancellationToken.None); + + result.Success.Should().BeTrue(); + result.EntityId.Should().Be(suggestionId.ToString()); + } + + [Fact] + public async Task ManageCalendarSyncTool_DismissesImport() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success()); + var tool = new ManageCalendarSyncTool(mediator); + + var result = await tool.ExecuteAsync(Parse("""{"action":"dismiss_import"}"""), UserId, CancellationToken.None); + + result.Success.Should().BeTrue(); + result.EntityName.Should().Be("Dismissed calendar import prompt"); + } + + [Fact] + public async Task ManageCalendarSyncTool_RunsSync() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Success(new CalendarAutoSyncResult(1, 2, GoogleCalendarAutoSyncStatus.Idle))); + var tool = new ManageCalendarSyncTool(mediator); + + var result = await tool.ExecuteAsync(Parse("""{"action":"run_sync"}"""), UserId, CancellationToken.None); + + result.Success.Should().BeTrue(); + result.EntityName.Should().Be("Calendar sync requested"); + } + + [Fact] + public async Task ManageCalendarSyncTool_HandlesRunSyncFailure() + { + var mediator = Substitute.For(); + mediator.Send(Arg.Any(), Arg.Any()) + .Returns(Result.Failure("sync_failed")); + var tool = new ManageCalendarSyncTool(mediator); + + var result = await tool.ExecuteAsync(Parse("""{"action":"run_sync"}"""), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("sync_failed"); + } + + [Fact] + public async Task ManageCalendarSyncTool_ReturnsFailureForUnsupportedAction() + { + var tool = new ManageCalendarSyncTool(Substitute.For()); + + var result = await tool.ExecuteAsync(Parse("""{"action":"unknown"}"""), UserId, CancellationToken.None); + + result.Success.Should().BeFalse(); + result.Error.Should().Be("Unsupported action 'unknown'."); + } + + private static JsonElement Parse(string json) + { + return JsonDocument.Parse(json).RootElement.Clone(); + } +} diff --git a/tests/Orbit.Domain.Tests/Entities/AgentSessionAndSyncEntityTests.cs b/tests/Orbit.Domain.Tests/Entities/AgentSessionAndSyncEntityTests.cs new file mode 100644 index 00000000..501824df --- /dev/null +++ b/tests/Orbit.Domain.Tests/Entities/AgentSessionAndSyncEntityTests.cs @@ -0,0 +1,203 @@ +using System.Reflection; +using System.Text.Json; +using FluentAssertions; +using Orbit.Domain.Entities; +using Orbit.Domain.Models; + +namespace Orbit.Domain.Tests.Entities; + +public class AgentSessionAndSyncEntityTests +{ + [Fact] + public void AgentStepUpChallengeState_CreateAndVerifyLifecycle_Works() + { + var userId = Guid.NewGuid(); + var pendingOperationId = Guid.NewGuid(); + var expiresAtUtc = DateTime.UtcNow.AddMinutes(10); + + var state = AgentStepUpChallengeState.Create(userId, pendingOperationId, "hash", expiresAtUtc); + + state.UserId.Should().Be(userId); + state.PendingOperationId.Should().Be(pendingOperationId); + state.CodeHash.Should().Be("hash"); + state.AttemptCount.Should().Be(0); + state.IsExpired(expiresAtUtc.AddMinutes(-1)).Should().BeFalse(); + state.CanVerify(3, expiresAtUtc.AddMinutes(-1)).Should().BeTrue(); + + state.RecordFailedAttempt(); + state.AttemptCount.Should().Be(1); + + state.MarkVerified(); + state.VerifiedAtUtc.Should().NotBeNull(); + state.CanVerify(3, expiresAtUtc.AddMinutes(-1)).Should().BeFalse(); + state.IsExpired(expiresAtUtc.AddMinutes(1)).Should().BeTrue(); + } + + [Fact] + public void UserSession_CreateRotateRevokeAndValidation_Works() + { + var userId = Guid.NewGuid(); + var expiresAtUtc = DateTime.UtcNow.AddDays(7); + + UserSession.Create(Guid.Empty, "hash", expiresAtUtc).IsFailure.Should().BeTrue(); + UserSession.Create(userId, "", expiresAtUtc).IsFailure.Should().BeTrue(); + + var session = UserSession.Create(userId, "hash", expiresAtUtc).Value; + + session.UserId.Should().Be(userId); + session.TokenHash.Should().Be("hash"); + session.CanUse(expiresAtUtc.AddMinutes(-1)).Should().BeTrue(); + session.CanUse(expiresAtUtc.AddMinutes(1)).Should().BeFalse(); + + var rotatedExpiry = expiresAtUtc.AddDays(1); + var usedAtUtc = DateTime.UtcNow.AddHours(1); + session.Rotate("new-hash", rotatedExpiry, usedAtUtc); + + session.TokenHash.Should().Be("new-hash"); + session.ExpiresAtUtc.Should().Be(rotatedExpiry); + session.LastUsedAtUtc.Should().Be(usedAtUtc); + + var revokedAtUtc = DateTime.UtcNow.AddHours(2); + session.Revoke(revokedAtUtc); + session.RevokedAtUtc.Should().Be(revokedAtUtc); + session.CanUse(revokedAtUtc.AddMinutes(-1)).Should().BeFalse(); + + session.Revoke(revokedAtUtc.AddMinutes(5)); + session.RevokedAtUtc.Should().Be(revokedAtUtc); + } + + [Fact] + public void GoogleCalendarSyncSuggestion_CreateDismissAndImport_Works() + { + var userId = Guid.NewGuid(); + var suggestion = GoogleCalendarSyncSuggestion.Create( + userId, + "event-123", + "Team Sync", + new DateTime(2026, 4, 14, 12, 0, 0, DateTimeKind.Utc), + """{"id":"event-123"}""", + new DateTime(2026, 4, 13, 8, 0, 0, DateTimeKind.Utc)); + + suggestion.UserId.Should().Be(userId); + suggestion.GoogleEventId.Should().Be("event-123"); + suggestion.Title.Should().Be("Team Sync"); + suggestion.RawEventJson.Should().Contain("event-123"); + + suggestion.MarkDismissed(new DateTime(2026, 4, 14, 13, 0, 0, DateTimeKind.Utc)); + suggestion.DismissedAtUtc.Should().Be(new DateTime(2026, 4, 14, 13, 0, 0, DateTimeKind.Utc)); + + var habitId = Guid.NewGuid(); + suggestion.MarkImported(habitId, new DateTime(2026, 4, 14, 14, 0, 0, DateTimeKind.Utc)); + suggestion.ImportedHabitId.Should().Be(habitId); + suggestion.ImportedAtUtc.Should().Be(new DateTime(2026, 4, 14, 14, 0, 0, DateTimeKind.Utc)); + } + + [Fact] + public void DistributedRateLimitBucket_CreateAndIncrement_Works() + { + var windowStartUtc = DateTime.UtcNow; + var windowEndsAtUtc = windowStartUtc.AddMinutes(1); + + var bucket = DistributedRateLimitBucket.Create("auth", "user:1", windowStartUtc, windowEndsAtUtc); + + bucket.PolicyName.Should().Be("auth"); + bucket.PartitionKey.Should().Be("user:1"); + bucket.WindowStartUtc.Should().Be(windowStartUtc); + bucket.WindowEndsAtUtc.Should().Be(windowEndsAtUtc); + bucket.Count.Should().Be(1); + + var previousUpdatedAtUtc = bucket.UpdatedAtUtc; + bucket.Increment(); + + bucket.Count.Should().Be(2); + bucket.UpdatedAtUtc.Should().BeOnOrAfter(previousUpdatedAtUtc); + } + + [Fact] + public void AgentContracts_RecordsAndCatalogConstants_AreAccessible() + { + var requestSchema = Parse("""{"type":"object"}"""); + var responseSchema = Parse("""{"type":"object"}"""); + + var capability = new AgentCapability( + AgentCapabilityIds.HabitsWrite, + "Habits Write", + "Write habits", + "habits", + AgentScopes.WriteHabits, + AgentRiskClass.Low, + true, + false, + AgentConfirmationRequirement.FreshConfirmation, + ChatToolNames: ["create_habit"]); + var operation = new AgentOperation( + "create_habit", + "Create habit", + "Create a habit", + capability.Id, + AgentRiskClass.Low, + AgentConfirmationRequirement.FreshConfirmation, + true, + true, + requestSchema, + responseSchema); + var surface = new AppSurface("today", "Today", "Today surface", ["Open Today"], ["Pinned"], [capability.Id], ["HabitsController.Get"]); + var field = new UserDataFieldDescriptor("title", "Habit title", true, true); + var catalog = new UserDataCatalogEntry("habits", "Habits", "Habit catalog", "medium", "keep", true, true, [field]); + var pendingOperation = new PendingAgentOperation(Guid.NewGuid(), capability.Id, "Create habit", "Create a new habit", AgentRiskClass.High, AgentConfirmationRequirement.StepUp, DateTime.UtcNow.AddMinutes(10)); + var confirmation = new PendingAgentOperationConfirmation(pendingOperation.Id, "agc_token", DateTime.UtcNow.AddMinutes(5)); + var execution = new PendingAgentOperationExecution(pendingOperation.Id, capability.Id, operation.Id, requestSchema, AgentExecutionSurface.Chat, AgentConfirmationRequirement.StepUp); + var challenge = new AgentStepUpChallenge(Guid.NewGuid(), pendingOperation.Id, DateTime.UtcNow.AddMinutes(10)); + var evaluationContext = new AgentPolicyEvaluationContext(capability.Id, Guid.NewGuid(), AgentExecutionSurface.Mcp, AgentAuthMethod.Jwt, [AgentScopes.WriteHabits], "execute_agent_operation_v2", "Create habit"); + var policyDecision = new AgentPolicyDecision(AgentPolicyDecisionStatus.ConfirmationRequired, capability, PendingOperation: pendingOperation); + var executeRequest = new AgentExecuteOperationRequest(Guid.NewGuid(), operation.Id, requestSchema, AgentExecutionSurface.Mcp, AgentAuthMethod.Jwt, [AgentScopes.WriteHabits], ConfirmationToken: confirmation.ConfirmationToken); + var operationResult = new AgentOperationResult(operation.Id, "execute_agent_operation_v2", AgentRiskClass.High, AgentConfirmationRequirement.StepUp, AgentOperationStatus.PendingConfirmation, PendingOperationId: pendingOperation.Id, Payload: new { habitId = Guid.NewGuid() }); + var denial = new AgentPolicyDenial(operation.Id, "execute_agent_operation_v2", AgentRiskClass.High, AgentConfirmationRequirement.StepUp, "denied"); + var executeResponse = new AgentExecuteOperationResponse(operationResult, pendingOperation, denial); + var clientContext = new AgentClientContext("ios", "en-US", "24h", "today", true); + var snapshot = new AgentContextSnapshot("pro", "en-US", "UTC", true, true, 1, "dark", "sunset", true, true, "Idle", ["beta"], ["Health"], ["Morning"], ["Run"], ["Lose weight"], clientContext); + var auditEntry = new AgentAuditEntry(Guid.NewGuid(), capability.Id, "execute_agent_operation_v2", AgentExecutionSurface.Mcp, AgentAuthMethod.Jwt, AgentRiskClass.High, AgentPolicyDecisionStatus.Allowed, AgentOperationStatus.Succeeded, Summary: "Created habit"); + + capability.ChatToolNames.Should().Contain("create_habit"); + operation.RequestSchema.GetProperty("type").GetString().Should().Be("object"); + surface.HowToSteps.Should().ContainSingle(); + catalog.Fields.Should().ContainSingle().Which.Name.Should().Be("title"); + pendingOperation.ConfirmationRequirement.Should().Be(AgentConfirmationRequirement.StepUp); + confirmation.ConfirmationToken.Should().Be("agc_token"); + execution.Surface.Should().Be(AgentExecutionSurface.Chat); + challenge.PendingOperationId.Should().Be(pendingOperation.Id); + evaluationContext.GrantedScopes.Should().Contain(AgentScopes.WriteHabits); + policyDecision.PendingOperation.Should().Be(pendingOperation); + executeRequest.ConfirmationToken.Should().Be("agc_token"); + operationResult.PendingOperationId.Should().Be(pendingOperation.Id); + executeResponse.PolicyDenial.Should().Be(denial); + clientContext.Platform.Should().Be("ios"); + snapshot.ClientContext.Should().Be(clientContext); + auditEntry.Summary.Should().Be("Created habit"); + + AgentScopes.ClaudeDefaultScopes.Should().Contain(AgentScopes.ChatInteract); + AgentScopes.ClaudeDefaultScopes.Should().OnlyHaveUniqueItems(); + + var capabilityIds = typeof(AgentCapabilityIds) + .GetFields(BindingFlags.Public | BindingFlags.Static) + .Select(fieldInfo => fieldInfo.GetValue(null)) + .OfType() + .ToList(); + var scopes = typeof(AgentScopes) + .GetFields(BindingFlags.Public | BindingFlags.Static) + .Where(fieldInfo => fieldInfo.FieldType == typeof(string)) + .Select(fieldInfo => fieldInfo.GetValue(null)) + .OfType() + .ToList(); + + capabilityIds.Should().OnlyContain(value => !string.IsNullOrWhiteSpace(value)); + capabilityIds.Should().OnlyHaveUniqueItems(); + scopes.Should().OnlyContain(value => !string.IsNullOrWhiteSpace(value)); + scopes.Should().OnlyHaveUniqueItems(); + } + + private static JsonElement Parse(string json) + { + return JsonDocument.Parse(json).RootElement.Clone(); + } +} diff --git a/tests/Orbit.Domain.Tests/Entities/AgentStateAndAuditTests.cs b/tests/Orbit.Domain.Tests/Entities/AgentStateAndAuditTests.cs new file mode 100644 index 00000000..8dc495aa --- /dev/null +++ b/tests/Orbit.Domain.Tests/Entities/AgentStateAndAuditTests.cs @@ -0,0 +1,158 @@ +using FluentAssertions; +using Orbit.Domain.Entities; +using Orbit.Domain.Models; + +namespace Orbit.Domain.Tests.Entities; + +public class AgentStateAndAuditTests +{ + [Fact] + public void PendingAgentOperationState_Create_UsesDefaultsForBlankArguments() + { + var capability = CreateCapability(AgentCapabilityIds.HabitsDelete, AgentRiskClass.Destructive, AgentConfirmationRequirement.FreshConfirmation); + + var state = PendingAgentOperationState.Create(new PendingAgentOperationStateCreateRequest + { + UserId = Guid.NewGuid(), + Capability = capability, + OperationId = "delete_habit", + ArgumentsJson = " ", + Summary = "Delete habit", + OperationFingerprint = "delete_habit:{}", + Surface = AgentExecutionSurface.Chat, + ExpiresAtUtc = DateTime.UtcNow.AddMinutes(10) + }); + + state.ArgumentsJson.Should().Be("{}"); + state.DisplayName.Should().Be(capability.DisplayName); + state.RiskClass.Should().Be(AgentRiskClass.Destructive); + } + + [Fact] + public void PendingAgentOperationState_IsUsable_RequiresConfirmation() + { + var state = CreateState(); + + state.IsUsable(state.CapabilityId, state.OperationFingerprint, requireStepUp: false, DateTime.UtcNow) + .Should().BeFalse(); + } + + [Fact] + public void PendingAgentOperationState_IsUsable_RequiresStepUpWhenRequested() + { + var state = CreateState(); + state.SetConfirmationTokenHash("hash"); + + state.IsUsable(state.CapabilityId, state.OperationFingerprint, requireStepUp: true, DateTime.UtcNow) + .Should().BeFalse(); + } + + [Fact] + public void PendingAgentOperationState_IsUsable_ReturnsTrueAfterConfirmationAndStepUp() + { + var state = CreateState(); + state.SetConfirmationTokenHash("hash"); + state.MarkStepUpSatisfied(); + + state.IsUsable(state.CapabilityId, state.OperationFingerprint, requireStepUp: true, DateTime.UtcNow) + .Should().BeTrue(); + } + + [Fact] + public void PendingAgentOperationState_IsUsable_ReturnsFalseAfterConsumed() + { + var state = CreateState(); + state.SetConfirmationTokenHash("hash"); + state.MarkConsumed(); + + state.IsUsable(state.CapabilityId, state.OperationFingerprint, requireStepUp: false, DateTime.UtcNow) + .Should().BeFalse(); + } + + [Fact] + public void AgentAuditLog_Create_TruncatesLongFields() + { + var entry = new AgentAuditEntry( + Guid.NewGuid(), + new string('c', 150), + new string('s', 150), + AgentExecutionSurface.Mcp, + AgentAuthMethod.Jwt, + AgentRiskClass.High, + AgentPolicyDecisionStatus.Denied, + AgentOperationStatus.Failed, + new string('i', 150), + new string('u', 600), + new string('t', 150), + new string('n', 250), + new string('a', 5000), + new string('e', 600), + AgentPolicyDecisionStatus.Allowed, + new string('r', 600)); + + var log = AgentAuditLog.Create(entry); + + log.CapabilityId.Length.Should().Be(100); + log.SourceName.Length.Should().Be(100); + log.CorrelationId!.Length.Should().Be(100); + log.Summary!.Length.Should().Be(500); + log.TargetId!.Length.Should().Be(100); + log.TargetName!.Length.Should().Be(200); + log.RedactedArguments!.Length.Should().Be(4000); + log.Error!.Length.Should().Be(500); + log.ShadowReason!.Length.Should().Be(500); + } + + [Fact] + public void AgentAuditLog_Create_PreservesWhitespaceOnlyValues() + { + var entry = new AgentAuditEntry( + Guid.NewGuid(), + "capability", + "source", + AgentExecutionSurface.Chat, + AgentAuthMethod.ApiKey, + AgentRiskClass.Low, + AgentPolicyDecisionStatus.Allowed, + AgentOperationStatus.Succeeded, + Summary: " "); + + var log = AgentAuditLog.Create(entry); + + log.Summary.Should().Be(" "); + log.PolicyDecision.Should().Be(AgentPolicyDecisionStatus.Allowed); + log.OutcomeStatus.Should().Be(AgentOperationStatus.Succeeded); + } + + private static PendingAgentOperationState CreateState() + { + return PendingAgentOperationState.Create(new PendingAgentOperationStateCreateRequest + { + UserId = Guid.NewGuid(), + Capability = CreateCapability(AgentCapabilityIds.ApiKeysManage, AgentRiskClass.High, AgentConfirmationRequirement.StepUp), + OperationId = "create_api_key", + ArgumentsJson = """{"name":"Claude"}""", + Summary = "Create key", + OperationFingerprint = "create_api_key:{\"name\":\"Claude\"}", + Surface = AgentExecutionSurface.Chat, + ExpiresAtUtc = DateTime.UtcNow.AddMinutes(10) + }); + } + + private static AgentCapability CreateCapability( + string id, + AgentRiskClass riskClass, + AgentConfirmationRequirement confirmationRequirement) + { + return new AgentCapability( + id, + id, + id, + "test", + "scope", + riskClass, + true, + false, + confirmationRequirement); + } +} diff --git a/tests/Orbit.Infrastructure.Tests/Controllers/AiControllerTests.cs b/tests/Orbit.Infrastructure.Tests/Controllers/AiControllerTests.cs index 3171767c..f340fe82 100644 --- a/tests/Orbit.Infrastructure.Tests/Controllers/AiControllerTests.cs +++ b/tests/Orbit.Infrastructure.Tests/Controllers/AiControllerTests.cs @@ -5,6 +5,7 @@ using Microsoft.AspNetCore.Mvc; using NSubstitute; using Orbit.Api.Controllers; +using Orbit.Domain.Common; using Orbit.Domain.Interfaces; using Orbit.Domain.Models; @@ -40,6 +41,267 @@ public AiControllerTests() }; } + [Fact] + public async Task GetCapabilitiesMetadata_ReturnsForbidWhenPolicyDenies() + { + _policyEvaluator.Evaluate(Arg.Any()) + .Returns(new AgentPolicyDecision(AgentPolicyDecisionStatus.Denied, null)); + + var result = await _controller.GetCapabilitiesMetadata(CancellationToken.None); + + result.Should().BeOfType(); + await _auditService.DidNotReceiveWithAnyArgs().RecordAsync(default!, default); + } + + [Fact] + public async Task GetCapabilitiesMetadata_ReturnsCatalogAndAudits() + { + IReadOnlyList capabilities = + [ + new AgentCapability( + AgentCapabilityIds.CatalogCapabilitiesRead, + "Capabilities", + "Read capability catalog", + "catalog", + AgentScopes.CatalogRead, + AgentRiskClass.Low, + false, + false, + AgentConfirmationRequirement.None) + ]; + _policyEvaluator.Evaluate(Arg.Any()) + .Returns(new AgentPolicyDecision(AgentPolicyDecisionStatus.Allowed, null)); + _catalogService.GetCapabilities().Returns(capabilities); + + var result = await _controller.GetCapabilitiesMetadata(CancellationToken.None); + + var ok = result.Should().BeOfType().Subject; + ok.Value.Should().Be(capabilities); + await _auditService.Received(1).RecordAsync( + Arg.Is(entry => + entry.UserId == UserId && + entry.SourceName == nameof(AiController.GetCapabilitiesMetadata) && + entry.PolicyDecision == AgentPolicyDecisionStatus.Allowed), + Arg.Any()); + } + + [Fact] + public async Task GetOperationsMetadata_ReturnsCatalogAndAudits() + { + IReadOnlyList operations = + [ + new AgentOperation( + "list_habits", + "List habits", + "Read habits", + AgentCapabilityIds.HabitsRead, + AgentRiskClass.Low, + AgentConfirmationRequirement.None, + false, + true, + Parse("{}"), + Parse("{}")) + ]; + _policyEvaluator.Evaluate(Arg.Any()) + .Returns(new AgentPolicyDecision(AgentPolicyDecisionStatus.Allowed, null)); + _catalogService.GetOperations().Returns(operations); + + var result = await _controller.GetOperationsMetadata(CancellationToken.None); + + var ok = result.Should().BeOfType().Subject; + ok.Value.Should().Be(operations); + await _auditService.Received(1).RecordAsync( + Arg.Is(entry => entry.SourceName == nameof(AiController.GetOperationsMetadata)), + Arg.Any()); + } + + [Fact] + public async Task GetUserDataCatalog_ReturnsCatalogAndAudits() + { + IReadOnlyList dataCatalog = + [ + new UserDataCatalogEntry("profile", "Profile", "Profile data", "medium", "keep", true, true, []) + ]; + _policyEvaluator.Evaluate(Arg.Any()) + .Returns(new AgentPolicyDecision(AgentPolicyDecisionStatus.Allowed, null)); + _catalogService.GetUserDataCatalog().Returns(dataCatalog); + + var result = await _controller.GetUserDataCatalog(CancellationToken.None); + + var ok = result.Should().BeOfType().Subject; + ok.Value.Should().Be(dataCatalog); + await _auditService.Received(1).RecordAsync( + Arg.Is(entry => entry.SourceName == nameof(AiController.GetUserDataCatalog)), + Arg.Any()); + } + + [Fact] + public async Task GetAppSurfaces_ReturnsCatalogAndAudits() + { + IReadOnlyList surfaces = + [ + new AppSurface("today", "Today", "Today screen", [], [], [], []) + ]; + _policyEvaluator.Evaluate(Arg.Any()) + .Returns(new AgentPolicyDecision(AgentPolicyDecisionStatus.Allowed, null)); + _catalogService.GetSurfaces().Returns(surfaces); + + var result = await _controller.GetAppSurfaces(CancellationToken.None); + + var ok = result.Should().BeOfType().Subject; + ok.Value.Should().Be(surfaces); + await _auditService.Received(1).RecordAsync( + Arg.Is(entry => entry.SourceName == nameof(AiController.GetAppSurfaces)), + Arg.Any()); + } + + [Fact] + public async Task ConfirmPendingOperation_ForApiKeyUser_ReturnsForbid() + { + SetUser(isApiKey: true); + + var result = await _controller.ConfirmPendingOperation(Guid.NewGuid(), CancellationToken.None); + + result.Should().BeOfType(); + } + + [Fact] + public async Task ConfirmPendingOperation_NotFound_ReturnsNotFoundAndAudits() + { + var pendingOperationId = Guid.NewGuid(); + _pendingOperationStore.Confirm(UserId, pendingOperationId).Returns((PendingAgentOperationConfirmation?)null); + + var result = await _controller.ConfirmPendingOperation(pendingOperationId, CancellationToken.None); + + result.Should().BeOfType(); + await _auditService.Received(1).RecordAsync( + Arg.Is(entry => + entry.TargetId == pendingOperationId.ToString() && + entry.OutcomeStatus == AgentOperationStatus.Failed && + entry.Error == "pending_operation_not_found"), + Arg.Any()); + } + + [Fact] + public async Task ConfirmPendingOperation_ReturnsConfirmationToken() + { + var pendingOperationId = Guid.NewGuid(); + var confirmation = new PendingAgentOperationConfirmation( + pendingOperationId, + "agc_token", + DateTime.UtcNow.AddMinutes(5)); + _pendingOperationStore.Confirm(UserId, pendingOperationId).Returns(confirmation); + + var result = await _controller.ConfirmPendingOperation(pendingOperationId, CancellationToken.None); + + var ok = result.Should().BeOfType().Subject; + var response = ok.Value.Should().BeOfType().Subject; + response.PendingOperationId.Should().Be(pendingOperationId); + response.ConfirmationToken.Should().Be("agc_token"); + } + + [Fact] + public async Task MarkPendingOperationStepUp_ForApiKeyUser_ReturnsForbid() + { + SetUser(isApiKey: true); + + var result = await _controller.MarkPendingOperationStepUp(Guid.NewGuid(), new AiController.StepUpChallengeRequest(), CancellationToken.None); + + result.Should().BeOfType(); + } + + [Fact] + public async Task MarkPendingOperationStepUp_ReturnsBadRequestOnFailure() + { + var pendingOperationId = Guid.NewGuid(); + _stepUpService.IssueChallengeAsync(UserId, pendingOperationId, "pt-BR", Arg.Any()) + .Returns(Result.Failure("step_up_failed")); + + var result = await _controller.MarkPendingOperationStepUp( + pendingOperationId, + new AiController.StepUpChallengeRequest("pt-BR"), + CancellationToken.None); + + result.Should().BeOfType(); + await _auditService.Received(1).RecordAsync( + Arg.Is(entry => entry.Error == "step_up_failed"), + Arg.Any()); + } + + [Fact] + public async Task MarkPendingOperationStepUp_ReturnsChallenge() + { + var pendingOperationId = Guid.NewGuid(); + var challenge = new AgentStepUpChallenge(Guid.NewGuid(), pendingOperationId, DateTime.UtcNow.AddMinutes(10)); + _stepUpService.IssueChallengeAsync(UserId, pendingOperationId, "en", Arg.Any()) + .Returns(Result.Success(challenge)); + + var result = await _controller.MarkPendingOperationStepUp( + pendingOperationId, + new AiController.StepUpChallengeRequest(), + CancellationToken.None); + + var ok = result.Should().BeOfType().Subject; + ok.Value.Should().Be(challenge); + } + + [Fact] + public async Task VerifyPendingOperationStepUp_ForApiKeyUser_ReturnsForbid() + { + SetUser(isApiKey: true); + + var result = await _controller.VerifyPendingOperationStepUp( + Guid.NewGuid(), + new AiController.VerifyStepUpRequest(Guid.NewGuid(), "123456"), + CancellationToken.None); + + result.Should().BeOfType(); + } + + [Fact] + public async Task VerifyPendingOperationStepUp_ReturnsBadRequestOnFailure() + { + var pendingOperationId = Guid.NewGuid(); + var challengeId = Guid.NewGuid(); + _stepUpService.VerifyChallengeAsync(UserId, pendingOperationId, challengeId, "123456", Arg.Any()) + .Returns(Result.Failure("verify_failed")); + + var result = await _controller.VerifyPendingOperationStepUp( + pendingOperationId, + new AiController.VerifyStepUpRequest(challengeId, "123456"), + CancellationToken.None); + + result.Should().BeOfType(); + await _auditService.Received(1).RecordAsync( + Arg.Is(entry => entry.Error == "verify_failed"), + Arg.Any()); + } + + [Fact] + public async Task VerifyPendingOperationStepUp_ReturnsPendingOperation() + { + var pendingOperationId = Guid.NewGuid(); + var challengeId = Guid.NewGuid(); + var pendingOperation = new PendingAgentOperation( + pendingOperationId, + AgentCapabilityIds.HabitsDelete, + "Delete habit", + "Delete a habit", + AgentRiskClass.High, + AgentConfirmationRequirement.StepUp, + DateTime.UtcNow.AddMinutes(10)); + _stepUpService.VerifyChallengeAsync(UserId, pendingOperationId, challengeId, "123456", Arg.Any()) + .Returns(Result.Success(pendingOperation)); + + var result = await _controller.VerifyPendingOperationStepUp( + pendingOperationId, + new AiController.VerifyStepUpRequest(challengeId, "123456"), + CancellationToken.None); + + var ok = result.Should().BeOfType().Subject; + ok.Value.Should().Be(pendingOperation); + } + [Fact] public async Task ExecutePendingOperation_NotFound_ReturnsNotFound() { @@ -99,4 +361,37 @@ await _operationExecutor.Received(1).ExecuteAsync( request.Arguments.GetProperty("habit_id").GetString() == "habit-123"), Arg.Any()); } + + [Fact] + public async Task ExecutePendingOperation_ForApiKeyUser_ReturnsForbid() + { + SetUser(isApiKey: true); + + var result = await _controller.ExecutePendingOperation( + Guid.NewGuid(), + new AiController.ExecutePendingOperationRequest("agc_token"), + CancellationToken.None); + + result.Should().BeOfType(); + } + + private void SetUser(bool isApiKey = false) + { + var claims = new List { new(ClaimTypes.NameIdentifier, UserId.ToString()) }; + if (isApiKey) + claims.Add(new Claim("auth_method", "api_key")); + + _controller.ControllerContext = new ControllerContext + { + HttpContext = new DefaultHttpContext + { + User = new ClaimsPrincipal(new ClaimsIdentity(claims, "Test")) + } + }; + } + + private static JsonElement Parse(string json) + { + return JsonDocument.Parse(json).RootElement.Clone(); + } } diff --git a/tests/Orbit.Infrastructure.Tests/Mcp/AgentToolsTests.cs b/tests/Orbit.Infrastructure.Tests/Mcp/AgentToolsTests.cs new file mode 100644 index 00000000..f18d65d7 --- /dev/null +++ b/tests/Orbit.Infrastructure.Tests/Mcp/AgentToolsTests.cs @@ -0,0 +1,267 @@ +using System.Security.Claims; +using System.Text.Json; +using FluentAssertions; +using NSubstitute; +using Orbit.Api.Mcp.Tools; +using Orbit.Domain.Common; +using Orbit.Domain.Interfaces; +using Orbit.Domain.Models; + +namespace Orbit.Infrastructure.Tests.Mcp; + +public class AgentToolsTests +{ + private readonly IAgentCatalogService _catalogService = Substitute.For(); + private readonly IAgentOperationExecutor _operationExecutor = Substitute.For(); + private readonly IPendingAgentOperationStore _pendingOperationStore = Substitute.For(); + private readonly IAgentStepUpService _stepUpService = Substitute.For(); + private readonly AgentTools _tools; + private static readonly Guid UserId = Guid.NewGuid(); + + public AgentToolsTests() + { + _tools = new AgentTools(_catalogService, _operationExecutor, _pendingOperationStore, _stepUpService); + } + + [Fact] + public void ListMethods_ReturnCatalogValues() + { + var capability = new AgentCapability( + AgentCapabilityIds.HabitsRead, + "Habits", + "Read habits", + "habits", + AgentScopes.ReadHabits, + AgentRiskClass.Low, + false, + false, + AgentConfirmationRequirement.None); + var operation = new AgentOperation( + "list_habits", + "List habits", + "Read habits", + AgentCapabilityIds.HabitsRead, + AgentRiskClass.Low, + AgentConfirmationRequirement.None, + false, + true, + Parse("{}"), + Parse("{}")); + var surface = new AppSurface("today", "Today", "Today view", [], [], [], []); + var dataCatalog = new UserDataCatalogEntry("habits", "Habits", "Habit data", "low", "keep", true, true, []); + + _catalogService.GetCapabilities().Returns([capability]); + _catalogService.GetOperations().Returns([operation]); + _catalogService.GetSurfaces().Returns([surface]); + _catalogService.GetUserDataCatalog().Returns([dataCatalog]); + + _tools.ListCapabilities().Should().ContainSingle().Which.Should().Be(capability); + _tools.ListOperations().Should().ContainSingle().Which.Should().Be(operation); + _tools.ListAppSurfaces().Should().ContainSingle().Which.Should().Be(surface); + _tools.ListUserDataCatalog().Should().ContainSingle().Which.Should().Be(dataCatalog); + } + + [Fact] + public async Task ExecuteAgentOperation_UsesDefaultArgumentsAndUserContext() + { + var user = CreateUser(); + var response = new AgentExecuteOperationResponse(new AgentOperationResult( + "list_habits", + "list_habits", + AgentRiskClass.Low, + AgentConfirmationRequirement.None, + AgentOperationStatus.Succeeded)); + _operationExecutor.ExecuteAsync(Arg.Any(), Arg.Any()) + .Returns(response); + + var result = await _tools.ExecuteAgentOperation(user, "list_habits", null, "token-123", CancellationToken.None); + + result.Should().Be(response); + await _operationExecutor.Received(1).ExecuteAsync( + Arg.Is(request => + request.UserId == UserId && + request.OperationId == "list_habits" && + request.Surface == AgentExecutionSurface.Mcp && + request.AuthMethod == AgentAuthMethod.Jwt && + request.ConfirmationToken == "token-123" && + request.Arguments.ValueKind == JsonValueKind.Object && + !request.Arguments.EnumerateObject().Any()), + Arg.Any()); + } + + [Fact] + public void ConfirmAgentOperation_ThrowsForApiKeyCredentials() + { + var user = CreateUser(isApiKey: true); + + var act = () => _tools.ConfirmAgentOperation(user, Guid.NewGuid().ToString()); + + act.Should().Throw() + .WithMessage("API key credentials cannot confirm pending operations."); + } + + [Fact] + public void ConfirmAgentOperation_ThrowsForInvalidGuid() + { + var act = () => _tools.ConfirmAgentOperation(CreateUser(), "bad-guid"); + + act.Should().Throw() + .WithMessage("*pendingOperationId must be a valid GUID.*"); + } + + [Fact] + public void ConfirmAgentOperation_ReturnsStoreResult() + { + var pendingOperationId = Guid.NewGuid(); + var confirmation = new PendingAgentOperationConfirmation( + pendingOperationId, + "agc_token", + DateTime.UtcNow.AddMinutes(5)); + _pendingOperationStore.Confirm(UserId, pendingOperationId).Returns(confirmation); + + var result = _tools.ConfirmAgentOperation(CreateUser(), pendingOperationId.ToString()); + + result.Should().Be(confirmation); + } + + [Fact] + public async Task StepUpAgentOperation_ThrowsForApiKeyCredentials() + { + var act = () => _tools.StepUpAgentOperation(CreateUser(isApiKey: true), Guid.NewGuid().ToString(), cancellationToken: CancellationToken.None); + + var assertions = await act.Should().ThrowAsync(); + assertions.WithMessage("API key credentials cannot satisfy step-up authorization."); + } + + [Fact] + public async Task StepUpAgentOperation_ThrowsForInvalidGuid() + { + var act = () => _tools.StepUpAgentOperation(CreateUser(), "bad-guid", cancellationToken: CancellationToken.None); + + var assertions = await act.Should().ThrowAsync(); + assertions.WithMessage("*pendingOperationId must be a valid GUID.*"); + } + + [Fact] + public async Task StepUpAgentOperation_ReturnsChallenge() + { + var pendingOperationId = Guid.NewGuid(); + var challenge = new AgentStepUpChallenge(Guid.NewGuid(), pendingOperationId, DateTime.UtcNow.AddMinutes(10)); + _stepUpService.IssueChallengeAsync(UserId, pendingOperationId, "pt-BR", Arg.Any()) + .Returns(Result.Success(challenge)); + + var result = await _tools.StepUpAgentOperation(CreateUser(), pendingOperationId.ToString(), "pt-BR", CancellationToken.None); + + result.Should().Be(challenge); + } + + [Fact] + public async Task StepUpAgentOperation_ThrowsWhenChallengeCannotBeIssued() + { + var pendingOperationId = Guid.NewGuid(); + _stepUpService.IssueChallengeAsync(UserId, pendingOperationId, "en", Arg.Any()) + .Returns(Result.Failure("issue_failed")); + + var act = () => _tools.StepUpAgentOperation(CreateUser(), pendingOperationId.ToString(), cancellationToken: CancellationToken.None); + + await act.Should().ThrowAsync() + .WithMessage("issue_failed"); + } + + [Fact] + public async Task VerifyStepUpAgentOperation_ThrowsForApiKeyCredentials() + { + var act = () => _tools.VerifyStepUpAgentOperation( + CreateUser(isApiKey: true), + Guid.NewGuid().ToString(), + Guid.NewGuid().ToString(), + "123456", + CancellationToken.None); + + var assertions = await act.Should().ThrowAsync(); + assertions.WithMessage("API key credentials cannot satisfy step-up authorization."); + } + + [Fact] + public async Task VerifyStepUpAgentOperation_ThrowsForInvalidPendingOperationId() + { + var act = () => _tools.VerifyStepUpAgentOperation(CreateUser(), "bad-guid", Guid.NewGuid().ToString(), "123456", CancellationToken.None); + + var assertions = await act.Should().ThrowAsync(); + assertions.WithMessage("*pendingOperationId must be a valid GUID.*"); + } + + [Fact] + public async Task VerifyStepUpAgentOperation_ThrowsForInvalidChallengeId() + { + var act = () => _tools.VerifyStepUpAgentOperation(CreateUser(), Guid.NewGuid().ToString(), "bad-guid", "123456", CancellationToken.None); + + var assertions = await act.Should().ThrowAsync(); + assertions.WithMessage("*challengeId must be a valid GUID.*"); + } + + [Fact] + public async Task VerifyStepUpAgentOperation_ReturnsPendingOperation() + { + var pendingOperationId = Guid.NewGuid(); + var challengeId = Guid.NewGuid(); + var pendingOperation = new PendingAgentOperation( + pendingOperationId, + AgentCapabilityIds.HabitsDelete, + "Delete habit", + "Delete a habit", + AgentRiskClass.High, + AgentConfirmationRequirement.StepUp, + DateTime.UtcNow.AddMinutes(10)); + _stepUpService.VerifyChallengeAsync(UserId, pendingOperationId, challengeId, "123456", Arg.Any()) + .Returns(Result.Success(pendingOperation)); + + var result = await _tools.VerifyStepUpAgentOperation( + CreateUser(), + pendingOperationId.ToString(), + challengeId.ToString(), + "123456", + CancellationToken.None); + + result.Should().Be(pendingOperation); + } + + [Fact] + public async Task VerifyStepUpAgentOperation_ThrowsWhenVerificationFails() + { + var pendingOperationId = Guid.NewGuid(); + var challengeId = Guid.NewGuid(); + _stepUpService.VerifyChallengeAsync(UserId, pendingOperationId, challengeId, "123456", Arg.Any()) + .Returns(Result.Failure("verify_failed")); + + var act = () => _tools.VerifyStepUpAgentOperation( + CreateUser(), + pendingOperationId.ToString(), + challengeId.ToString(), + "123456", + CancellationToken.None); + + await act.Should().ThrowAsync() + .WithMessage("verify_failed"); + } + + private static ClaimsPrincipal CreateUser(bool isApiKey = false) + { + var claims = new List + { + new(ClaimTypes.NameIdentifier, UserId.ToString()), + new("scope", AgentScopes.ReadHabits), + new("scope", AgentScopes.WriteHabits) + }; + + if (isApiKey) + claims.Add(new Claim("auth_method", "api_key")); + + return new ClaimsPrincipal(new ClaimsIdentity(claims, "Test")); + } + + private static JsonElement Parse(string json) + { + return JsonDocument.Parse(json).RootElement.Clone(); + } +} diff --git a/tests/Orbit.Infrastructure.Tests/Services/AgentExecutionAndSanitizerTests.cs b/tests/Orbit.Infrastructure.Tests/Services/AgentExecutionAndSanitizerTests.cs new file mode 100644 index 00000000..8ca1b3cb --- /dev/null +++ b/tests/Orbit.Infrastructure.Tests/Services/AgentExecutionAndSanitizerTests.cs @@ -0,0 +1,328 @@ +using System.Text.Json; +using FluentAssertions; +using NSubstitute; +using Orbit.Application.Chat.Tools; +using Orbit.Domain.Interfaces; +using Orbit.Domain.Models; +using Orbit.Infrastructure.Services; +using Orbit.Infrastructure.Services.Prompts; + +namespace Orbit.Infrastructure.Tests.Services; + +public class AgentExecutionAndSanitizerTests +{ + private static readonly JsonElement EmptySchema = JsonDocument.Parse("{}").RootElement.Clone(); + private static readonly Guid UserId = Guid.NewGuid(); + + [Fact] + public void PromptDataSanitizer_SanitizeInline_CollapsesWhitespace() + { + var result = PromptDataSanitizer.SanitizeInline(" hello\t\r\nworld "); + + result.Should().Be("hello world"); + } + + [Fact] + public void PromptDataSanitizer_SanitizeBlock_PreservesSingleNewlines() + { + var result = PromptDataSanitizer.SanitizeBlock("first\r\n\r\nsecond\tthird"); + + result.Should().Be("first\nsecond third"); + } + + [Fact] + public void PromptDataSanitizer_QuoteInline_EscapesQuotesAndSlashes() + { + var result = PromptDataSanitizer.QuoteInline("say \"hi\" \\ now"); + + result.Should().Be("\"say \\\"hi\\\" \\\\ now\""); + } + + [Fact] + public void PromptDataSanitizer_SanitizeInline_TruncatesLongValues() + { + var result = PromptDataSanitizer.SanitizeInline("abcdefghijklmnopqrstuvwxyz", 10); + + result.Should().Be("abcdefg..."); + } + + [Fact] + public async Task AgentOperationExecutor_ReturnsUnsupportedWhenOperationIsMissing() + { + var catalog = Substitute.For(); + catalog.GetOperation("missing").Returns((AgentOperation?)null); + var executor = CreateExecutor(catalog); + + var response = await executor.ExecuteAsync(new AgentExecuteOperationRequest( + UserId, + "missing", + Parse("""{"id":"1"}"""), + AgentExecutionSurface.Chat, + AgentAuthMethod.Jwt)); + + response.Operation.Status.Should().Be(AgentOperationStatus.UnsupportedByPolicy); + response.PolicyDenial!.Reason.Should().Be("unsupported_by_policy"); + } + + [Fact] + public async Task AgentOperationExecutor_DeniesDirectUserFlowOperations() + { + var catalog = Substitute.For(); + var capability = CreateCapability(AgentCapabilityIds.AuthManage, AgentScopes.ManageAuth, AgentRiskClass.Low, AgentConfirmationRequirement.None, isMutation: true); + var operation = CreateOperation("send_auth_code", capability.Id, isMutation: true, isAgentExecutable: false, AgentConfirmationRequirement.None, AgentRiskClass.Low); + catalog.GetOperation(operation.Id).Returns(operation); + catalog.GetCapability(capability.Id).Returns(capability); + var auditService = Substitute.For(); + var executor = CreateExecutor(catalog, auditService: auditService); + + var response = await executor.ExecuteAsync(new AgentExecuteOperationRequest( + UserId, + operation.Id, + default, + AgentExecutionSurface.Chat, + AgentAuthMethod.Jwt)); + + response.Operation.Status.Should().Be(AgentOperationStatus.Denied); + response.PolicyDenial!.Reason.Should().Be("direct_user_flow_required"); + await auditService.Received(1).RecordAsync( + Arg.Is(entry => + entry.CapabilityId == capability.Id && + entry.OutcomeStatus == AgentOperationStatus.Denied && + entry.Error == "direct_user_flow_required"), + Arg.Any()); + } + + [Fact] + public async Task AgentOperationExecutor_DeniesWhenTargetIsNotOwned() + { + var catalog = Substitute.For(); + var capability = CreateCapability(AgentCapabilityIds.HabitsDelete, AgentScopes.DeleteHabits, AgentRiskClass.Destructive, AgentConfirmationRequirement.None, isMutation: true); + var operation = CreateOperation("delete_habit", capability.Id, isMutation: true, isAgentExecutable: true, AgentConfirmationRequirement.None, AgentRiskClass.Destructive); + catalog.GetOperation(operation.Id).Returns(operation); + catalog.GetCapability(capability.Id).Returns(capability); + var ownership = Substitute.For(); + ownership.GetDenialReasonAsync(operation.Id, UserId, Arg.Any(), Arg.Any()) + .Returns("target_not_owned:delete_habit:habit"); + var executor = CreateExecutor(catalog, ownershipService: ownership); + + var response = await executor.ExecuteAsync(new AgentExecuteOperationRequest( + UserId, + operation.Id, + Parse("""{"habit_id":"123"}"""), + AgentExecutionSurface.Chat, + AgentAuthMethod.Jwt)); + + response.Operation.Status.Should().Be(AgentOperationStatus.Denied); + response.PolicyDenial!.Reason.Should().Be("target_not_owned:delete_habit:habit"); + } + + [Fact] + public async Task AgentOperationExecutor_PropagatesPolicyDenial() + { + var catalog = Substitute.For(); + var capability = CreateCapability(AgentCapabilityIds.HabitsWrite, AgentScopes.WriteHabits, AgentRiskClass.Low, AgentConfirmationRequirement.None, isMutation: true); + var operation = CreateOperation("create_habit", capability.Id, isMutation: true, isAgentExecutable: true, AgentConfirmationRequirement.None, AgentRiskClass.Low); + catalog.GetOperation(operation.Id).Returns(operation); + catalog.GetCapability(capability.Id).Returns(capability); + var policy = Substitute.For(); + policy.Evaluate(Arg.Any()) + .Returns(new AgentPolicyDecision(AgentPolicyDecisionStatus.Denied, capability, "missing_scope:write_habits")); + var executor = CreateExecutor(catalog, policyEvaluator: policy); + + var response = await executor.ExecuteAsync(new AgentExecuteOperationRequest( + UserId, + operation.Id, + Parse("""{"title":"Test"}"""), + AgentExecutionSurface.Chat, + AgentAuthMethod.ApiKey, + ["read_habits"])); + + response.Operation.Status.Should().Be(AgentOperationStatus.Denied); + response.PolicyDenial!.Reason.Should().Be("missing_scope:write_habits"); + } + + [Fact] + public async Task AgentOperationExecutor_ReturnsPendingConfirmation() + { + var catalog = Substitute.For(); + var capability = CreateCapability(AgentCapabilityIds.HabitsDelete, AgentScopes.DeleteHabits, AgentRiskClass.Destructive, AgentConfirmationRequirement.FreshConfirmation, isMutation: true); + var operation = CreateOperation("delete_habit", capability.Id, isMutation: true, isAgentExecutable: true, AgentConfirmationRequirement.FreshConfirmation, AgentRiskClass.Destructive); + catalog.GetOperation(operation.Id).Returns(operation); + catalog.GetCapability(capability.Id).Returns(capability); + var pending = new PendingAgentOperation(Guid.NewGuid(), capability.Id, capability.DisplayName, "Delete habit", capability.RiskClass, capability.ConfirmationRequirement, DateTime.UtcNow.AddMinutes(5)); + var policy = Substitute.For(); + policy.Evaluate(Arg.Any()) + .Returns(new AgentPolicyDecision(AgentPolicyDecisionStatus.ConfirmationRequired, capability, "confirmation_required", pending)); + var executor = CreateExecutor(catalog, policyEvaluator: policy); + + var response = await executor.ExecuteAsync(new AgentExecuteOperationRequest( + UserId, + operation.Id, + Parse("""{"habit_id":"123"}"""), + AgentExecutionSurface.Chat, + AgentAuthMethod.Jwt)); + + response.Operation.Status.Should().Be(AgentOperationStatus.PendingConfirmation); + response.PendingOperation.Should().Be(pending); + } + + [Fact] + public async Task AgentOperationExecutor_UsesFallbackScopesForJwtRequests() + { + var catalog = Substitute.For(); + var capability = CreateCapability(AgentCapabilityIds.HabitsWrite, AgentScopes.WriteHabits, AgentRiskClass.Low, AgentConfirmationRequirement.None, isMutation: true); + var operation = CreateOperation("create_habit", capability.Id, isMutation: true, isAgentExecutable: true, AgentConfirmationRequirement.None, AgentRiskClass.Low); + catalog.GetOperation(operation.Id).Returns(operation); + catalog.GetCapability(capability.Id).Returns(capability); + catalog.GetCapabilities().Returns([capability]); + var capturedContext = default(AgentPolicyEvaluationContext); + var policy = Substitute.For(); + policy.Evaluate(Arg.Do(context => capturedContext = context)) + .Returns(new AgentPolicyDecision(AgentPolicyDecisionStatus.Allowed, capability)); + var executor = CreateExecutor( + catalog, + policyEvaluator: policy, + toolRegistry: new AiToolRegistry([new StubTool(operation.Id, (_, _, _) => Task.FromResult(new ToolResult(true, EntityId: "1", EntityName: "Habit")))])); + + var response = await executor.ExecuteAsync(new AgentExecuteOperationRequest( + UserId, + operation.Id, + Parse("""{"title":"Test"}"""), + AgentExecutionSurface.Chat, + AgentAuthMethod.Jwt)); + + response.Operation.Status.Should().Be(AgentOperationStatus.Succeeded); + capturedContext.Should().NotBeNull(); + capturedContext.GrantedScopes.Should().Contain(AgentScopes.WriteHabits); + } + + [Fact] + public async Task AgentOperationExecutor_ReturnsToolFailure() + { + var catalog = Substitute.For(); + var capability = CreateCapability(AgentCapabilityIds.HabitsWrite, AgentScopes.WriteHabits, AgentRiskClass.Low, AgentConfirmationRequirement.None, isMutation: true); + var operation = CreateOperation("create_habit", capability.Id, isMutation: true, isAgentExecutable: true, AgentConfirmationRequirement.None, AgentRiskClass.Low); + catalog.GetOperation(operation.Id).Returns(operation); + catalog.GetCapability(capability.Id).Returns(capability); + var policy = Substitute.For(); + policy.Evaluate(Arg.Any()) + .Returns(new AgentPolicyDecision(AgentPolicyDecisionStatus.Allowed, capability)); + var executor = CreateExecutor( + catalog, + policyEvaluator: policy, + toolRegistry: new AiToolRegistry([new StubTool(operation.Id, (_, _, _) => Task.FromResult(new ToolResult(false, EntityId: UserId.ToString(), Error: "tool_failed")))])); + + var response = await executor.ExecuteAsync(new AgentExecuteOperationRequest( + UserId, + operation.Id, + Parse("""{"title":"Test"}"""), + AgentExecutionSurface.Chat, + AgentAuthMethod.Jwt)); + + response.Operation.Status.Should().Be(AgentOperationStatus.Failed); + response.Operation.PolicyReason.Should().Be("tool_failed"); + } + + [Fact] + public async Task AgentOperationExecutor_MapsUnexpectedExceptions() + { + var catalog = Substitute.For(); + var capability = CreateCapability(AgentCapabilityIds.HabitsWrite, AgentScopes.WriteHabits, AgentRiskClass.Low, AgentConfirmationRequirement.None, isMutation: true); + var operation = CreateOperation("create_habit", capability.Id, isMutation: true, isAgentExecutable: true, AgentConfirmationRequirement.None, AgentRiskClass.Low); + catalog.GetOperation(operation.Id).Returns(operation); + catalog.GetCapability(capability.Id).Returns(capability); + var policy = Substitute.For(); + policy.Evaluate(Arg.Any()) + .Returns(new AgentPolicyDecision(AgentPolicyDecisionStatus.Allowed, capability)); + var executor = CreateExecutor( + catalog, + policyEvaluator: policy, + toolRegistry: new AiToolRegistry([new StubTool(operation.Id, (_, _, _) => throw new InvalidOperationException("boom"))])); + + var response = await executor.ExecuteAsync(new AgentExecuteOperationRequest( + UserId, + operation.Id, + Parse("""{"title":"Test"}"""), + AgentExecutionSurface.Chat, + AgentAuthMethod.Jwt)); + + response.Operation.Status.Should().Be(AgentOperationStatus.Failed); + response.Operation.PolicyReason.Should().Be("unexpected_error"); + } + + private static AgentOperationExecutor CreateExecutor( + IAgentCatalogService catalogService, + IAgentPolicyEvaluator? policyEvaluator = null, + IAgentAuditService? auditService = null, + IAgentTargetOwnershipService? ownershipService = null, + AiToolRegistry? toolRegistry = null) + { + if (ownershipService is null) + { + ownershipService = Substitute.For(); + ownershipService.GetDenialReasonAsync(Arg.Any(), Arg.Any(), Arg.Any(), Arg.Any()) + .Returns((string?)null); + } + + return new AgentOperationExecutor( + catalogService, + policyEvaluator ?? Substitute.For(), + auditService ?? Substitute.For(), + ownershipService, + toolRegistry ?? new AiToolRegistry([])); + } + + private static AgentCapability CreateCapability( + string id, + string scope, + AgentRiskClass riskClass, + AgentConfirmationRequirement confirmationRequirement, + bool isMutation) + { + return new AgentCapability( + id, + id, + id, + "test", + scope, + riskClass, + isMutation, + false, + confirmationRequirement); + } + + private static AgentOperation CreateOperation( + string id, + string capabilityId, + bool isMutation, + bool isAgentExecutable, + AgentConfirmationRequirement confirmationRequirement, + AgentRiskClass riskClass) + { + return new AgentOperation( + id, + id, + id, + capabilityId, + riskClass, + confirmationRequirement, + isMutation, + isAgentExecutable, + EmptySchema, + EmptySchema); + } + + private static JsonElement Parse(string json) + { + return JsonDocument.Parse(json).RootElement.Clone(); + } + + private sealed class StubTool(string name, Func> executeAsync) : IAiTool + { + public string Name => name; + public string Description => name; + public bool IsReadOnly => false; + public object GetParameterSchema() => new { }; + public Task ExecuteAsync(JsonElement args, Guid userId, CancellationToken ct) => executeAsync(args, userId, ct); + } +} diff --git a/tests/Orbit.Infrastructure.Tests/Services/FeatureFlagAndAgentSupportTests.cs b/tests/Orbit.Infrastructure.Tests/Services/FeatureFlagAndAgentSupportTests.cs new file mode 100644 index 00000000..d73636cb --- /dev/null +++ b/tests/Orbit.Infrastructure.Tests/Services/FeatureFlagAndAgentSupportTests.cs @@ -0,0 +1,197 @@ +using System.Net; +using System.Security.Claims; +using System.Text.Json; +using FluentAssertions; +using Microsoft.AspNetCore.Http; +using Microsoft.AspNetCore.Mvc; +using Microsoft.AspNetCore.Mvc.Abstractions; +using Microsoft.AspNetCore.Mvc.Filters; +using Microsoft.AspNetCore.Routing; +using Microsoft.EntityFrameworkCore; +using NSubstitute; +using Orbit.Api.RateLimiting; +using Orbit.Domain.Entities; +using Orbit.Domain.Enums; +using Orbit.Domain.Interfaces; +using Orbit.Domain.Models; +using Orbit.Infrastructure.Persistence; +using Orbit.Infrastructure.Services; + +namespace Orbit.Infrastructure.Tests.Services; + +public class FeatureFlagAndAgentSupportTests : IDisposable +{ + private readonly OrbitDbContext _dbContext; + + public FeatureFlagAndAgentSupportTests() + { + var options = new DbContextOptionsBuilder() + .UseInMemoryDatabase($"FeatureFlagAndAgentSupportTests_{Guid.NewGuid()}") + .Options; + + _dbContext = new OrbitDbContext(options); + } + + public void Dispose() + { + _dbContext.Dispose(); + GC.SuppressFinalize(this); + } + + [Fact] + public async Task FeatureFlagService_ReturnsEnabledFlagsForMatchingPlan() + { + var user = User.Create("Thomas", "thomas@example.com").Value; + user.SetStripeSubscription("sub_123", DateTime.UtcNow.AddDays(30), SubscriptionInterval.Yearly); + _dbContext.Users.Add(user); + _dbContext.AppFeatureFlags.Add(AppFeatureFlag.Create("basic", true, null, "Basic")); + _dbContext.AppFeatureFlags.Add(AppFeatureFlag.Create("pro_only", true, "pro", "Pro")); + _dbContext.AppFeatureFlags.Add(AppFeatureFlag.Create("disabled", false, null, "Disabled")); + await _dbContext.SaveChangesAsync(); + + var service = new FeatureFlagService(_dbContext); + + var result = await service.GetEnabledKeysForUserAsync(user.Id); + + result.Should().Equal("basic", "pro_only"); + } + + [Fact] + public async Task FeatureFlagService_ReturnsEmptyWhenUserIsMissing() + { + var service = new FeatureFlagService(_dbContext); + + var result = await service.GetEnabledKeysForUserAsync(Guid.NewGuid()); + + result.Should().BeEmpty(); + } + + [Fact] + public async Task AgentTargetOwnershipService_ReturnsNullWhenNoTargetsAreProvided() + { + var service = new AgentTargetOwnershipService(_dbContext); + + var result = await service.GetDenialReasonAsync("delete_habit", Guid.NewGuid(), Parse("{}")); + + result.Should().BeNull(); + } + + [Fact] + public async Task AgentTargetOwnershipService_ReturnsDenialWhenHabitIsNotOwned() + { + var owner = User.Create("Owner", "owner@example.com").Value; + var otherUser = User.Create("Other", "other@example.com").Value; + var habit = Habit.Create(new HabitCreateParams( + owner.Id, + "Exercise", + FrequencyUnit.Day, + 1, + DueDate: DateOnly.FromDateTime(DateTime.UtcNow))).Value; + + _dbContext.Users.AddRange(owner, otherUser); + _dbContext.Habits.Add(habit); + await _dbContext.SaveChangesAsync(); + + var service = new AgentTargetOwnershipService(_dbContext); + + var result = await service.GetDenialReasonAsync( + "delete_habit", + otherUser.Id, + Parse($$"""{"habit_id":"{{habit.Id}}"}""")); + + result.Should().Be("target_not_owned:delete_habit:habit"); + } + + [Fact] + public async Task AgentAuditService_PersistsAuditLog() + { + var service = new AgentAuditService(_dbContext); + var entry = new AgentAuditEntry( + Guid.NewGuid(), + AgentCapabilityIds.HabitsRead, + "list_habits", + AgentExecutionSurface.Mcp, + AgentAuthMethod.Jwt, + AgentRiskClass.Low, + AgentPolicyDecisionStatus.Allowed, + AgentOperationStatus.Succeeded, + Summary: "Read habits"); + + await service.RecordAsync(entry); + + _dbContext.AgentAuditLogs.Should().ContainSingle(log => + log.CapabilityId == AgentCapabilityIds.HabitsRead && + log.SourceName == "list_habits"); + } + + [Fact] + public void DistributedRateLimitAttribute_CreatesFilterWithResolvedService() + { + var rateLimitService = Substitute.For(); + var serviceProvider = Substitute.For(); + serviceProvider.GetService(typeof(IDistributedRateLimitService)).Returns(rateLimitService); + var attribute = new DistributedRateLimitAttribute("chat"); + + var filter = attribute.CreateInstance(serviceProvider); + + filter.Should().BeOfType(); + } + + [Fact] + public async Task DistributedRateLimitFilter_UsesAuthenticatedUserPartitionAndCallsNext() + { + var userId = Guid.NewGuid(); + var rateLimitService = Substitute.For(); + rateLimitService.TryAcquireAsync("chat", $"user:{userId}", Arg.Any()) + .Returns(new DistributedRateLimitDecision(true, 20, 1, DateTime.UtcNow.AddSeconds(30))); + var filter = new DistributedRateLimitFilter("chat", rateLimitService); + var context = CreateActionExecutingContext(new DefaultHttpContext + { + User = new ClaimsPrincipal(new ClaimsIdentity( + [ + new Claim(ClaimTypes.NameIdentifier, userId.ToString()) + ], "Test")) + }); + var nextCalled = false; + + await filter.OnActionExecutionAsync(context, () => + { + nextCalled = true; + return Task.FromResult(new ActionExecutedContext(context, [], new object())); + }); + + nextCalled.Should().BeTrue(); + context.Result.Should().BeNull(); + } + + [Fact] + public async Task DistributedRateLimitFilter_ReturnsTooManyRequestsForAnonymousIp() + { + var rateLimitService = Substitute.For(); + var retryAt = DateTime.UtcNow.AddSeconds(15); + rateLimitService.TryAcquireAsync("auth", "ip:203.0.113.10", Arg.Any()) + .Returns(new DistributedRateLimitDecision(false, 5, 5, retryAt)); + var filter = new DistributedRateLimitFilter("auth", rateLimitService); + var httpContext = new DefaultHttpContext(); + httpContext.Connection.RemoteIpAddress = IPAddress.Parse("203.0.113.10"); + var context = CreateActionExecutingContext(httpContext); + + await filter.OnActionExecutionAsync(context, () => + Task.FromResult(new ActionExecutedContext(context, [], new object()))); + + var result = context.Result.Should().BeOfType().Subject; + result.StatusCode.Should().Be(StatusCodes.Status429TooManyRequests); + httpContext.Response.Headers.RetryAfter.Should().NotBeEmpty(); + } + + private static ActionExecutingContext CreateActionExecutingContext(HttpContext httpContext) + { + var actionContext = new ActionContext(httpContext, new RouteData(), new ActionDescriptor()); + return new ActionExecutingContext(actionContext, [], new Dictionary(), new object()); + } + + private static JsonElement Parse(string json) + { + return JsonDocument.Parse(json).RootElement.Clone(); + } +}