From 414c758cae08058d1a18dd360c345eb6ddec0881 Mon Sep 17 00:00:00 2001 From: Aaron Burtle Date: Fri, 28 Aug 2026 06:07:06 -0700 Subject: [PATCH 01/19] expand code coverage --- .../AuthorizationResolverUnitTests.cs | 7 + .../REST/RestAuthorizationHandlerUnitTests.cs | 135 ++++ .../GraphQLBuilder/GraphQLUtilsTests.cs | 83 +++ .../Mcp/AggregateRecordsToolTests.cs | 453 ++++++++++++ .../Mcp/CreateRecordToolUnitTests.cs | 306 ++++++++ .../Mcp/DeleteRecordToolUnitTests.cs | 324 +++++++++ .../Mcp/DescribeEntitiesFilteringTests.cs | 42 ++ .../Mcp/DynamicCustomToolTests.cs | 392 +++++++++-- .../Mcp/ExecuteEntityToolTests.cs | 318 ++++++++- .../Mcp/McpArgumentParserTests.cs | 50 ++ .../Mcp/McpAuthorizationHelperTests.cs | 29 + .../Mcp/McpMetadataHelperTests.cs | 23 + .../Mcp/McpResponseBuilderTests.cs | 10 + .../Mcp/McpServerConfigurationTests.cs | 34 + .../ApplicationNameTelemetryTests.cs | 81 +++ .../AutoentityConverterCoverageTests.cs | 64 ++ .../BaseSqlQueryStructureHelperTests.cs | 228 ++++++ .../ConfigObjectModelCoverageTests.cs | 223 ++++++ .../ConfigureJwtBearerOptionsTests.cs | 88 +++ .../UnitTests/CosmosEngineHelperTests.cs | 183 +++++ .../CosmosSqlMetadataProviderHelperTests.cs | 238 +++++++ .../UnitTests/DatabaseObjectUnitTests.cs | 167 +++++ ...izationVariableReplacementSettingsTests.cs | 32 + .../UnitTests/DmlToolsConfigConverterTests.cs | 12 + .../UnitTests/DwSqlQueryBuilderHelperTests.cs | 172 +++++ .../EmbeddingTelemetryHelperTests.cs | 119 ++++ .../EmbeddingsOptionsConverterTests.cs | 51 ++ .../EntityApiOptionsConverterCoverageTests.cs | 137 ++++ .../EntityHealthOptionsConverterTests.cs | 132 ++++ .../UnitTests/EntitySourceConverterTests.cs | 108 +++ .../UnitTests/ExecutionHelperScalarTests.cs | 80 +++ .../UnitTests/GraphQLFilterParserUnitTests.cs | 93 +++ .../UnitTests/McpLogNotificationTests.cs | 86 +++ .../McpStdioServerContentBlockTests.cs | 177 ++++- .../UnitTests/McpStdioServerProtocolTests.cs | 611 ++++++++++++++++ .../UnitTests/McpTelemetryTests.cs | 43 ++ .../MultipleCreateOrderHelperEdgeTests.cs | 154 +++++ .../UnitTests/PureUtilityCoverageTests.cs | 101 +++ .../UnitTests/QueryExecutorHelperTests.cs | 432 ++++++++++++ .../UnitTests/RuntimeConfigHelperTests.cs | 214 ++++++ .../RuntimeConfigValidatorUnitTests.cs | 253 ++++++- .../RuntimeOptionsConverterCoverageTests.cs | 137 ++++ .../SqlMetadataProviderHelperTests.cs | 398 +++++++++++ .../UnitTests/SqlMutationEngineHelperTests.cs | 651 ++++++++++++++++++ .../UnitTests/SqlPaginationUtilUnitTests.cs | 122 ++++ .../UnitTests/SqlQueryEngineHelperTests.cs | 187 +++++ .../UnitTests/SqlQueryExecutorUnitTests.cs | 106 ++- .../UnitTests/SqlQueryStructureHelperTests.cs | 178 +++++ .../UnitTests/SqlQueryStructuresModelTests.cs | 53 ++ .../UnitTests/TypeHelperTests.cs | 10 + 50 files changed, 8267 insertions(+), 60 deletions(-) create mode 100644 src/Service.Tests/Mcp/CreateRecordToolUnitTests.cs create mode 100644 src/Service.Tests/Mcp/DeleteRecordToolUnitTests.cs create mode 100644 src/Service.Tests/Mcp/McpServerConfigurationTests.cs create mode 100644 src/Service.Tests/UnitTests/AutoentityConverterCoverageTests.cs create mode 100644 src/Service.Tests/UnitTests/BaseSqlQueryStructureHelperTests.cs create mode 100644 src/Service.Tests/UnitTests/ConfigObjectModelCoverageTests.cs create mode 100644 src/Service.Tests/UnitTests/ConfigureJwtBearerOptionsTests.cs create mode 100644 src/Service.Tests/UnitTests/CosmosEngineHelperTests.cs create mode 100644 src/Service.Tests/UnitTests/CosmosSqlMetadataProviderHelperTests.cs create mode 100644 src/Service.Tests/UnitTests/DatabaseObjectUnitTests.cs create mode 100644 src/Service.Tests/UnitTests/DwSqlQueryBuilderHelperTests.cs create mode 100644 src/Service.Tests/UnitTests/EmbeddingTelemetryHelperTests.cs create mode 100644 src/Service.Tests/UnitTests/EntityApiOptionsConverterCoverageTests.cs create mode 100644 src/Service.Tests/UnitTests/EntityHealthOptionsConverterTests.cs create mode 100644 src/Service.Tests/UnitTests/EntitySourceConverterTests.cs create mode 100644 src/Service.Tests/UnitTests/ExecutionHelperScalarTests.cs create mode 100644 src/Service.Tests/UnitTests/McpStdioServerProtocolTests.cs create mode 100644 src/Service.Tests/UnitTests/MultipleCreateOrderHelperEdgeTests.cs create mode 100644 src/Service.Tests/UnitTests/PureUtilityCoverageTests.cs create mode 100644 src/Service.Tests/UnitTests/QueryExecutorHelperTests.cs create mode 100644 src/Service.Tests/UnitTests/RuntimeConfigHelperTests.cs create mode 100644 src/Service.Tests/UnitTests/RuntimeOptionsConverterCoverageTests.cs create mode 100644 src/Service.Tests/UnitTests/SqlMetadataProviderHelperTests.cs create mode 100644 src/Service.Tests/UnitTests/SqlMutationEngineHelperTests.cs create mode 100644 src/Service.Tests/UnitTests/SqlQueryEngineHelperTests.cs create mode 100644 src/Service.Tests/UnitTests/SqlQueryStructureHelperTests.cs create mode 100644 src/Service.Tests/UnitTests/SqlQueryStructuresModelTests.cs diff --git a/src/Service.Tests/Authorization/AuthorizationResolverUnitTests.cs b/src/Service.Tests/Authorization/AuthorizationResolverUnitTests.cs index 1ab41a7c1a..9dc9cfd717 100644 --- a/src/Service.Tests/Authorization/AuthorizationResolverUnitTests.cs +++ b/src/Service.Tests/Authorization/AuthorizationResolverUnitTests.cs @@ -34,6 +34,13 @@ public class AuthorizationResolverUnitTests private const string TEST_AUTHENTICATION_TYPE = "TestAuth"; private const string TEST_CLAIMTYPE_NAME = "TestName"; + [TestMethod] + public void GetRolesForOperation_NullEntityNameThrows() + { + Assert.ThrowsException(() => + IAuthorizationResolver.GetRolesForOperation(null!, EntityActionOperation.Read, null)); + } + #region Role Context Tests /// /// When the client role header is present, validates result when diff --git a/src/Service.Tests/Authorization/REST/RestAuthorizationHandlerUnitTests.cs b/src/Service.Tests/Authorization/REST/RestAuthorizationHandlerUnitTests.cs index 8db19814d0..d326b08c88 100644 --- a/src/Service.Tests/Authorization/REST/RestAuthorizationHandlerUnitTests.cs +++ b/src/Service.Tests/Authorization/REST/RestAuthorizationHandlerUnitTests.cs @@ -1,9 +1,11 @@ // Copyright (c) Microsoft Corporation. // Licensed under the MIT License. +using System; using System.Collections; using System.Collections.Generic; using System.Security.Claims; +using System.Text.Json; using System.Threading.Tasks; using Azure.DataApiBuilder.Auth; using Azure.DataApiBuilder.Config.DatabasePrimitives; @@ -288,6 +290,132 @@ public async Task FindColumnPermissionsTests(string[] columnsRequestedInput, CollectionAssert.AreEquivalent(expected: (ICollection)allowedColumns, actual: stubRestRequestContext.FieldsToBeReturned, message: "FieldsToBeReturned not subset of allowed columns."); } + [TestMethod] + public async Task MultipleRequirementsAreRejected() + { + AuthorizationHandlerContext context = new( + new IAuthorizationRequirement[] { new RoleContextPermissionsRequirement(), new ColumnsPermissionsRequirement() }, + new ClaimsPrincipal(), + AuthorizationHelpers.TEST_ENTITY); + RestAuthorizationHandler handler = CreateHandler(new Mock().Object, CreateHttpContext()); + + await Assert.ThrowsExceptionAsync(() => handler.HandleAsync(context)); + } + + [TestMethod] + public async Task MissingHttpContextIsRejected() + { + AuthorizationHandlerContext context = new( + new IAuthorizationRequirement[] { new RoleContextPermissionsRequirement() }, + new ClaimsPrincipal(), + AuthorizationHelpers.TEST_ENTITY); + RestAuthorizationHandler handler = CreateHandler(new Mock().Object, null); + + await Assert.ThrowsExceptionAsync(() => handler.HandleAsync(context)); + } + + [TestMethod] + public async Task UnsupportedHttpVerbIsRejected() + { + await Assert.ThrowsExceptionAsync(() => IsAuthorizationSuccessfulAsync( + new EntityRoleOperationPermissionsRequirement(), + AuthorizationHelpers.TEST_ENTITY, + new Mock().Object, + CreateHttpContext("OPTIONS"))); + } + + [TestMethod] + public async Task DeleteColumnRequirementSucceedsWithoutColumnChecks() + { + bool result = await IsAuthorizationSuccessfulAsync( + new ColumnsPermissionsRequirement(), + CreateRestRequestContext(Array.Empty()), + new Mock().Object, + CreateHttpContext(HttpConstants.DELETE)); + + Assert.IsTrue(result); + } + + [DataTestMethod] + [DataRow(true, true)] + [DataRow(false, false)] + public async Task EmptyInsertColumnsDependOnAccessibleFields(bool hasAccessibleFields, bool expected) + { + Mock resolver = new(); + resolver.Setup(x => x.GetAllowedExposedColumns( + AuthorizationHelpers.TEST_ENTITY, + AuthorizationHelpers.TEST_ROLE, + EntityActionOperation.Create)) + .Returns(hasAccessibleFields ? new[] { "id" } : Array.Empty()); + using JsonDocument payload = JsonDocument.Parse("{}"); + RestRequestContext context = new InsertRequestContext( + AuthorizationHelpers.TEST_ENTITY, + new DatabaseTable { TableDefinition = new SourceDefinition() }, + payload.RootElement, + EntityActionOperation.Insert); + + bool result = await IsAuthorizationSuccessfulAsync( + new ColumnsPermissionsRequirement(), + context, + resolver.Object, + CreateHttpContext(HttpConstants.POST)); + + Assert.AreEqual(expected, result); + } + + [TestMethod] + public async Task InvalidColumnsRequirementResourceIsRejected() + { + await Assert.ThrowsExceptionAsync(() => IsAuthorizationSuccessfulAsync( + new ColumnsPermissionsRequirement(), + new object(), + new Mock().Object, + CreateHttpContext())); + } + + [DataTestMethod] + [DataRow(true, true)] + [DataRow(false, false)] + public async Task StoredProcedureRequirementUsesResolverDecision(bool permitted, bool expected) + { + Mock resolver = new(); + resolver.Setup(x => x.IsStoredProcedureExecutionPermitted( + AuthorizationHelpers.TEST_ENTITY, + AuthorizationHelpers.TEST_ROLE, + SupportedHttpVerb.Post)) + .Returns(permitted); + + bool result = await IsAuthorizationSuccessfulAsync( + new StoredProcedurePermissionsRequirement(), + AuthorizationHelpers.TEST_ENTITY, + resolver.Object, + CreateHttpContext(HttpConstants.POST)); + + Assert.AreEqual(expected, result); + } + + [TestMethod] + public async Task StoredProcedureRequirementFailsForNullResource() + { + bool result = await IsAuthorizationSuccessfulAsync( + new StoredProcedurePermissionsRequirement(), + null, + new Mock().Object, + CreateHttpContext(HttpConstants.POST)); + + Assert.IsFalse(result); + } + + [TestMethod] + public async Task InvalidStoredProcedureResourceIsRejected() + { + await Assert.ThrowsExceptionAsync(() => IsAuthorizationSuccessfulAsync( + new StoredProcedurePermissionsRequirement(), + new object(), + new Mock().Object, + CreateHttpContext(HttpConstants.POST))); + } + #region Helper Methods /// /// Setup request and authorization context and get Authorization result @@ -315,6 +443,13 @@ private static async Task IsAuthorizationSuccessfulAsync( return context.HasSucceeded; } + private static RestAuthorizationHandler CreateHandler(IAuthorizationResolver resolver, HttpContext? httpContext) + { + Mock accessor = new(); + accessor.Setup(x => x.HttpContext).Returns(httpContext); + return new RestAuthorizationHandler(resolver, accessor.Object, new Mock>().Object); + } + /// /// Create Mock HttpContext object for use in test fixture. /// diff --git a/src/Service.Tests/GraphQLBuilder/GraphQLUtilsTests.cs b/src/Service.Tests/GraphQLBuilder/GraphQLUtilsTests.cs index e0c9a207f6..f08c4fe125 100644 --- a/src/Service.Tests/GraphQLBuilder/GraphQLUtilsTests.cs +++ b/src/Service.Tests/GraphQLBuilder/GraphQLUtilsTests.cs @@ -1,10 +1,12 @@ // Copyright (c) Microsoft Corporation. // Licensed under the MIT License. +using System; using System.Collections.Generic; using Azure.DataApiBuilder.Config.ObjectModel; using Azure.DataApiBuilder.Service.Exceptions; using Azure.DataApiBuilder.Service.GraphQLBuilder; +using Azure.DataApiBuilder.Service.GraphQLBuilder.Directives; using HotChocolate.Language; using Microsoft.VisualStudio.TestTools.UnitTesting; @@ -149,6 +151,87 @@ public void CreateAuthorizationDirective_WithRoles_ReturnsAuthorizeDirective() Assert.AreEqual(GraphQLUtils.AUTHORIZE_DIRECTIVE_ARGUMENT_ROLES, directive.Arguments[0].Name.Value); } + [DataTestMethod] + [DataRow(SyntaxKind.IntValue, true)] + [DataRow(SyntaxKind.EnumValue, true)] + [DataRow(SyntaxKind.ObjectValue, false)] + public void IsScalarField_ReturnsExpected(SyntaxKind kind, bool expected) + { + Assert.AreEqual(expected, GraphQLUtils.IsScalarField(kind)); + } + + [TestMethod] + public void LinkingEntityName_RoundTripsSourceAndTarget() + { + string name = GraphQLUtils.GenerateLinkingEntityName("Book", "Author"); + Tuple decoded = GraphQLUtils.GetSourceAndTargetEntityNameFromLinkingEntityName(name); + + Assert.AreEqual("Book", decoded.Item1); + Assert.AreEqual("Author", decoded.Item2); + Assert.ThrowsException( + () => GraphQLUtils.GetSourceAndTargetEntityNameFromLinkingEntityName("Book$Author")); + Assert.ThrowsException( + () => GraphQLUtils.GetSourceAndTargetEntityNameFromLinkingEntityName("LinkingEntity$Book")); + } + + [TestMethod] + public void GetFieldNodeForGivenFieldName_ReturnsValueOrThrows() + { + List fields = new() { new("id", new IntValueNode(7)) }; + + Assert.AreEqual(7, ((IntValueNode)GraphQLUtils.GetFieldNodeForGivenFieldName(fields, "id")).ToInt32()); + Assert.ThrowsException(() => GraphQLUtils.GetFieldNodeForGivenFieldName(fields, "missing")); + } + + [TestMethod] + public void RelationshipHelpers_HandleConfiguredMissingAndAbsentRelationships() + { + EntityRelationship relationship = new( + Cardinality.Many, + "Author", + Array.Empty(), + Array.Empty(), + "book_author", + Array.Empty(), + Array.Empty()); + Entity entity = CreateEntity(new Dictionary { ["authors"] = relationship }); + + Assert.IsTrue(GraphQLUtils.IsMToNRelationship(entity, "authors")); + Assert.IsFalse(GraphQLUtils.IsMToNRelationship(entity, "missing")); + Assert.AreEqual("Author", GraphQLUtils.GetRelationshipTargetEntityName(entity, "Book", "authors")); + Assert.ThrowsException( + () => GraphQLUtils.GetRelationshipTargetEntityName(entity, "Book", "missing")); + Assert.ThrowsException( + () => GraphQLUtils.GetRelationshipTargetEntityName(CreateEntity(null), "Book", "authors")); + } + + [TestMethod] + public void RelationshipDirective_ExtractsTargetCardinalityAndDirective() + { + FieldDefinitionNode field = ParseObjectType( + "type Book { author: Author @relationship(target: \"Author\", cardinality: \"many\") }").Fields[0]; + + Assert.AreEqual("Author", RelationshipDirectiveType.Target(field)); + Assert.AreEqual(Cardinality.Many, RelationshipDirectiveType.Cardinality(field)); + Assert.IsNotNull(RelationshipDirectiveType.GetDirective(field)); + + FieldDefinitionNode plainField = ParseObjectType("type Book { author: Author }").Fields[0]; + Assert.AreEqual("Author", RelationshipDirectiveType.Target(plainField)); + Assert.ThrowsException(() => RelationshipDirectiveType.Cardinality(plainField)); + } + + private static Entity CreateEntity(Dictionary? relationships) + { + return new Entity( + Source: new EntitySource("books", EntitySourceType.Table, null, null), + GraphQL: new EntityGraphQLOptions("Book", "Books"), + Fields: null, + Rest: new EntityRestOptions(), + Permissions: Array.Empty(), + Mappings: null, + Relationships: relationships); + } + private static ObjectTypeDefinitionNode ParseObjectType(string sdl) { DocumentNode document = Utf8GraphQLParser.Parse(sdl); diff --git a/src/Service.Tests/Mcp/AggregateRecordsToolTests.cs b/src/Service.Tests/Mcp/AggregateRecordsToolTests.cs index e1ab32e267..e57f48d1d6 100644 --- a/src/Service.Tests/Mcp/AggregateRecordsToolTests.cs +++ b/src/Service.Tests/Mcp/AggregateRecordsToolTests.cs @@ -3,17 +3,25 @@ using System; using System.Collections.Generic; +using System.Reflection; +using System.Runtime.CompilerServices; using System.Text; using System.Text.Json; +using System.Text.Json.Nodes; using System.Threading; using System.Threading.Tasks; using Azure.DataApiBuilder.Auth; +using Azure.DataApiBuilder.Config.DatabasePrimitives; using Azure.DataApiBuilder.Config.ObjectModel; using Azure.DataApiBuilder.Core.Authorization; using Azure.DataApiBuilder.Core.Configurations; +using Azure.DataApiBuilder.Core.Models; +using Azure.DataApiBuilder.Core.Resolvers; +using Azure.DataApiBuilder.Core.Services; using Azure.DataApiBuilder.Mcp.BuiltInTools; using Azure.DataApiBuilder.Mcp.Model; using Azure.DataApiBuilder.Mcp.Utils; +using Azure.DataApiBuilder.Service.GraphQLBuilder.GraphQLTypes; using Microsoft.AspNetCore.Http; using Microsoft.Extensions.DependencyInjection; using Microsoft.VisualStudio.TestTools.UnitTesting; @@ -103,6 +111,209 @@ public void GetToolMetadata_DescriptionDocumentsWorkflowAndAlias() #endregion + #region Metadata and Query Construction Tests + + [DataTestMethod] + [DataRow(EntitySourceType.Table, false)] + [DataRow(EntitySourceType.View, false)] + [DataRow(EntitySourceType.StoredProcedure, true)] + public void ValidateEntitySourceType_OnlyRejectsStoredProcedures(EntitySourceType sourceType, bool expectError) + { + DatabaseObject databaseObject = sourceType switch + { + EntitySourceType.Table => new DatabaseTable("dbo", "books"), + EntitySourceType.View => new DatabaseView("dbo", "books"), + _ => new DatabaseStoredProcedure("dbo", "books") + }; + databaseObject.SourceType = sourceType; + + CallToolResult? result = InvokePrivate( + "ValidateEntitySourceType", + "Book", + databaseObject, + "aggregate_records", + null); + + Assert.AreEqual(expectError, result is not null); + if (result is not null) + { + AssertErrorResult(result, "InvalidEntity"); + } + } + + [TestMethod] + public void ValidateFieldsExist_ValidatesAggregateAndGroupByFields() + { + Mock metadata = new(); + metadata.Setup(x => x.TryGetBackingColumn("Book", "price", out It.Ref.IsAny)) + .Returns((string _, string field, out string? backing) => + { + backing = field == "price" ? "book_price" : null; + return field == "price"; + }); + metadata.Setup(x => x.TryGetBackingColumn("Book", "category", out It.Ref.IsAny)) + .Returns((string _, string field, out string? backing) => + { + backing = field == "category" ? "book_category" : null; + return field == "category"; + }); + + AggregateRecordsTool.AggregateArguments valid = CreateAggregateArguments( + field: "price", + groupby: new() { "category" }); + Assert.IsNull(InvokePrivate( + "ValidateFieldsExist", valid, "Book", metadata.Object, "aggregate_records", null)); + + AggregateRecordsTool.AggregateArguments badField = valid with { Field = "missing" }; + CallToolResult fieldError = InvokePrivate( + "ValidateFieldsExist", badField, "Book", metadata.Object, "aggregate_records", null); + AssertErrorResult(fieldError, "FieldNotFound"); + + AggregateRecordsTool.AggregateArguments badGroup = valid with { Groupby = new() { "missing" } }; + CallToolResult groupError = InvokePrivate( + "ValidateFieldsExist", badGroup, "Book", metadata.Object, "aggregate_records", null); + AssertErrorResult(groupError, "FieldNotFound"); + } + + [TestMethod] + public void ResolveBackingField_HandlesMappedFieldsAndCountStarPrimaryKeys() + { + SourceDefinition definition = new(); + Mock metadata = new(); + metadata.Setup(x => x.TryGetBackingColumn("Book", "price", out It.Ref.IsAny)) + .Returns((string _, string _, out string? backing) => + { + backing = "book_price"; + return true; + }); + metadata.Setup(x => x.TryGetBackingColumn("Book", "missing", out It.Ref.IsAny)) + .Returns(false); + metadata.Setup(x => x.GetSourceDefinition("Book")).Returns(definition); + + object?[] mappedArgs = + { + CreateAggregateArguments(field: "price"), "Book", metadata.Object, "aggregate_records", null, null + }; + string mapped = InvokePrivateWithMutableArguments("ResolveBackingField", mappedArgs); + Assert.AreEqual("book_price", mapped); + Assert.IsNull(mappedArgs[4]); + + object?[] missingArgs = + { + CreateAggregateArguments(field: "missing"), "Book", metadata.Object, "aggregate_records", null, null + }; + Assert.IsNull(InvokePrivateWithMutableArguments("ResolveBackingField", missingArgs)); + AssertErrorResult((CallToolResult)missingArgs[4]!, "FieldNotFound"); + + AggregateRecordsTool.AggregateArguments countStar = CreateAggregateArguments( + function: "count", + field: "*", + isCountStar: true); + object?[] noPrimaryKeyArgs = { countStar, "Book", metadata.Object, "aggregate_records", null, null }; + Assert.IsNull(InvokePrivateWithMutableArguments("ResolveBackingField", noPrimaryKeyArgs)); + AssertErrorResult((CallToolResult)noPrimaryKeyArgs[4]!, "InvalidEntity"); + + definition.PrimaryKey.Add("book_id"); + object?[] primaryKeyArgs = { countStar, "Book", metadata.Object, "aggregate_records", null, null }; + Assert.AreEqual("book_id", InvokePrivateWithMutableArguments("ResolveBackingField", primaryKeyArgs)); + } + + [TestMethod] + public void BuildAggregationStructure_AddsGroupsAggregationHavingAndPaginationState() + { + SqlQueryStructure structure = CreateUninitializedQueryStructure(); + Mock metadata = new(); + metadata.Setup(x => x.TryGetBackingColumn("Book", "category", out It.Ref.IsAny)) + .Returns((string _, string _, out string? backing) => + { + backing = "book_category"; + return true; + }); + AggregateRecordsTool.AggregateArguments args = CreateAggregateArguments( + function: "sum", + field: "price", + distinct: true, + first: 2, + groupby: new() { "category" }, + havingOperators: new() { ["gt"] = 10, ["lte"] = 100 }, + havingInValues: new() { 20, 40 }); + + InvokePrivate( + "BuildAggregationStructure", + args, + structure, + new DatabaseTable("dbo", "books"), + "book_price", + "sum_price", + "Book", + metadata.Object); + + Assert.AreEqual(1, structure.Columns.Count); + Assert.AreEqual(1, structure.GroupByMetadata.Fields.Count); + Assert.AreEqual(1, structure.GroupByMetadata.Aggregations.Count); + Assert.IsTrue(structure.GroupByMetadata.RequestedAggregations); + Assert.IsNotNull(structure.GroupByMetadata.Aggregations[0].HavingPredicates); + Assert.AreEqual(4, structure.Parameters.Count); + Assert.IsTrue(structure.IsListQuery); + Assert.AreEqual(0, structure.OrderByColumns.Count); + } + + [TestMethod] + public void BuildAggregationStructure_InvalidGroupByMapping_ThrowsDataApiBuilderException() + { + SqlQueryStructure structure = CreateUninitializedQueryStructure(); + Mock metadata = new(); + metadata.Setup(x => x.TryGetBackingColumn("Book", "missing", out It.Ref.IsAny)) + .Returns(false); + AggregateRecordsTool.AggregateArguments args = CreateAggregateArguments(groupby: new() { "missing" }); + + TargetInvocationException exception = Assert.ThrowsException(() => + InvokePrivate( + "BuildAggregationStructure", + args, + structure, + new DatabaseTable("dbo", "books"), + "book_price", + "sum_price", + "Book", + metadata.Object)); + + Assert.IsInstanceOfType(exception.InnerException); + } + + [DataTestMethod] + [DataRow("SELECT TOP 10 value FOR JSON PATH", true, "desc", false, "SELECT value ORDER BY SUM([table0].[price]) DESC OFFSET @param0 ROWS FETCH NEXT @param1 ROWS ONLY FOR JSON PATH")] + [DataRow("SELECT value FOR JSON PATH", false, "asc", true, "SELECT value ORDER BY SUM(DISTINCT [table0].[price]) ASC FOR JSON PATH")] + [DataRow("SELECT value", false, "desc", false, "SELECT value ORDER BY SUM([table0].[price]) DESC")] + public void ApplyOrderByAndPagination_RewritesSql( + string sql, + bool paginate, + string orderby, + bool distinct, + string expected) + { + SqlQueryStructure structure = CreateUninitializedQueryStructure(); + Mock builder = new(); + builder.Setup(x => x.QuoteIdentifier(It.IsAny())) + .Returns((string value) => $"[{value}]"); + AggregateRecordsTool.AggregateArguments args = CreateAggregateArguments( + function: "sum", + field: "price", + distinct: distinct, + orderby: orderby, + first: paginate ? 9 : null, + after: paginate ? Convert.ToBase64String(Encoding.UTF8.GetBytes("4")) : null, + groupby: new() { "category" }); + + string result = InvokePrivate( + "ApplyOrderByAndPagination", sql, args, structure, builder.Object, "price"); + + Assert.AreEqual(expected, result); + Assert.AreEqual(paginate ? 2 : 0, structure.Parameters.Count); + } + + #endregion + #region Configuration Tests [DataTestMethod] @@ -181,6 +392,50 @@ public async Task AggregateRecords_InvalidFieldFunctionCombination_ReturnsInvali $"Error message must contain '{expectedInMessage}'. Actual: '{message}'"); } + [DataTestMethod] + [DataRow("{\"entity\":\"Book\",\"function\":\"count\",\"field\":\"\"}", "EntityNotFound")] + [DataRow("{\"entity\":\"Book\",\"function\":\"sum\",\"field\":\"\"}", "InvalidArguments")] + [DataRow("{\"entity\":\"Book\",\"function\":\"count\",\"field\":\"id\",\"distinct\":\"yes\"}", "InvalidArguments")] + [DataRow("{\"entity\":\"Book\",\"function\":\"count\",\"first\":0}", "InvalidArguments")] + [DataRow("{\"entity\":\"Book\",\"function\":\"count\",\"first\":1}", "InvalidArguments")] + [DataRow("{\"entity\":\"Book\",\"function\":\"count\",\"after\":\"MA==\"}", "InvalidArguments")] + [DataRow("{\"entity\":\"Book\",\"function\":\"count\",\"groupby\":[\"title\"],\"after\":\"MA==\"}", "InvalidArguments")] + [DataRow("{\"entity\":\"Book\",\"function\":\"count\",\"groupby\":[\"title\"],\"having\":{\"in\":[]}}", "InvalidArguments")] + public async Task AggregateRecords_AdditionalArgumentEdges_ReturnExpectedError(string json, string expectedError) + { + CallToolResult result = await ExecuteToolAsync(CreateDefaultServiceProvider(), json); + + AssertErrorResult(result, expectedError); + } + + [TestMethod] + public async Task AggregateRecords_ValidHavingOperators_PassArgumentValidation() + { + const string json = """ + { + "entity": "Book", + "function": "count", + "groupby": ["title", "TITLE", ""], + "having": { + "eq": 1, + "neq": 2, + "gt": 3, + "gte": 4, + "lt": 5, + "lte": 6, + "in": [7, 8] + } + } + """; + + CallToolResult result = await ExecuteToolAsync(CreateDefaultServiceProvider(), json); + + JsonElement content = ParseContent(result); + Assert.AreNotEqual( + "InvalidArguments", + content.GetProperty("error").GetProperty("type").GetString()); + } + #endregion #region Input Validation Tests - GroupBy Dependencies @@ -372,6 +627,112 @@ public void DecodeCursorOffset_Base64Encoded_ReturnsExpectedOffset(string rawVal Assert.AreEqual(expectedOffset, AggregateRecordsTool.DecodeCursorOffset(cursor)); } + [DataTestMethod] + [DataRow(0, true, 1, null, false, "first")] + [DataRow(0, true, null, "MA==", false, "after")] + [DataRow(1, true, null, "MA==", false, "first")] + [DataRow(0, true, null, null, true, null)] + public void ValidateGroupByDependencies_HandlesEveryDependency( + int groupbyCount, + bool userProvidedOrderby, + int? first, + string? after, + bool expectSuccess, + string? expectedText) + { + CallToolResult? result = AggregateRecordsTool.ValidateGroupByDependencies( + groupbyCount, + ref userProvidedOrderby, + first, + after, + "aggregate_records", + logger: null); + + Assert.AreEqual(expectSuccess, result is null); + if (groupbyCount == 0) + { + Assert.IsFalse(userProvidedOrderby); + } + + if (result is not null && expectedText is not null) + { + StringAssert.Contains(GetText(result), expectedText); + } + } + + [TestMethod] + public void BuildSimpleResponse_NullEmptyAndPopulatedResults_ReturnExpectedArrays() + { + foreach (JsonArray? input in new JsonArray?[] + { + null, + new JsonArray(), + new JsonArray(new JsonObject { ["count"] = 3 }) + }) + { + CallToolResult result = InvokePrivateResponseBuilder( + "BuildSimpleResponse", + input, + "Book", + "count", + null); + JsonElement response = ParseContent(result); + JsonElement values = response.GetProperty("result"); + + Assert.AreEqual(JsonValueKind.Array, values.ValueKind); + Assert.AreEqual(1, values.GetArrayLength()); + if (input is null || input.Count == 0) + { + Assert.AreEqual(JsonValueKind.Null, values[0].GetProperty("count").ValueKind); + } + else + { + Assert.AreEqual(3, values[0].GetProperty("count").GetInt32()); + } + } + } + + [TestMethod] + public void BuildPaginatedResponse_TrimsLookaheadAndAdvancesCursor() + { + JsonArray input = new( + new JsonObject { ["group"] = "a" }, + new JsonObject { ["group"] = "b" }, + new JsonObject { ["group"] = "c" }); + string after = Convert.ToBase64String(Encoding.UTF8.GetBytes("5")); + + CallToolResult result = InvokePrivateResponseBuilder( + "BuildPaginatedResponse", + input, + 2, + after, + "Book", + null); + JsonElement page = ParseContent(result).GetProperty("result"); + + Assert.AreEqual(2, page.GetProperty("items").GetArrayLength()); + Assert.IsTrue(page.GetProperty("hasNextPage").GetBoolean()); + string endCursor = page.GetProperty("endCursor").GetString()!; + Assert.AreEqual("7", Encoding.UTF8.GetString(Convert.FromBase64String(endCursor))); + } + + [TestMethod] + public void BuildPaginatedResponse_EmptyResult_HasNoCursorOrNextPage() + { + CallToolResult result = InvokePrivateResponseBuilder( + "BuildPaginatedResponse", + null, + 2, + null, + "Book", + null); + JsonElement page = ParseContent(result).GetProperty("result"); + + Assert.AreEqual(0, page.GetProperty("items").GetArrayLength()); + Assert.IsFalse(page.GetProperty("hasNextPage").GetBoolean()); + Assert.AreEqual(JsonValueKind.Null, page.GetProperty("endCursor").ValueKind); + } + #endregion #region Timeout and Cancellation Tests @@ -549,6 +910,98 @@ private static JsonElement ParseContent(CallToolResult result) return JsonDocument.Parse(firstContent.Text).RootElement; } + private static string GetText(CallToolResult result) + { + return ((TextContentBlock)result.Content[0]).Text; + } + + private static CallToolResult InvokePrivateResponseBuilder(string methodName, params object?[] arguments) + { + MethodInfo method = typeof(AggregateRecordsTool).GetMethod( + methodName, + BindingFlags.NonPublic | BindingFlags.Static)!; + return (CallToolResult)method.Invoke(null, arguments)!; + } + + private static T InvokePrivate(string methodName, params object?[] arguments) + { + MethodInfo method = typeof(AggregateRecordsTool).GetMethod( + methodName, + BindingFlags.NonPublic | BindingFlags.Static)!; + return (T)method.Invoke(null, arguments)!; + } + + private static T InvokePrivateWithMutableArguments(string methodName, object?[] arguments) + { + MethodInfo method = typeof(AggregateRecordsTool).GetMethod( + methodName, + BindingFlags.NonPublic | BindingFlags.Static)!; + return (T)method.Invoke(null, arguments)!; + } + + private static AggregateRecordsTool.AggregateArguments CreateAggregateArguments( + string function = "sum", + string field = "price", + bool isCountStar = false, + bool distinct = false, + string orderby = "desc", + int? first = null, + string? after = null, + List? groupby = null, + Dictionary? havingOperators = null, + List? havingInValues = null) + { + return new AggregateRecordsTool.AggregateArguments( + EntityName: "Book", + Function: function, + Field: field, + IsCountStar: isCountStar, + Distinct: distinct, + Filter: null, + UserProvidedOrderby: true, + Orderby: orderby, + First: first, + After: after, + Groupby: groupby ?? new(), + HavingOperators: havingOperators, + HavingInValues: havingInValues); + } + + private static SqlQueryStructure CreateUninitializedQueryStructure() + { + SqlQueryStructure structure = (SqlQueryStructure)RuntimeHelpers.GetUninitializedObject(typeof(SqlQueryStructure)); + SetBackingField(structure, nameof(BaseQueryStructure.SourceAlias), "table0"); + SetBackingField(structure, nameof(BaseQueryStructure.Columns), new List()); + SetBackingField(structure, nameof(BaseQueryStructure.Counter), new IncrementingInteger()); + structure.Parameters = new(); + SetBackingField(structure, nameof(SqlQueryStructure.OrderByColumns), new List + { + new("dbo", "books", "id", "table0", OrderBy.DESC) + }); + SetBackingField(structure, nameof(SqlQueryStructure.GroupByMetadata), new GroupByMetadata()); + return structure; + } + + private static void SetBackingField(object instance, string propertyName, object value) + { + Type? type = instance.GetType(); + while (type is not null) + { + FieldInfo? field = type.GetField( + $"<{propertyName}>k__BackingField", + BindingFlags.Instance | BindingFlags.NonPublic); + if (field is not null) + { + field.SetValue(instance, value); + return; + } + + type = type.BaseType; + } + + Assert.Fail($"Backing field for property '{propertyName}' was not found."); + } + /// /// Asserts that the result is an error with the expected error type. /// Returns the error message for further assertions. diff --git a/src/Service.Tests/Mcp/CreateRecordToolUnitTests.cs b/src/Service.Tests/Mcp/CreateRecordToolUnitTests.cs new file mode 100644 index 0000000000..5015548470 --- /dev/null +++ b/src/Service.Tests/Mcp/CreateRecordToolUnitTests.cs @@ -0,0 +1,306 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +using System; +using System.Collections.Generic; +using System.Net; +using System.Text.Json; +using System.Threading; +using System.Threading.Tasks; +using Azure.DataApiBuilder.Auth; +using Azure.DataApiBuilder.Config.DatabasePrimitives; +using Azure.DataApiBuilder.Config.ObjectModel; +using Azure.DataApiBuilder.Core.Authorization; +using Azure.DataApiBuilder.Core.Configurations; +using Azure.DataApiBuilder.Core.Models; +using Azure.DataApiBuilder.Core.Resolvers; +using Azure.DataApiBuilder.Core.Resolvers.Factories; +using Azure.DataApiBuilder.Core.Services; +using Azure.DataApiBuilder.Core.Services.MetadataProviders; +using Azure.DataApiBuilder.Mcp.BuiltInTools; +using Azure.DataApiBuilder.Service.Exceptions; +using Microsoft.AspNetCore.Http; +using Microsoft.AspNetCore.Mvc; +using Microsoft.Extensions.DependencyInjection; +using Microsoft.VisualStudio.TestTools.UnitTesting; +using ModelContextProtocol.Protocol; +using Moq; + +namespace Azure.DataApiBuilder.Service.Tests.Mcp +{ + [TestClass] + public class CreateRecordToolUnitTests + { + private const string ENTITY_NAME = "Book"; + private const string ARGUMENTS = "{\"entity\":\"Book\",\"data\":{\"title\":\"Dune\"}}"; + + [TestMethod] + public async Task ExecuteAsync_CanceledToken_ReturnsError() + { + using CancellationTokenSource cancellationTokenSource = new(); + cancellationTokenSource.Cancel(); + using JsonDocument arguments = JsonDocument.Parse(ARGUMENTS); + + CallToolResult result = await new CreateRecordTool().ExecuteAsync( + arguments, + CreateServiceProvider(new CreatedResult("", new { id = 1 })), + cancellationTokenSource.Token); + + AssertErrorType(result, "Error"); + } + + [TestMethod] + public async Task ExecuteAsync_MissingHttpContext_UsesDefaultContextAndReturnsPermissionDenied() + { + CallToolResult result = await ExecuteAsync( + new CreatedResult("", new { id = 1 }), + includeHttpContext: false); + + AssertErrorType(result, "PermissionDenied"); + } + + [TestMethod] + public async Task ExecuteAsync_UnauthorizedOperation_ReturnsPermissionDenied() + { + CallToolResult result = await ExecuteAsync( + new CreatedResult("", new { id = 1 }), + authorizeOperation: false); + + AssertErrorType(result, "PermissionDenied"); + } + + [TestMethod] + public async Task ExecuteAsync_UnauthorizedColumn_ReturnsPermissionDenied() + { + CallToolResult result = await ExecuteAsync( + new CreatedResult("", new { id = 1 }), + authorizeColumns: false); + + AssertErrorType(result, "PermissionDenied"); + } + + [TestMethod] + public async Task ExecuteAsync_ColumnAuthorizationException_ReturnsValidationFailed() + { + DataApiBuilderException exception = new( + "column policy failed", + HttpStatusCode.BadRequest, + DataApiBuilderException.SubStatusCodes.BadRequest); + + CallToolResult result = await ExecuteAsync( + new CreatedResult("", new { id = 1 }), + columnAuthorizationException: exception); + + AssertErrorType(result, "ValidationFailed"); + } + + [TestMethod] + public async Task ExecuteAsync_StoredProcedure_ReturnsInvalidEntity() + { + DatabaseStoredProcedure storedProcedure = new("dbo", "create_book") + { + SourceType = EntitySourceType.StoredProcedure, + StoredProcedureDefinition = new() + }; + + CallToolResult result = await ExecuteAsync( + new CreatedResult("", new { id = 1 }), + dbObject: storedProcedure); + + AssertErrorType(result, "InvalidEntity"); + } + + [TestMethod] + public async Task ExecuteAsync_CreatedResult_ReturnsCreatedValue() + { + CallToolResult result = await ExecuteAsync(new CreatedResult("", new { id = 7 })); + + Assert.IsFalse(result.IsError == true); + StringAssert.Contains(GetText(result), "\"id\": 7"); + } + + [DataTestMethod] + [DataRow(400, true, "CreateFailed")] + [DataRow(500, true, "CreateFailed")] + [DataRow(403, false, "Unable to perform read-back")] + [DataRow(200, false, "Unable to perform read-back")] + public async Task ExecuteAsync_ObjectResult_MapsByStatus(int statusCode, bool isError, string expectedText) + { + ObjectResult mutationResult = new(new { detail = "result" }) { StatusCode = statusCode }; + + CallToolResult result = await ExecuteAsync(mutationResult); + + Assert.AreEqual(isError, result.IsError == true); + StringAssert.Contains(GetText(result), expectedText); + } + + [TestMethod] + public async Task ExecuteAsync_NullMutationResult_ReturnsUnexpectedError() + { + CallToolResult result = await ExecuteAsync(mutationOutcome: null); + + AssertErrorType(result, "UnexpectedError"); + } + + [TestMethod] + public async Task ExecuteAsync_UnexpectedResultType_ReturnsSuccess() + { + CallToolResult result = await ExecuteAsync(new NoContentResult()); + + Assert.IsFalse(result.IsError == true); + StringAssert.Contains(GetText(result), nameof(NoContentResult)); + } + + [TestMethod] + public async Task ExecuteAsync_MutationException_ReturnsError() + { + CallToolResult result = await ExecuteAsync(new InvalidOperationException("mutation failed")); + + AssertErrorType(result, "Error"); + StringAssert.Contains(GetText(result), "mutation failed"); + } + + private static async Task ExecuteAsync( + object? mutationOutcome, + DatabaseObject? dbObject = null, + bool includeHttpContext = true, + bool authorizeOperation = true, + bool authorizeColumns = true, + Exception? columnAuthorizationException = null) + { + using JsonDocument arguments = JsonDocument.Parse(ARGUMENTS); + return await new CreateRecordTool().ExecuteAsync( + arguments, + CreateServiceProvider( + mutationOutcome, + dbObject, + includeHttpContext, + authorizeOperation, + authorizeColumns, + columnAuthorizationException), + CancellationToken.None); + } + + private static IServiceProvider CreateServiceProvider( + object? mutationOutcome, + DatabaseObject? dbObject = null, + bool includeHttpContext = true, + bool authorizeOperation = true, + bool authorizeColumns = true, + Exception? columnAuthorizationException = null) + { + RuntimeConfig config = CreateConfig(); + RuntimeConfigProvider configProvider = TestHelper.GenerateInMemoryRuntimeConfigProvider(config); + ServiceCollection services = new(); + services.AddSingleton(configProvider); + + DatabaseObject resolvedObject = dbObject ?? new DatabaseView("dbo", "books_view") + { + SourceType = EntitySourceType.View, + ViewDefinition = new() + }; + + Mock metadataProvider = new(); + metadataProvider.Setup(x => x.EntityToDatabaseObject).Returns(new Dictionary + { + [ENTITY_NAME] = resolvedObject + }); + metadataProvider.Setup(x => x.GetDatabaseType()).Returns(DatabaseType.MSSQL); + + Mock metadataProviderFactory = new(); + metadataProviderFactory.Setup(x => x.GetMetadataProvider(It.IsAny())).Returns(metadataProvider.Object); + services.AddSingleton(metadataProviderFactory.Object); + + Mock authorizationResolver = new(); + authorizationResolver.Setup(x => x.IsValidRoleContext(It.IsAny())).Returns(true); + authorizationResolver.Setup(x => x.AreRoleAndOperationDefinedForEntity( + ENTITY_NAME, + AuthorizationResolver.ROLE_ANONYMOUS, + EntityActionOperation.Create)).Returns(authorizeOperation); + if (columnAuthorizationException is not null) + { + authorizationResolver.Setup(x => x.AreColumnsAllowedForOperation( + ENTITY_NAME, + AuthorizationResolver.ROLE_ANONYMOUS, + EntityActionOperation.Create, + It.IsAny>())).Throws(columnAuthorizationException); + } + else + { + authorizationResolver.Setup(x => x.AreColumnsAllowedForOperation( + ENTITY_NAME, + AuthorizationResolver.ROLE_ANONYMOUS, + EntityActionOperation.Create, + It.IsAny>())).Returns(authorizeColumns); + } + + services.AddSingleton(authorizationResolver.Object); + + DefaultHttpContext context = new(); + context.Request.Headers[AuthorizationResolver.CLIENT_ROLE_HEADER] = AuthorizationResolver.ROLE_ANONYMOUS; + services.AddSingleton(new HttpContextAccessor + { + HttpContext = includeHttpContext ? context : null + }); + + Mock mutationEngine = new(); + if (mutationOutcome is Exception exception) + { + mutationEngine.Setup(x => x.ExecuteAsync(It.IsAny())).ThrowsAsync(exception); + } + else + { + mutationEngine.Setup(x => x.ExecuteAsync(It.IsAny())) + .ReturnsAsync((IActionResult?)mutationOutcome); + } + + Mock mutationEngineFactory = new(); + mutationEngineFactory.Setup(x => x.GetMutationEngine(DatabaseType.MSSQL)).Returns(mutationEngine.Object); + services.AddSingleton(mutationEngineFactory.Object); + services.AddLogging(); + + return services.BuildServiceProvider(); + } + + private static RuntimeConfig CreateConfig() + { + Entity entity = new( + Source: new("books", EntitySourceType.View, null, null), + GraphQL: new("Book", "Books"), + Fields: null, + Rest: new(Enabled: true), + Permissions: new[] + { + new EntityPermission("anonymous", new[] + { + new EntityAction(EntityActionOperation.Create, null, null) + }) + }, + Mappings: null, + Relationships: null, + Mcp: null); + + return new RuntimeConfig( + Schema: "test-schema", + DataSource: new DataSource(DatabaseType.MSSQL, "", null), + Runtime: new( + Rest: new(), + GraphQL: new(), + Mcp: new(Enabled: true, Path: "/mcp", DmlTools: new(createRecord: true)), + Host: new(Cors: null, Authentication: null, Mode: HostMode.Development)), + Entities: new(new Dictionary { [ENTITY_NAME] = entity })); + } + + private static string GetText(CallToolResult result) + { + return ((TextContentBlock)result.Content[0]).Text; + } + + private static void AssertErrorType(CallToolResult result, string expectedType) + { + Assert.IsTrue(result.IsError == true, GetText(result)); + using JsonDocument document = JsonDocument.Parse(GetText(result)); + Assert.AreEqual(expectedType, document.RootElement.GetProperty("error").GetProperty("type").GetString()); + } + } +} diff --git a/src/Service.Tests/Mcp/DeleteRecordToolUnitTests.cs b/src/Service.Tests/Mcp/DeleteRecordToolUnitTests.cs new file mode 100644 index 0000000000..7bfc3fd453 --- /dev/null +++ b/src/Service.Tests/Mcp/DeleteRecordToolUnitTests.cs @@ -0,0 +1,324 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +using System; +using System.Collections.Generic; +using System.Data.Common; +using System.Net; +using System.Text.Json; +using System.Threading; +using System.Threading.Tasks; +using Azure.DataApiBuilder.Auth; +using Azure.DataApiBuilder.Config.DatabasePrimitives; +using Azure.DataApiBuilder.Config.ObjectModel; +using Azure.DataApiBuilder.Core.Authorization; +using Azure.DataApiBuilder.Core.Configurations; +using Azure.DataApiBuilder.Core.Models; +using Azure.DataApiBuilder.Core.Resolvers; +using Azure.DataApiBuilder.Core.Resolvers.Factories; +using Azure.DataApiBuilder.Core.Services; +using Azure.DataApiBuilder.Core.Services.MetadataProviders; +using Azure.DataApiBuilder.Mcp.BuiltInTools; +using Azure.DataApiBuilder.Service.Exceptions; +using Microsoft.AspNetCore.Http; +using Microsoft.AspNetCore.Mvc; +using Microsoft.Extensions.DependencyInjection; +using Microsoft.VisualStudio.TestTools.UnitTesting; +using ModelContextProtocol.Protocol; +using Moq; + +namespace Azure.DataApiBuilder.Service.Tests.Mcp +{ + [TestClass] + public class DeleteRecordToolUnitTests + { + private const string ENTITY_NAME = "Book"; + + [TestMethod] + public async Task ExecuteAsync_CanceledToken_ReturnsOperationCanceled() + { + using CancellationTokenSource cancellationTokenSource = new(); + cancellationTokenSource.Cancel(); + + CallToolResult result = await new DeleteRecordTool().ExecuteAsync( + arguments: null, + CreateServiceProvider(), + cancellationTokenSource.Token); + + AssertErrorType(result, "OperationCanceled"); + } + + [TestMethod] + public async Task ExecuteAsync_NullKey_ReturnsInvalidArguments() + { + CallToolResult result = await ExecuteAsync("{\"entity\":\"Book\",\"keys\":{\"id\":null}}", new NoContentResult()); + + AssertErrorType(result, "InvalidArguments"); + StringAssert.Contains(GetText(result), "cannot be null"); + } + + [TestMethod] + public async Task ExecuteAsync_StoredProcedure_ReturnsInvalidEntity() + { + DatabaseStoredProcedure storedProcedure = new("dbo", "get_book") + { + SourceType = EntitySourceType.StoredProcedure, + StoredProcedureDefinition = new() + }; + + CallToolResult result = await ExecuteAsync( + "{\"entity\":\"Book\",\"keys\":{\"id\":1}}", + new NoContentResult(), + dbObject: storedProcedure); + + AssertErrorType(result, "InvalidEntity"); + } + + [TestMethod] + public async Task ExecuteAsync_MissingHttpContext_ReturnsPermissionDenied() + { + CallToolResult result = await ExecuteAsync( + "{\"entity\":\"Book\",\"keys\":{\"id\":1}}", + new NoContentResult(), + includeHttpContext: false); + + AssertErrorType(result, "PermissionDenied"); + } + + [TestMethod] + public async Task ExecuteAsync_UnauthorizedOperation_ReturnsPermissionDenied() + { + CallToolResult result = await ExecuteAsync( + "{\"entity\":\"Book\",\"keys\":{\"id\":1}}", + new NoContentResult(), + authorizeOperation: false); + + AssertErrorType(result, "PermissionDenied"); + } + + [TestMethod] + public async Task ExecuteAsync_NoContentResult_ReturnsSuccess() + { + CallToolResult result = await ExecuteAsync("{\"entity\":\"Book\",\"keys\":{\"id\":1}}", new NoContentResult()); + + Assert.IsFalse(result.IsError == true); + StringAssert.Contains(GetText(result), "Record deleted successfully"); + StringAssert.Contains(GetText(result), "id=1"); + } + + [TestMethod] + public async Task ExecuteAsync_OkObjectResult_IncludesResult() + { + CallToolResult result = await ExecuteAsync( + "{\"entity\":\"Book\",\"keys\":{\"id\":1}}", + new OkObjectResult(new { deleted = 1 })); + + Assert.IsFalse(result.IsError == true); + StringAssert.Contains(GetText(result), "deleted"); + } + + [DataTestMethod] + [DataRow("Could not find item with id", "RecordNotFound")] + [DataRow("violates foreign key constraint", "ConstraintViolation")] + [DataRow("REFERENCE constraint failure", "ConstraintViolation")] + [DataRow("authorization failed", "PermissionDenied")] + [DataRow("invalid key type", "InvalidArguments")] + [DataRow("other DAB failure", "DataApiBuilderError")] + public async Task ExecuteAsync_DataApiBuilderException_MapsError(string message, string expectedError) + { + DataApiBuilderException exception = new( + message, + HttpStatusCode.BadRequest, + DataApiBuilderException.SubStatusCodes.BadRequest); + + CallToolResult result = await ExecuteAsync("{\"entity\":\"Book\",\"keys\":{\"id\":1}}", exception); + + AssertErrorType(result, expectedError); + } + + [DataTestMethod] + [DataRow("foreign key failure", "ConstraintViolation")] + [DataRow("record does not exist", "RecordNotFound")] + [DataRow("provider exploded", "DatabaseError")] + public async Task ExecuteAsync_DbException_MapsError(string message, string expectedError) + { + CallToolResult result = await ExecuteAsync( + "{\"entity\":\"Book\",\"keys\":{\"id\":1}}", + new FakeDbException(message)); + + AssertErrorType(result, expectedError); + } + + [DataTestMethod] + [DataRow("connection unavailable", "ConnectionError")] + [DataRow("Could not find record", "RecordNotFound")] + [DataRow("unexpected failure", "UnexpectedError")] + public async Task ExecuteAsync_GeneralException_MapsError(string message, string expectedError) + { + CallToolResult result = await ExecuteAsync( + "{\"entity\":\"Book\",\"keys\":{\"id\":1}}", + new InvalidOperationException(message)); + + AssertErrorType(result, expectedError); + } + + [TestMethod] + public async Task ExecuteAsync_Timeout_ReturnsTimeoutError() + { + CallToolResult result = await ExecuteAsync( + "{\"entity\":\"Book\",\"keys\":{\"id\":1}}", + new TimeoutException()); + + AssertErrorType(result, "TimeoutError"); + } + + [TestMethod] + public async Task ExecuteAsync_InvalidPrimaryKey_ReturnsDataApiBuilderError() + { + CallToolResult result = await ExecuteAsync( + "{\"entity\":\"Book\",\"keys\":{\"other\":1}}", + new NoContentResult()); + + AssertErrorType(result, "UnexpectedError"); + } + + private static async Task ExecuteAsync( + string arguments, + object mutationOutcome, + DatabaseObject? dbObject = null, + bool includeHttpContext = true, + bool authorizeOperation = true) + { + using JsonDocument document = JsonDocument.Parse(arguments); + IServiceProvider serviceProvider = CreateServiceProvider( + mutationOutcome, + dbObject, + includeHttpContext, + authorizeOperation); + return await new DeleteRecordTool().ExecuteAsync(document, serviceProvider, CancellationToken.None); + } + + private static IServiceProvider CreateServiceProvider( + object? mutationOutcome = null, + DatabaseObject? dbObject = null, + bool includeHttpContext = true, + bool authorizeOperation = true) + { + RuntimeConfig config = CreateConfig(); + RuntimeConfigProvider configProvider = TestHelper.GenerateInMemoryRuntimeConfigProvider(config); + ServiceCollection services = new(); + services.AddSingleton(configProvider); + + SourceDefinition sourceDefinition = new() { PrimaryKey = new() { "id" } }; + DatabaseObject resolvedObject = dbObject ?? new DatabaseTable("dbo", "books") + { + SourceType = EntitySourceType.Table, + TableDefinition = sourceDefinition + }; + + Mock metadataProvider = new(); + metadataProvider.Setup(x => x.EntityToDatabaseObject).Returns(new Dictionary + { + [ENTITY_NAME] = resolvedObject + }); + metadataProvider.Setup(x => x.GetSourceDefinition(ENTITY_NAME)).Returns(sourceDefinition); + string? idBackingColumn = "id"; + metadataProvider.Setup(x => x.TryGetBackingColumn(ENTITY_NAME, "id", out idBackingColumn)).Returns(true); + string? missingBackingColumn = null; + metadataProvider.Setup(x => x.TryGetBackingColumn(ENTITY_NAME, "other", out missingBackingColumn)).Returns(false); + + Mock metadataProviderFactory = new(); + metadataProviderFactory.Setup(x => x.GetMetadataProvider(It.IsAny())).Returns(metadataProvider.Object); + services.AddSingleton(metadataProviderFactory.Object); + + Mock authorizationResolver = new(); + authorizationResolver.Setup(x => x.IsValidRoleContext(It.IsAny())).Returns(true); + authorizationResolver.Setup(x => x.AreRoleAndOperationDefinedForEntity( + ENTITY_NAME, + AuthorizationResolver.ROLE_ANONYMOUS, + EntityActionOperation.Delete)).Returns(authorizeOperation); + services.AddSingleton(authorizationResolver.Object); + + DefaultHttpContext httpContext = new(); + httpContext.Request.Headers[AuthorizationResolver.CLIENT_ROLE_HEADER] = AuthorizationResolver.ROLE_ANONYMOUS; + services.AddSingleton(new HttpContextAccessor + { + HttpContext = includeHttpContext ? httpContext : null + }); + + Mock mutationEngine = new(); + if (mutationOutcome is Exception exception) + { + mutationEngine.Setup(x => x.ExecuteAsync(It.IsAny())).ThrowsAsync(exception); + } + else + { + mutationEngine.Setup(x => x.ExecuteAsync(It.IsAny())) + .ReturnsAsync((IActionResult?)mutationOutcome ?? new NoContentResult()); + } + + Mock mutationEngineFactory = new(); + mutationEngineFactory.Setup(x => x.GetMutationEngine(DatabaseType.MSSQL)).Returns(mutationEngine.Object); + services.AddSingleton(mutationEngineFactory.Object); + services.AddLogging(); + + return services.BuildServiceProvider(); + } + + private static RuntimeConfig CreateConfig() + { + Entity entity = new( + Source: new("books", EntitySourceType.Table, null, null), + GraphQL: new("Book", "Books"), + Fields: null, + Rest: new(Enabled: true), + Permissions: new[] + { + new EntityPermission("anonymous", new[] + { + new EntityAction(EntityActionOperation.Delete, null, null) + }) + }, + Mappings: null, + Relationships: null, + Mcp: null); + + return new RuntimeConfig( + Schema: "test-schema", + DataSource: new DataSource(DatabaseType.MSSQL, "", null), + Runtime: new( + Rest: new(), + GraphQL: new(), + Mcp: new(Enabled: true, Path: "/mcp", DmlTools: new(deleteRecord: true)), + Host: new(Cors: null, Authentication: null, Mode: HostMode.Development)), + Entities: new(new Dictionary { [ENTITY_NAME] = entity })); + } + + private static string GetText(CallToolResult result) + { + return ((TextContentBlock)result.Content[0]).Text; + } + + private static void AssertErrorType(CallToolResult result, string expectedType) + { + Assert.IsTrue(result.IsError == true, GetText(result)); + using JsonDocument document = JsonDocument.Parse(GetText(result)); + Assert.AreEqual(expectedType, document.RootElement.GetProperty("error").GetProperty("type").GetString()); + } + + private sealed class FakeDbException : DbException + { + public FakeDbException() + { + } + + public FakeDbException(string message) : base(message) + { + } + + public FakeDbException(string message, Exception innerException) : base(message, innerException) + { + } + } + } +} diff --git a/src/Service.Tests/Mcp/DescribeEntitiesFilteringTests.cs b/src/Service.Tests/Mcp/DescribeEntitiesFilteringTests.cs index 105dbfc907..c44a7af001 100644 --- a/src/Service.Tests/Mcp/DescribeEntitiesFilteringTests.cs +++ b/src/Service.Tests/Mcp/DescribeEntitiesFilteringTests.cs @@ -159,6 +159,48 @@ public async Task DescribeEntities_ReturnsNoEntitiesConfigured_WhenConfigHasNoEn Assert.IsTrue(message.Contains("No entities are configured")); } + [TestMethod] + public async Task DescribeEntities_ReturnsToolDisabled_WhenDescribeToolIsDisabled() + { + RuntimeConfig config = new( + Schema: "test-schema", + DataSource: new DataSource(DatabaseType: DatabaseType.MSSQL, ConnectionString: "", Options: null), + Runtime: new( + Rest: new(), + GraphQL: new(), + Mcp: new(Enabled: true, Path: "/mcp", DmlTools: DmlToolsConfig.FromBoolean(false)), + Host: new(Cors: null, Authentication: null, Mode: HostMode.Development)), + Entities: new(new Dictionary())); + + CallToolResult result = await new DescribeEntitiesTool().ExecuteAsync( + null, CreateServiceProvider(config), CancellationToken.None); + + AssertErrorResult(result, "ToolDisabled"); + } + + [TestMethod] + public async Task DescribeEntities_ReturnsOperationCanceled_WhenCancellationRequested() + { + using CancellationTokenSource cancellation = new(); + cancellation.Cancel(); + + CallToolResult result = await new DescribeEntitiesTool().ExecuteAsync( + null, CreateServiceProvider(CreateConfigWithNoEntities()), cancellation.Token); + + AssertErrorResult(result, "OperationCanceled"); + } + + [TestMethod] + public async Task DescribeEntities_ReturnsEntitiesNotFound_ForExplicitMissingFilter() + { + using JsonDocument arguments = JsonDocument.Parse("{\"entities\":[\"Missing\",\" \",42]}"); + + CallToolResult result = await new DescribeEntitiesTool().ExecuteAsync( + arguments, CreateServiceProvider(CreateConfigWithNoEntities()), CancellationToken.None); + + AssertErrorResult(result, "EntitiesNotFound"); + } + /// /// CRITICAL TEST: Verifies that stored procedures with BOTH custom-tool AND dml-tools enabled /// appear in describe_entities. This validates the truth table scenario: diff --git a/src/Service.Tests/Mcp/DynamicCustomToolTests.cs b/src/Service.Tests/Mcp/DynamicCustomToolTests.cs index b3debd48f1..f955fe1111 100644 --- a/src/Service.Tests/Mcp/DynamicCustomToolTests.cs +++ b/src/Service.Tests/Mcp/DynamicCustomToolTests.cs @@ -3,7 +3,9 @@ using System; using System.Collections.Generic; +using System.Data.Common; using System.Linq; +using System.Net; using System.Text.Json; using System.Threading; using System.Threading.Tasks; @@ -18,6 +20,7 @@ using Azure.DataApiBuilder.Core.Services; using Azure.DataApiBuilder.Core.Services.MetadataProviders; using Azure.DataApiBuilder.Mcp.Core; +using Azure.DataApiBuilder.Service.Exceptions; using Microsoft.AspNetCore.Http; using Microsoft.AspNetCore.Mvc; using Microsoft.Extensions.DependencyInjection; @@ -350,6 +353,225 @@ public async Task ExecuteAsync_ZeroParamSP_PassesEmptyParams() Assert.AreEqual(0, capturedContext!.ResolvedParameters.Count); } + [TestMethod] + public void IsEnabled_AlwaysReturnsTrue() + { + DynamicCustomTool tool = new(TEST_ENTITY, CreateTestStoredProcedureEntity()); + + Assert.IsTrue(tool.IsEnabled(CreateRuntimeConfig())); + } + + [TestMethod] + public async Task ExecuteAsync_CanceledToken_ReturnsOperationCanceled() + { + DynamicCustomTool tool = new(TEST_ENTITY, CreateTestStoredProcedureEntity()); + IServiceProvider serviceProvider = BuildExecutionServiceProvider(new()); + using CancellationTokenSource cancellationTokenSource = new(); + cancellationTokenSource.Cancel(); + + CallToolResult result = await tool.ExecuteAsync(null, serviceProvider, cancellationTokenSource.Token); + + AssertError(result, "OperationCanceled"); + } + + [TestMethod] + public async Task ExecuteAsync_EntityRemovedAfterRegistration_ReturnsEntityNotFound() + { + CallToolResult result = await ExecuteCustomToolAsync( + dbParameters: new(), + userParameters: null, + runtimeConfig: CreateRuntimeConfig(includeEntity: false)); + + AssertError(result, "EntityNotFound"); + } + + [TestMethod] + public async Task ExecuteAsync_EntityChangedToTable_ReturnsInvalidEntity() + { + CallToolResult result = await ExecuteCustomToolAsync( + dbParameters: new(), + userParameters: null, + runtimeConfig: CreateRuntimeConfig(sourceType: EntitySourceType.Table)); + + AssertError(result, "InvalidEntity"); + } + + [TestMethod] + public async Task ExecuteAsync_UnresolvableMetadata_ReturnsEntityNotFound() + { + CallToolResult result = await ExecuteCustomToolAsync( + dbParameters: new(), + userParameters: null, + registerMetadataFactory: false); + + AssertError(result, "EntityNotFound"); + } + + [TestMethod] + public async Task ExecuteAsync_InvalidRoleContext_ReturnsPermissionDenied() + { + CallToolResult result = await ExecuteCustomToolAsync( + dbParameters: new(), + userParameters: null, + validRoleContext: false); + + AssertError(result, "PermissionDenied"); + } + + [TestMethod] + public async Task ExecuteAsync_UnauthorizedOperation_ReturnsPermissionDenied() + { + CallToolResult result = await ExecuteCustomToolAsync( + dbParameters: new(), + userParameters: null, + authorizeOperation: false); + + AssertError(result, "PermissionDenied"); + } + + [TestMethod] + public async Task ExecuteAsync_MetadataObjectChangedToTable_ReturnsInvalidEntity() + { + CallToolResult result = await ExecuteCustomToolAsync( + dbParameters: new(), + userParameters: null, + metadataObject: new DatabaseTable("dbo", "books") { SourceType = EntitySourceType.Table }); + + AssertError(result, "InvalidEntity"); + } + + [TestMethod] + public async Task ExecuteAsync_ConvertsAllJsonParameterKinds() + { + Dictionary dbParameters = new() + { + ["truth"] = new(), + ["falsehood"] = new(), + ["nothing"] = new(), + ["complex"] = new() + }; + StoredProcedureRequestContext? capturedContext = null; + IServiceProvider serviceProvider = BuildExecutionServiceProvider( + dbParameters, + context => capturedContext = context); + DynamicCustomTool tool = new(TEST_ENTITY, CreateTestStoredProcedureEntity()); + using JsonDocument arguments = JsonDocument.Parse( + "{\"truth\":true,\"falsehood\":false,\"nothing\":null,\"complex\":{\"x\":1}}"); + + CallToolResult result = await tool.ExecuteAsync(arguments, serviceProvider, CancellationToken.None); + + AssertSuccess(result, "All JSON parameter kinds should be converted."); + Assert.IsNotNull(capturedContext); + Assert.AreEqual(true, capturedContext.ResolvedParameters["truth"]); + Assert.AreEqual(false, capturedContext.ResolvedParameters["falsehood"]); + Assert.IsNull(capturedContext.ResolvedParameters["nothing"]); + Assert.AreEqual("{\"x\":1}", capturedContext.ResolvedParameters["complex"]); + } + + [DataTestMethod] + [DataRow("dab", "ExecutionError")] + [DataRow("database", "DatabaseError")] + [DataRow("unexpected", "UnexpectedError")] + public async Task ExecuteAsync_ExecutionException_MapsError(string exceptionType, string expectedError) + { + Exception exception = exceptionType switch + { + "dab" => new DataApiBuilderException( + "DAB execution failed", + HttpStatusCode.BadRequest, + DataApiBuilderException.SubStatusCodes.BadRequest), + "database" => new FakeDbException("provider failed"), + _ => new InvalidOperationException("unexpected failure") + }; + + CallToolResult result = await ExecuteCustomToolAsync( + dbParameters: new(), + userParameters: null, + queryOutcome: exception); + + AssertError(result, expectedError); + } + + [DataTestMethod] + [DataRow("bad", "BadRequest")] + [DataRow("unauthorized", "PermissionDenied")] + [DataRow("unknown", "Stored procedure executed successfully")] + public async Task ExecuteAsync_NonSuccessResult_MapsResponse(string resultType, string expectedText) + { + IActionResult queryResult = resultType switch + { + "bad" => new BadRequestObjectResult((object?)null), + "unauthorized" => new UnauthorizedObjectResult("denied"), + _ => new NoContentResult() + }; + + CallToolResult result = await ExecuteCustomToolAsync( + dbParameters: new(), + userParameters: null, + queryOutcome: queryResult); + + StringAssert.Contains(GetFirstText(result), expectedText); + } + + [TestMethod] + public async Task ExecuteAsync_JsonDocumentObjectResult_WrapsValueInArray() + { + using JsonDocument queryDocument = JsonDocument.Parse("{\"id\":1}"); + CallToolResult result = await ExecuteCustomToolAsync( + dbParameters: new(), + userParameters: new() { ["ignored"] = null }, + queryOutcome: new OkObjectResult(queryDocument)); + + AssertError(result, "InvalidArguments"); + + result = await ExecuteCustomToolAsync( + dbParameters: new(), + userParameters: null, + queryOutcome: new OkObjectResult(queryDocument)); + + AssertSuccess(result, "JsonDocument results should be returned."); + using JsonDocument response = JsonDocument.Parse(GetFirstText(result)); + Assert.AreEqual(JsonValueKind.Array, response.RootElement.GetProperty("value").ValueKind); + } + + [TestMethod] + public async Task ExecuteAsync_PocoResult_IsSerialized() + { + CallToolResult result = await ExecuteCustomToolAsync( + dbParameters: new(), + userParameters: null, + queryOutcome: new OkObjectResult(new { count = 2 })); + + AssertSuccess(result, "POCO results should be serialized."); + StringAssert.Contains(GetFirstText(result), "count"); + } + + [TestMethod] + public void InitializeMetadata_UnresolvableOrNonStoredProcedureMetadata_FallsBackToConfig() + { + Entity entity = CreateTestStoredProcedureEntity(parameters: new[] { new ParameterMetadata { Name = "id" } }); + DynamicCustomTool missingMetadataTool = new("MissingMetadata", entity); + missingMetadataTool.InitializeMetadata(BuildMetadataServiceProvider("MissingMetadata", metadataObject: null)); + Assert.AreEqual(JsonValueKind.Array, ParseSchemaProperties(missingMetadataTool.GetToolMetadata()).GetProperty("id").GetProperty("type").ValueKind); + + DynamicCustomTool tableMetadataTool = new("TableMetadata", entity); + tableMetadataTool.InitializeMetadata(BuildMetadataServiceProvider( + "TableMetadata", + new DatabaseTable("dbo", "books") { SourceType = EntitySourceType.Table })); + Assert.AreEqual(JsonValueKind.Array, ParseSchemaProperties(tableMetadataTool.GetToolMetadata()).GetProperty("id").GetProperty("type").ValueKind); + } + + [TestMethod] + public void InitializeMetadata_UnknownSystemType_UsesPermissiveSchema() + { + JsonElement properties = InitializeAndGetSchemaProperties(new Dictionary + { + ["payload"] = new() { SystemType = typeof(object) } + }); + + Assert.AreEqual(JsonValueKind.Array, properties.GetProperty("payload").GetProperty("type").ValueKind); + } + #endregion #region Execution Helpers @@ -363,11 +585,23 @@ public async Task ExecuteAsync_ZeroParamSP_PassesEmptyParams() private static async Task ExecuteCustomToolAsync( Dictionary dbParameters, Dictionary? userParameters, - Action? captureContext = null) + Action? captureContext = null, + object? queryOutcome = null, + RuntimeConfig? runtimeConfig = null, + DatabaseObject? metadataObject = null, + bool registerMetadataFactory = true, + bool validRoleContext = true, + bool authorizeOperation = true) { IServiceProvider sp = BuildExecutionServiceProvider( dbParameters: dbParameters, - captureContext: captureContext); + captureContext: captureContext, + queryOutcome: queryOutcome, + runtimeConfig: runtimeConfig, + metadataObject: metadataObject, + registerMetadataFactory: registerMetadataFactory, + validRoleContext: validRoleContext, + authorizeOperation: authorizeOperation); Entity entity = CreateTestStoredProcedureEntity(); DynamicCustomTool tool = new(TEST_ENTITY, entity); @@ -385,38 +619,15 @@ private static async Task ExecuteCustomToolAsync( /// private static IServiceProvider BuildExecutionServiceProvider( Dictionary dbParameters, - Action? captureContext = null) + Action? captureContext = null, + object? queryOutcome = null, + RuntimeConfig? runtimeConfig = null, + DatabaseObject? metadataObject = null, + bool registerMetadataFactory = true, + bool validRoleContext = true, + bool authorizeOperation = true) { - Entity entity = new( - Source: new(SP_SOURCE_OBJECT, EntitySourceType.StoredProcedure, Parameters: null, KeyFields: null), - GraphQL: new(TEST_ENTITY, TEST_ENTITY), - Rest: new(Enabled: true), - Fields: null, - Permissions: new[] - { - new EntityPermission( - Role: "anonymous", - Actions: new[] - { - new EntityAction(Action: EntityActionOperation.Execute, Fields: null, Policy: null) - }) - }, - Relationships: null, - Mappings: null, - Mcp: new EntityMcpOptions(customToolEnabled: true, dmlToolsEnabled: null)); - - Dictionary entities = new() { [TEST_ENTITY] = entity }; - - RuntimeConfig config = new( - Schema: "test-schema", - DataSource: new DataSource(DatabaseType: DatabaseType.MSSQL, ConnectionString: "", Options: null), - Runtime: new( - Rest: new(), - GraphQL: new(), - Mcp: new(Enabled: true, Path: "/mcp", DmlTools: null), - Host: new(Cors: null, Authentication: null, Mode: HostMode.Development) - ), - Entities: new(entities)); + RuntimeConfig config = runtimeConfig ?? CreateRuntimeConfig(); ServiceCollection services = new(); @@ -425,11 +636,11 @@ private static IServiceProvider BuildExecutionServiceProvider( // Mock authorization resolver Mock mockAuthResolver = new(); - mockAuthResolver.Setup(x => x.IsValidRoleContext(It.IsAny())).Returns(true); + mockAuthResolver.Setup(x => x.IsValidRoleContext(It.IsAny())).Returns(validRoleContext); mockAuthResolver .Setup(x => x.AreRoleAndOperationDefinedForEntity( It.IsAny(), It.IsAny(), It.IsAny())) - .Returns(true); + .Returns(authorizeOperation); services.AddSingleton(mockAuthResolver.Object); // Mock HttpContext @@ -439,7 +650,7 @@ private static IServiceProvider BuildExecutionServiceProvider( services.AddSingleton(httpContextAccessor); // Mock metadata provider with DatabaseStoredProcedure - DatabaseObject dbObject = new DatabaseStoredProcedure("dbo", SP_SOURCE_OBJECT) + DatabaseObject dbObject = metadataObject ?? new DatabaseStoredProcedure("dbo", SP_SOURCE_OBJECT) { SourceType = EntitySourceType.StoredProcedure, StoredProcedureDefinition = new StoredProcedureDefinition @@ -458,18 +669,35 @@ private static IServiceProvider BuildExecutionServiceProvider( mockMetadataProviderFactory .Setup(x => x.GetMetadataProvider(It.IsAny())) .Returns(mockSqlMetadataProvider.Object); - services.AddSingleton(mockMetadataProviderFactory.Object); + if (registerMetadataFactory) + { + services.AddSingleton(mockMetadataProviderFactory.Object); + } // Mock query engine Mock mockQueryEngine = new(); - mockQueryEngine - .Setup(x => x.ExecuteAsync(It.IsAny(), It.IsAny())) - .Returns((StoredProcedureRequestContext ctx, string ds) => + if (queryOutcome is Exception exception) + { + mockQueryEngine + .Setup(x => x.ExecuteAsync(It.IsAny(), It.IsAny())) + .ThrowsAsync(exception); + } + else + { + mockQueryEngine + .Setup(x => x.ExecuteAsync(It.IsAny(), It.IsAny())) + .Returns((StoredProcedureRequestContext ctx, string ds) => { captureContext?.Invoke(ctx); + if (queryOutcome is IActionResult configuredResult) + { + return Task.FromResult(configuredResult); + } + using JsonDocument doc = JsonDocument.Parse("[]"); return Task.FromResult(new OkObjectResult(doc.RootElement.Clone())); }); + } Mock mockQueryEngineFactory = new(); mockQueryEngineFactory @@ -482,6 +710,69 @@ private static IServiceProvider BuildExecutionServiceProvider( return services.BuildServiceProvider(); } + private static RuntimeConfig CreateRuntimeConfig( + bool includeEntity = true, + EntitySourceType sourceType = EntitySourceType.StoredProcedure) + { + Dictionary entities = new(); + if (includeEntity) + { + entities[TEST_ENTITY] = new Entity( + Source: new(SP_SOURCE_OBJECT, sourceType, Parameters: null, KeyFields: null), + GraphQL: new(TEST_ENTITY, TEST_ENTITY), + Rest: new(Enabled: true), + Fields: null, + Permissions: new[] + { + new EntityPermission( + Role: "anonymous", + Actions: new[] + { + new EntityAction(Action: EntityActionOperation.Execute, Fields: null, Policy: null) + }) + }, + Relationships: null, + Mappings: null, + Mcp: new EntityMcpOptions(customToolEnabled: true, dmlToolsEnabled: null)); + } + + return new RuntimeConfig( + Schema: "test-schema", + DataSource: new DataSource(DatabaseType: DatabaseType.MSSQL, ConnectionString: "", Options: null), + Runtime: new( + Rest: new(), + GraphQL: new(), + Mcp: new(Enabled: true, Path: "/mcp", DmlTools: null), + Host: new(Cors: null, Authentication: null, Mode: HostMode.Development)), + Entities: new(entities)); + } + + private static IServiceProvider BuildMetadataServiceProvider(string entityName, DatabaseObject? metadataObject) + { + RuntimeConfig config = new( + Schema: "test-schema", + DataSource: new DataSource(DatabaseType.MSSQL, "", null), + Entities: new(new Dictionary + { + [entityName] = CreateTestStoredProcedureEntity() + })); + ServiceCollection services = new(); + services.AddSingleton(TestHelper.GenerateInMemoryRuntimeConfigProvider(config)); + + Mock metadataProvider = new(); + Dictionary objects = new(); + if (metadataObject is not null) + { + objects[entityName] = metadataObject; + } + + metadataProvider.Setup(x => x.EntityToDatabaseObject).Returns(objects); + Mock factory = new(); + factory.Setup(x => x.GetMetadataProvider(It.IsAny())).Returns(metadataProvider.Object); + services.AddSingleton(factory.Object); + return services.BuildServiceProvider(); + } + private static void AssertSuccess(CallToolResult result, string message) { Assert.IsTrue(result.IsError != true, @@ -500,6 +791,27 @@ private static string GetFirstText(CallToolResult result) : string.Empty; } + private static void AssertError(CallToolResult result, string expectedError) + { + Assert.IsTrue(result.IsError == true, GetFirstText(result)); + StringAssert.Contains(GetFirstText(result), expectedError); + } + + private sealed class FakeDbException : DbException + { + public FakeDbException() + { + } + + public FakeDbException(string message) : base(message) + { + } + + public FakeDbException(string message, Exception innerException) : base(message, innerException) + { + } + } + #endregion /// diff --git a/src/Service.Tests/Mcp/ExecuteEntityToolTests.cs b/src/Service.Tests/Mcp/ExecuteEntityToolTests.cs index 4cc7b352fd..379ba882c1 100644 --- a/src/Service.Tests/Mcp/ExecuteEntityToolTests.cs +++ b/src/Service.Tests/Mcp/ExecuteEntityToolTests.cs @@ -3,6 +3,8 @@ using System; using System.Collections.Generic; +using System.Data.Common; +using System.Net; using System.Text.Json; using System.Threading; using System.Threading.Tasks; @@ -17,6 +19,7 @@ using Azure.DataApiBuilder.Core.Services; using Azure.DataApiBuilder.Core.Services.MetadataProviders; using Azure.DataApiBuilder.Mcp.BuiltInTools; +using Azure.DataApiBuilder.Service.Exceptions; using Microsoft.AspNetCore.Http; using Microsoft.AspNetCore.Mvc; using Microsoft.Extensions.DependencyInjection; @@ -283,6 +286,241 @@ public async Task ExecuteEntity_ReturnsError_WhenEntityNotFound() #endregion + #region Result and error handling tests + + [TestMethod] + public async Task ExecuteEntity_CanceledToken_ReturnsOperationCanceled() + { + IServiceProvider serviceProvider = BuildServiceProvider( + TEST_ENTITY, + SP_SOURCE_OBJECT, + EntitySourceType.StoredProcedure, + new()); + using JsonDocument arguments = JsonDocument.Parse("{\"entity\":\"GetBook\"}"); + using CancellationTokenSource cancellationTokenSource = new(); + cancellationTokenSource.Cancel(); + + CallToolResult result = await new ExecuteEntityTool().ExecuteAsync( + arguments, + serviceProvider, + cancellationTokenSource.Token); + + AssertError(result, "OperationCanceled"); + } + + [TestMethod] + public async Task ExecuteEntity_NullArguments_ReturnsInvalidArguments() + { + IServiceProvider serviceProvider = BuildServiceProvider( + TEST_ENTITY, + SP_SOURCE_OBJECT, + EntitySourceType.StoredProcedure, + new()); + + CallToolResult result = await new ExecuteEntityTool().ExecuteAsync(null, serviceProvider, CancellationToken.None); + + AssertError(result, "InvalidArguments"); + } + + [TestMethod] + public async Task ExecuteEntity_MissingHttpContext_ReturnsPermissionDenied() + { + CallToolResult result = await ExecuteWithMockedEngineAsync( + TEST_ENTITY, + new(), + null, + includeHttpContext: false); + + AssertError(result, "PermissionDenied"); + } + + [TestMethod] + public async Task ExecuteEntity_UnauthorizedOperation_ReturnsPermissionDenied() + { + CallToolResult result = await ExecuteWithMockedEngineAsync( + TEST_ENTITY, + new(), + null, + authorizeOperation: false); + + AssertError(result, "PermissionDenied"); + } + + [TestMethod] + public async Task ExecuteEntity_MetadataObjectIsNotStoredProcedure_ReturnsInvalidEntity() + { + DatabaseTable table = new("dbo", "books") { SourceType = EntitySourceType.Table }; + CallToolResult result = await ExecuteWithMockedEngineAsync( + TEST_ENTITY, + new(), + null, + metadataObject: table); + + AssertError(result, "InvalidEntity"); + } + + [TestMethod] + public async Task ExecuteEntity_ConvertsJsonParameterKinds() + { + Dictionary parameters = new() + { + ["text"] = new(), + ["integer"] = new(), + ["number"] = new(), + ["truth"] = new(), + ["falsehood"] = new(), + ["nothing"] = new(), + ["complex"] = new() + }; + StoredProcedureRequestContext? capturedContext = null; + IServiceProvider serviceProvider = BuildServiceProvider( + TEST_ENTITY, + SP_SOURCE_OBJECT, + EntitySourceType.StoredProcedure, + parameters, + captureContext: context => capturedContext = context); + using JsonDocument arguments = JsonDocument.Parse( + "{\"entity\":\"GetBook\",\"parameters\":{" + + "\"text\":\"value\",\"integer\":42,\"number\":1.5," + + "\"truth\":true,\"falsehood\":false,\"nothing\":null,\"complex\":{\"x\":1}}}"); + + CallToolResult result = await new ExecuteEntityTool().ExecuteAsync(arguments, serviceProvider, CancellationToken.None); + + AssertSuccess(result, "All supported JSON parameter kinds should be converted."); + Assert.IsNotNull(capturedContext); + Assert.AreEqual("value", capturedContext.ResolvedParameters["text"]); + Assert.AreEqual(42L, capturedContext.ResolvedParameters["integer"]); + Assert.AreEqual(1.5m, capturedContext.ResolvedParameters["number"]); + Assert.AreEqual(true, capturedContext.ResolvedParameters["truth"]); + Assert.AreEqual(false, capturedContext.ResolvedParameters["falsehood"]); + Assert.IsNull(capturedContext.ResolvedParameters["nothing"]); + Assert.AreEqual("{\"x\":1}", capturedContext.ResolvedParameters["complex"]); + } + + [DataTestMethod] + [DataRow("permission denied", "PermissionDenied")] + [DataRow("invalid parameter type", "InvalidArguments")] + [DataRow("other DAB error", "DataApiBuilderError")] + public async Task ExecuteEntity_DataApiBuilderException_MapsError(string message, string expectedError) + { + DataApiBuilderException exception = new( + message, + HttpStatusCode.BadRequest, + DataApiBuilderException.SubStatusCodes.BadRequest); + + CallToolResult result = await ExecuteWithMockedEngineAsync( + TEST_ENTITY, + new(), + null, + queryOutcome: exception); + + AssertError(result, expectedError); + } + + [DataTestMethod] + [DataRow("provider failure", "DatabaseError")] + [DataRow("connection unavailable", "ConnectionError")] + [DataRow("unexpected", "DatabaseError")] + public async Task ExecuteEntity_ExecutionException_MapsError(string message, string expectedError) + { + Exception exception = message switch + { + "provider failure" => new FakeDbException(message), + "connection unavailable" => new InvalidOperationException(message), + _ => new Exception(message) + }; + + CallToolResult result = await ExecuteWithMockedEngineAsync( + TEST_ENTITY, + new(), + null, + queryOutcome: exception); + + AssertError(result, expectedError); + } + + [TestMethod] + public async Task ExecuteEntity_Timeout_ReturnsTimeoutError() + { + CallToolResult result = await ExecuteWithMockedEngineAsync( + TEST_ENTITY, + new(), + null, + queryOutcome: new TimeoutException()); + + AssertError(result, "TimeoutError"); + } + + [TestMethod] + public async Task ExecuteEntity_BadRequestResult_ReturnsError() + { + CallToolResult result = await ExecuteWithMockedEngineAsync( + TEST_ENTITY, + new(), + null, + queryOutcome: new BadRequestObjectResult("bad input")); + + AssertError(result, "BadRequest"); + } + + [TestMethod] + public async Task ExecuteEntity_UnauthorizedResult_ReturnsPermissionDenied() + { + CallToolResult result = await ExecuteWithMockedEngineAsync( + TEST_ENTITY, + new(), + null, + queryOutcome: new UnauthorizedObjectResult("denied")); + + AssertError(result, "PermissionDenied"); + } + + [TestMethod] + public async Task ExecuteEntity_NonJsonResult_IsSerialized() + { + CallToolResult result = await ExecuteWithMockedEngineAsync( + TEST_ENTITY, + new(), + null, + queryOutcome: new OkObjectResult(new { count = 2 })); + + AssertSuccess(result, "POCO results should be serialized."); + StringAssert.Contains(GetFirstText(result), "count"); + } + + [TestMethod] + public async Task ExecuteEntity_ObjectJsonResult_IsWrappedInArray() + { + using JsonDocument document = JsonDocument.Parse("{\"id\":1}"); + CallToolResult result = await ExecuteWithMockedEngineAsync( + TEST_ENTITY, + new(), + null, + queryOutcome: new OkObjectResult(document.RootElement.Clone())); + + AssertSuccess(result, "Object JSON results should be wrapped in an array."); + using JsonDocument response = JsonDocument.Parse(GetFirstText(result)); + JsonElement value = response.RootElement.GetProperty("value"); + Assert.AreEqual(JsonValueKind.Array, value.ValueKind); + Assert.AreEqual(1, value.GetArrayLength()); + Assert.AreEqual(1, value[0].GetProperty("id").GetInt32()); + } + + [TestMethod] + public async Task ExecuteEntity_UnknownResult_ReturnsEmptyValue() + { + CallToolResult result = await ExecuteWithMockedEngineAsync( + TEST_ENTITY, + new(), + null, + queryOutcome: new NoContentResult()); + + AssertSuccess(result, "Unknown result types should produce an empty value array."); + StringAssert.Contains(GetFirstText(result), "\"value\": []"); + } + + #endregion + #region Helpers /// @@ -293,14 +531,22 @@ private static async Task ExecuteWithMockedEngineAsync( string entityName, Dictionary dbParameters, Dictionary? userParameters, - Action? captureContext = null) + Action? captureContext = null, + object? queryOutcome = null, + DatabaseObject? metadataObject = null, + bool includeHttpContext = true, + bool authorizeOperation = true) { IServiceProvider sp = BuildServiceProvider( entityName: entityName, sourceObject: SP_SOURCE_OBJECT, sourceType: EntitySourceType.StoredProcedure, dbParameters: dbParameters, - captureContext: captureContext); + captureContext: captureContext, + queryOutcome: queryOutcome, + metadataObject: metadataObject, + includeHttpContext: includeHttpContext, + authorizeOperation: authorizeOperation); ExecuteEntityTool tool = new(); @@ -324,7 +570,11 @@ private static IServiceProvider BuildServiceProvider( string sourceObject, EntitySourceType sourceType, Dictionary dbParameters, - Action? captureContext = null) + Action? captureContext = null, + object? queryOutcome = null, + DatabaseObject? metadataObject = null, + bool includeHttpContext = true, + bool authorizeOperation = true) { Entity entity = new( Source: new(sourceObject, sourceType, Parameters: null, KeyFields: null), @@ -368,17 +618,20 @@ private static IServiceProvider BuildServiceProvider( mockAuthResolver .Setup(x => x.AreRoleAndOperationDefinedForEntity( It.IsAny(), It.IsAny(), It.IsAny())) - .Returns(true); + .Returns(authorizeOperation); services.AddSingleton(mockAuthResolver.Object); // Mock HttpContext with anonymous role header DefaultHttpContext httpContext = new(); httpContext.Request.Headers[AuthorizationResolver.CLIENT_ROLE_HEADER] = "anonymous"; - IHttpContextAccessor httpContextAccessor = new HttpContextAccessor { HttpContext = httpContext }; + IHttpContextAccessor httpContextAccessor = new HttpContextAccessor + { + HttpContext = includeHttpContext ? httpContext : null + }; services.AddSingleton(httpContextAccessor); // Mock metadata provider with DB object - DatabaseObject dbObject = sourceType == EntitySourceType.StoredProcedure + DatabaseObject dbObject = metadataObject ?? (sourceType == EntitySourceType.StoredProcedure ? new DatabaseStoredProcedure("dbo", sourceObject) { SourceType = EntitySourceType.StoredProcedure, @@ -387,7 +640,7 @@ private static IServiceProvider BuildServiceProvider( Parameters = dbParameters } } - : new DatabaseTable("dbo", sourceObject) { SourceType = EntitySourceType.Table }; + : new DatabaseTable("dbo", sourceObject) { SourceType = EntitySourceType.Table }); Mock mockSqlMetadataProvider = new(); mockSqlMetadataProvider @@ -403,15 +656,33 @@ private static IServiceProvider BuildServiceProvider( // Mock query engine factory Mock mockQueryEngine = new(); - mockQueryEngine - .Setup(x => x.ExecuteAsync(It.IsAny(), It.IsAny())) - .Returns((StoredProcedureRequestContext ctx, string ds) => + if (queryOutcome is Exception exception) + { + mockQueryEngine + .Setup(x => x.ExecuteAsync(It.IsAny(), It.IsAny())) + .ThrowsAsync(exception); + } + else + { + mockQueryEngine + .Setup(x => x.ExecuteAsync(It.IsAny(), It.IsAny())) + .Returns((StoredProcedureRequestContext ctx, string ds) => { captureContext?.Invoke(ctx); - // Return empty JSON array result - using JsonDocument doc = JsonDocument.Parse("[]"); - return Task.FromResult(new OkObjectResult(doc.RootElement.Clone())); + IActionResult result; + if (queryOutcome is IActionResult configuredResult) + { + result = configuredResult; + } + else + { + using JsonDocument doc = JsonDocument.Parse("[]"); + result = new OkObjectResult(doc.RootElement.Clone()); + } + + return Task.FromResult(result); }); + } Mock mockQueryEngineFactory = new(); mockQueryEngineFactory @@ -442,6 +713,27 @@ private static string GetFirstText(CallToolResult result) : string.Empty; } + private static void AssertError(CallToolResult result, string expectedError) + { + Assert.IsTrue(result.IsError == true, GetFirstText(result)); + StringAssert.Contains(GetFirstText(result), expectedError); + } + + private sealed class FakeDbException : DbException + { + public FakeDbException() + { + } + + public FakeDbException(string message) : base(message) + { + } + + public FakeDbException(string message, Exception innerException) : base(message, innerException) + { + } + } + #endregion } } diff --git a/src/Service.Tests/Mcp/McpArgumentParserTests.cs b/src/Service.Tests/Mcp/McpArgumentParserTests.cs index 6054cd7983..5813bbc6b0 100644 --- a/src/Service.Tests/Mcp/McpArgumentParserTests.cs +++ b/src/Service.Tests/Mcp/McpArgumentParserTests.cs @@ -80,6 +80,13 @@ public void TryParseEntityAndData_DataNotObject_ReturnsFalse() StringAssert.Contains(error, "JSON object"); } + [TestMethod] + public void TryParseEntityAndData_InvalidEntityReturnsEarly() + { + Assert.IsFalse(McpArgumentParser.TryParseEntityAndData(Parse("{}"), out _, out _, out string error)); + StringAssert.Contains(error, "entity"); + } + [TestMethod] public void TryParseEntityAndKeys_Valid_ReturnsTrue() { @@ -169,6 +176,23 @@ public void TryParseEntityKeysAndFields_EmptyFields_ReturnsFalse() StringAssert.Contains(error, "field must be provided"); } + [TestMethod] + public void TryParseEntityKeysAndFields_InvalidKeysReturnsEarly() + { + Assert.IsFalse(McpArgumentParser.TryParseEntityKeysAndFields( + Parse(@"{ ""entity"": ""Book"" }"), out _, out _, out _, out string error)); + StringAssert.Contains(error, "keys"); + } + + [TestMethod] + public void TryParseEntityKeysAndFields_FieldsNotObjectReturnsFalse() + { + Assert.IsFalse(McpArgumentParser.TryParseEntityKeysAndFields( + Parse(@"{ ""entity"": ""Book"", ""keys"": { ""id"": 1 }, ""fields"": [] }"), + out _, out _, out _, out string error)); + StringAssert.Contains(error, "JSON object"); + } + [TestMethod] public void TryParseExecuteArguments_NonObjectRoot_ReturnsFalse() { @@ -217,5 +241,31 @@ public void TryParseExecuteArguments_NoParameters_ReturnsEmptyDictionary() Assert.AreEqual("GetBooks", entity); Assert.AreEqual(0, parameters.Count); } + + [TestMethod] + public void TryParseExecuteArguments_InvalidEntityReturnsEarly() + { + Assert.IsFalse(McpArgumentParser.TryParseExecuteArguments(Parse("{}"), out _, out _, out string error)); + StringAssert.Contains(error, "entity"); + } + + [TestMethod] + public void TryParseExecuteArguments_NonObjectParametersAreIgnored() + { + Assert.IsTrue(McpArgumentParser.TryParseExecuteArguments( + Parse(@"{ ""entity"": ""GetBooks"", ""parameters"": [] }"), out _, out Dictionary parameters, out _)); + Assert.AreEqual(0, parameters.Count); + } + + [TestMethod] + public void TryParseExecuteArguments_FalseAndCompositeParametersAreConverted() + { + Assert.IsTrue(McpArgumentParser.TryParseExecuteArguments( + Parse(@"{ ""entity"": ""GetBooks"", ""parameters"": { ""disabled"": false, ""items"": [1] } }"), + out _, out Dictionary parameters, out _)); + + Assert.AreEqual(false, parameters["disabled"]); + Assert.AreEqual("[1]", parameters["items"]); + } } } diff --git a/src/Service.Tests/Mcp/McpAuthorizationHelperTests.cs b/src/Service.Tests/Mcp/McpAuthorizationHelperTests.cs index c5d7ef0443..7c5dccfbb7 100644 --- a/src/Service.Tests/Mcp/McpAuthorizationHelperTests.cs +++ b/src/Service.Tests/Mcp/McpAuthorizationHelperTests.cs @@ -1,6 +1,8 @@ // Copyright (c) Microsoft Corporation. // Licensed under the MIT License. +using System; +using System.Collections.Generic; using Azure.DataApiBuilder.Auth; using Azure.DataApiBuilder.Config.ObjectModel; using Azure.DataApiBuilder.Core.Authorization; @@ -116,5 +118,32 @@ public void TryResolveAuthorizedRole_NoAllowedRole_ReturnsFalse() Assert.IsNull(effectiveRole); StringAssert.Contains(error, "permission"); } + + [TestMethod] + public void AreColumnsAuthorizedForOperation_EmptyColumnsRequireNoAuthorization() + { + Mock resolver = new(); + + Assert.IsTrue(McpAuthorizationHelper.AreColumnsAuthorizedForOperation( + resolver.Object, "Book", "writer", EntityActionOperation.Create, Array.Empty(), out string error)); + Assert.AreEqual(string.Empty, error); + resolver.VerifyNoOtherCalls(); + } + + [DataTestMethod] + [DataRow(true)] + [DataRow(false)] + public void AreColumnsAuthorizedForOperation_ReturnsResolverDecision(bool allowed) + { + Mock resolver = new(); + resolver.Setup(r => r.AreColumnsAllowedForOperation( + "Book", "writer", EntityActionOperation.Create, It.IsAny>())).Returns(allowed); + + bool result = McpAuthorizationHelper.AreColumnsAuthorizedForOperation( + resolver.Object, "Book", "writer", EntityActionOperation.Create, new[] { "title" }, out string error); + + Assert.AreEqual(allowed, result); + Assert.AreEqual(allowed, string.IsNullOrEmpty(error)); + } } } diff --git a/src/Service.Tests/Mcp/McpMetadataHelperTests.cs b/src/Service.Tests/Mcp/McpMetadataHelperTests.cs index 1602bdc6bc..a26e936e21 100644 --- a/src/Service.Tests/Mcp/McpMetadataHelperTests.cs +++ b/src/Service.Tests/Mcp/McpMetadataHelperTests.cs @@ -3,12 +3,14 @@ using System; using System.Collections.Generic; +using System.Net; using System.Threading; using Azure.DataApiBuilder.Config.DatabasePrimitives; using Azure.DataApiBuilder.Config.ObjectModel; using Azure.DataApiBuilder.Core.Services; using Azure.DataApiBuilder.Core.Services.MetadataProviders; using Azure.DataApiBuilder.Mcp.Utils; +using Azure.DataApiBuilder.Service.Exceptions; using Microsoft.Extensions.DependencyInjection; using Microsoft.VisualStudio.TestTools.UnitTesting; using Moq; @@ -135,6 +137,27 @@ public void TryResolveDatabaseObject_Failure_ReturnsNull() StringAssert.Contains(error, "is not defined in the configuration"); } + [DataTestMethod] + [DataRow(DataApiBuilderException.SubStatusCodes.DataSourceNotFound, "is not defined in the configuration")] + [DataRow(DataApiBuilderException.SubStatusCodes.UnexpectedError, "factory failure")] + public void TryResolveMetadata_FactoryExceptionReturnsError( + DataApiBuilderException.SubStatusCodes subStatusCode, + string expectedError) + { + RuntimeConfig config = CreateConfig(includeBookEntity: true); + Mock factory = new(); + factory.Setup(x => x.GetMetadataProvider(It.IsAny())).Throws( + new DataApiBuilderException("factory failure", HttpStatusCode.InternalServerError, subStatusCode)); + ServiceCollection services = new(); + services.AddSingleton(factory.Object); + + bool result = McpMetadataHelper.TryResolveMetadata( + ENTITY_NAME, config, services.BuildServiceProvider(), out _, out _, out _, out string error); + + Assert.IsFalse(result); + StringAssert.Contains(error, expectedError); + } + #region Helpers private static RuntimeConfig CreateConfig(bool includeBookEntity) diff --git a/src/Service.Tests/Mcp/McpResponseBuilderTests.cs b/src/Service.Tests/Mcp/McpResponseBuilderTests.cs index 6b502a9d63..429222ce99 100644 --- a/src/Service.Tests/Mcp/McpResponseBuilderTests.cs +++ b/src/Service.Tests/Mcp/McpResponseBuilderTests.cs @@ -121,5 +121,15 @@ public void GetJsonValue_Number_PreservesNumericValue() Assert.AreEqual(42m, Convert.ToDecimal(McpResponseBuilder.GetJsonValue(number))); } + + [TestMethod] + public void GetJsonValue_ReturnsTypedAndRawValues() + { + Assert.AreEqual("text", McpResponseBuilder.GetJsonValue(JsonDocument.Parse("\"text\"").RootElement)); + Assert.AreEqual(true, McpResponseBuilder.GetJsonValue(JsonDocument.Parse("true").RootElement)); + Assert.AreEqual(false, McpResponseBuilder.GetJsonValue(JsonDocument.Parse("false").RootElement)); + Assert.IsNull(McpResponseBuilder.GetJsonValue(JsonDocument.Parse("null").RootElement)); + Assert.AreEqual("{\"id\":1}", McpResponseBuilder.GetJsonValue(JsonDocument.Parse("{\"id\":1}").RootElement)); + } } } diff --git a/src/Service.Tests/Mcp/McpServerConfigurationTests.cs b/src/Service.Tests/Mcp/McpServerConfigurationTests.cs new file mode 100644 index 0000000000..b52417e974 --- /dev/null +++ b/src/Service.Tests/Mcp/McpServerConfigurationTests.cs @@ -0,0 +1,34 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +using System; +using Azure.DataApiBuilder.Mcp.Core; +using Microsoft.Extensions.DependencyInjection; +using Microsoft.Extensions.Options; +using Microsoft.VisualStudio.TestTools.UnitTesting; +using ModelContextProtocol.Server; + +namespace Azure.DataApiBuilder.Service.Tests.Mcp +{ + [TestClass] + public class McpServerConfigurationTests + { + [DataTestMethod] + [DataRow(null, null)] + [DataRow(" ", null)] + [DataRow("Use discovered tools only.", "Use discovered tools only.")] + public void ConfigureMcpServer_ConfiguresServerOptions(string? instructions, string? expectedInstructions) + { + ServiceCollection services = new(); + + IServiceProvider provider = services.ConfigureMcpServer(instructions).BuildServiceProvider(); + McpServerOptions options = provider.GetRequiredService>().Value; + + Assert.AreEqual(McpProtocolDefaults.MCP_SERVER_NAME, options.ServerInfo!.Name); + Assert.AreEqual(McpProtocolDefaults.MCP_SERVER_VERSION, options.ServerInfo.Version); + Assert.IsNotNull(options.Capabilities); + Assert.IsNotNull(options.Capabilities.Tools); + Assert.AreEqual(expectedInstructions, options.ServerInstructions); + } + } +} diff --git a/src/Service.Tests/UnitTests/ApplicationNameTelemetryTests.cs b/src/Service.Tests/UnitTests/ApplicationNameTelemetryTests.cs index 4bf662640e..9a7a94207a 100644 --- a/src/Service.Tests/UnitTests/ApplicationNameTelemetryTests.cs +++ b/src/Service.Tests/UnitTests/ApplicationNameTelemetryTests.cs @@ -416,6 +416,87 @@ public void Decode_TruncatedPayload_DoesNotThrowAndDecodesPartial() Assert.IsTrue(lines.Any(l => l.StartsWith("Version:", StringComparison.Ordinal))); } + [TestMethod] + public void Decode_TruncatedBeforeRuntimeAndEntity_SkipsMissingSections() + { + IReadOnlyList lines = ApplicationNameTelemetry.Decode("dab_oss_1.2.3+RTSN+"); + + Assert.IsTrue(lines.Any(l => l.Contains("Protocol: R (REST)", StringComparison.Ordinal))); + Assert.IsFalse(lines.Any(l => l.StartsWith("Runtime >", StringComparison.Ordinal))); + Assert.IsFalse(lines.Any(l => l.StartsWith("Entity >", StringComparison.Ordinal))); + } + + [TestMethod] + public void Decode_NewerPayloadFlag_IsReportedAsUnrecognizedPosition() + { + IReadOnlyList lines = ApplicationNameTelemetry.Decode("dab_oss_1.2.3+XXXXZ|||+"); + + Assert.IsTrue(lines.Any(l => l.Contains("Context > [position 5]: Z", StringComparison.Ordinal))); + } + + [DataTestMethod] + [DataRow("RTSN", "Protocol: R (REST)", "Object: T (Table)", "Source: S (SQL)", "Role: N (Anonymous)")] + [DataRow("GVDA", "Protocol: G (GraphQL)", "Object: V (View)", "Source: D (DWSQL)", "Role: A (Authenticated)")] + [DataRow("MSPC", "Protocol: M (MCP)", "Object: S (Stored Procedure)", "Source: P (PostgreSQL)", "Role: C (Custom)")] + [DataRow("ZPMZ", "Protocol: Z (unrecognized)", "Object: P (Persisted Document)", "Source: M (MySQL)", "Role: Z (unrecognized)")] + [DataRow("XXXZ", "Protocol: X (not applicable)", "Object: X (not applicable)", "Source: X (not applicable)", "Role: Z (unrecognized)")] + [DataRow("ZZCZ", "Protocol: Z (unrecognized)", "Object: Z (unrecognized)", "Source: C (CosmosDB)", "Role: Z (unrecognized)")] + public void Decode_ContextAlphabet_DescribesEveryValue( + string context, + string expectedProtocol, + string expectedObject, + string expectedSource, + string expectedRole) + { + IReadOnlyList lines = ApplicationNameTelemetry.Decode($"dab_oss_1.2.3+{context}|||+"); + string decoded = string.Join('\n', lines); + + StringAssert.Contains(decoded, expectedProtocol); + StringAssert.Contains(decoded, expectedObject); + StringAssert.Contains(decoded, expectedSource); + StringAssert.Contains(decoded, expectedRole); + } + + [DataTestMethod] + [DataRow('U', "Unauthenticated")] + [DataRow('E', "EntraId")] + [DataRow('C', "Custom")] + [DataRow('S', "Simulator")] + [DataRow('A', "AppService")] + [DataRow('W', "StaticWebApps")] + [DataRow('Z', "unrecognized")] + public void Decode_RuntimeAlphabet_DescribesAuthProvider(char authProvider, string expectedDescription) + { + char[] runtime = new string('M', 20).ToCharArray(); + runtime[3] = 'Z'; + runtime[17] = authProvider; + + IReadOnlyList lines = ApplicationNameTelemetry.Decode( + $"dab_oss_1.2.3+XXXX||{new string(runtime)}|+"); + string decoded = string.Join('\n', lines); + + StringAssert.Contains(decoded, $"auth.provider: {authProvider} ({expectedDescription})"); + StringAssert.Contains(decoded, "runtime.host.mode: Z (unrecognized)"); + } + + [TestMethod] + public void Decode_UnknownFlagValue_IsReportedAsUnrecognized() + { + IReadOnlyList lines = ApplicationNameTelemetry.Decode("dab_oss_1.2.3+XXXX||Z|+"); + + Assert.IsTrue(lines.Any(l => l.Contains("runtime.rest.enabled: Z (unrecognized)", StringComparison.Ordinal))); + } + + [TestMethod] + public void EncodeTelemetryString_CosmosPostgreSql_UsesCosmosSourceCode() + { + string telemetry = ApplicationNameTelemetry.EncodeTelemetryString( + BuildConfig(), + Source(DatabaseType.CosmosDB_PostgreSQL)); + + Assert.AreEqual('C', Sections(telemetry).context[2]); + } + /// Verifies the friendly response when no supported telemetry marker is present. [TestMethod] public void Decode_NoMarker_ReturnsFriendlyMessage() diff --git a/src/Service.Tests/UnitTests/AutoentityConverterCoverageTests.cs b/src/Service.Tests/UnitTests/AutoentityConverterCoverageTests.cs new file mode 100644 index 0000000000..f8193642fb --- /dev/null +++ b/src/Service.Tests/UnitTests/AutoentityConverterCoverageTests.cs @@ -0,0 +1,64 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +using System.Text.Json; +using Azure.DataApiBuilder.Config; +using Azure.DataApiBuilder.Config.ObjectModel; +using Microsoft.VisualStudio.TestTools.UnitTesting; + +namespace Azure.DataApiBuilder.Service.Tests.UnitTests +{ + [TestClass] + public class AutoentityConverterCoverageTests + { + private static JsonSerializerOptions Options => RuntimeConfigLoader.GetSerializationOptions(); + + [TestMethod] + public void Autoentity_WithPatterns_RoundTripsUserProvidedValues() + { + const string Json = """ + { + "patterns": { + "include": ["dbo.*", null], + "exclude": ["dbo.internal_*"], + "name": "generated_{object}" + }, + "permissions": [] + } + """; + + Autoentity? autoentity = JsonSerializer.Deserialize(Json, Options); + string serialized = JsonSerializer.Serialize(autoentity, Options); + using JsonDocument document = JsonDocument.Parse(serialized); + JsonElement patterns = document.RootElement.GetProperty("patterns"); + + Assert.IsNotNull(autoentity); + CollectionAssert.AreEqual(new[] { "dbo.*" }, autoentity.Patterns.Include); + CollectionAssert.AreEqual(new[] { "dbo.internal_*" }, autoentity.Patterns.Exclude); + Assert.AreEqual("generated_{object}", autoentity.Patterns.Name); + Assert.AreEqual(1, patterns.GetProperty("include").GetArrayLength()); + Assert.AreEqual("generated_{object}", patterns.GetProperty("name").GetString()); + } + + [DataTestMethod] + [DataRow("42", typeof(Autoentity))] + [DataRow("42", typeof(AutoentityPatterns))] + [DataRow("42", typeof(AutoentityTemplate))] + [DataRow("{\"unexpected\":true}", typeof(Autoentity))] + [DataRow("{\"unexpected\":true}", typeof(AutoentityPatterns))] + [DataRow("{\"unexpected\":true}", typeof(AutoentityTemplate))] + public void AutoentityConverters_InvalidInputThrows(string json, System.Type targetType) + { + Assert.ThrowsException(() => JsonSerializer.Deserialize(json, targetType, Options)); + } + + [DataTestMethod] + [DataRow("{\"include\":true}")] + [DataRow("{\"exclude\":42}")] + public void AutoentityPatterns_NonArrayPatternThrows(string json) + { + Assert.ThrowsException(() => + JsonSerializer.Deserialize(json, Options)); + } + } +} diff --git a/src/Service.Tests/UnitTests/BaseSqlQueryStructureHelperTests.cs b/src/Service.Tests/UnitTests/BaseSqlQueryStructureHelperTests.cs new file mode 100644 index 0000000000..d7a399e06e --- /dev/null +++ b/src/Service.Tests/UnitTests/BaseSqlQueryStructureHelperTests.cs @@ -0,0 +1,228 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +using System; +using System.Collections.Generic; +using System.Reflection; +using Azure.DataApiBuilder.Auth; +using Azure.DataApiBuilder.Config.DatabasePrimitives; +using Azure.DataApiBuilder.Config.ObjectModel; +using Azure.DataApiBuilder.Core.Models; +using Azure.DataApiBuilder.Core.Resolvers; +using Azure.DataApiBuilder.Core.Services; +using Azure.DataApiBuilder.Service.Exceptions; +using HotChocolate.Language; +using Microsoft.VisualStudio.TestTools.UnitTesting; +using Moq; + +namespace Azure.DataApiBuilder.Service.Tests.UnitTests +{ + [TestClass] + public class BaseSqlQueryStructureHelperTests + { + [DataTestMethod] + [DataRow("text", typeof(string), "text")] + [DataRow("255", typeof(byte), (byte)255)] + [DataRow("-12", typeof(short), (short)-12)] + [DataRow("123", typeof(int), 123)] + [DataRow("123456789", typeof(long), 123456789L)] + [DataRow("true", typeof(bool), true)] + [DataRow("7d4ee078-a85c-4a95-82b6-4bf6c3f3cfe8", typeof(Guid), "7d4ee078-a85c-4a95-82b6-4bf6c3f3cfe8")] + public void ParseParamAsSystemType_ParsesSupportedScalarTypes(string value, Type targetType, object expected) + { + object result = InvokeParse(value, targetType); + + if (targetType == typeof(Guid)) + { + Assert.AreEqual(Guid.Parse((string)expected), result); + } + else + { + Assert.AreEqual(expected, result); + } + } + + [TestMethod] + public void ParseParamAsSystemType_ParsesBinaryDateAndFloatingPointTypes() + { + CollectionAssert.AreEqual(new byte[] { 1, 2, 3 }, (byte[])InvokeParse("AQID", typeof(byte[]))); + Assert.AreEqual(1.25f, InvokeParse("1.25", typeof(float))); + Assert.AreEqual(2.5d, InvokeParse("2.5", typeof(double))); + Assert.AreEqual(3.75m, InvokeParse("3.75", typeof(decimal))); + Assert.AreEqual(new TimeOnly(12, 34, 56), InvokeParse("12:34:56", typeof(TimeOnly))); + } + + [TestMethod] + public void ParseParamAsSystemType_ParsesDatesAndArrays() + { + DateTime dateTime = (DateTime)InvokeParse("2025-01-02T12:00:00+03:00", typeof(DateTime)); + Assert.AreEqual(DateTimeKind.Utc, dateTime.Kind); + Assert.AreEqual(9, dateTime.Hour); + + DateTimeOffset offset = (DateTimeOffset)InvokeParse("2025-01-02T12:00:00+03:00", typeof(DateTimeOffset)); + Assert.AreEqual(TimeSpan.FromHours(3), offset.Offset); + + object[] values = (object[])InvokeParse("[1.5,2.25]", typeof(float[])); + CollectionAssert.AreEqual(new object[] { 1.5f, 2.25f }, values); + } + + [TestMethod] + public void ParseParamAsSystemType_UnsupportedOrMalformedArray_Throws() + { + TargetInvocationException unsupported = Assert.ThrowsException( + () => InvokeParse("value", typeof(Uri))); + Assert.IsInstanceOfType(unsupported.InnerException); + + TargetInvocationException malformed = Assert.ThrowsException( + () => InvokeParse("not-json", typeof(float[]))); + Assert.IsInstanceOfType(malformed.InnerException); + } + + [TestMethod] + public void GetSubArgumentNamesFromGQLMutArguments_ReturnsObjectFieldNames() + { + Dictionary parameters = new() + { + ["item"] = new List + { + new("id", new IntValueNode(1)), + new("title", new StringValueNode("book")) + } + }; + + List names = BaseSqlQueryStructure.GetSubArgumentNamesFromGQLMutArguments("item", parameters); + + CollectionAssert.AreEqual(new[] { "id", "title" }, names); + } + + [DataTestMethod] + [DataRow(true)] + [DataRow(false)] + public void GetSubArgumentNamesFromGQLMutArguments_InvalidArguments_Throw(bool includeWrongFormat) + { + Dictionary parameters = new(); + if (includeWrongFormat) + { + parameters["item"] = "unexpected"; + } + + DataApiBuilderException exception = Assert.ThrowsException(() => + BaseSqlQueryStructure.GetSubArgumentNamesFromGQLMutArguments("item", parameters)); + StringAssert.Contains(exception.Message, includeWrongFormat ? "Unexpected" : "Expected"); + } + + [TestMethod] + public void GetColumnSystemType_ReturnsKnownTypeAndRejectsUnknownColumn() + { + (TestSqlQueryStructure structure, _) = CreateStructure(EntitySourceType.Table, isDevelopment: false); + + Assert.AreEqual(typeof(int), structure.GetColumnSystemType("id")); + Assert.ThrowsException(() => structure.GetColumnSystemType("missing")); + } + + [TestMethod] + public void AddJoinPredicatesForRelationship_MissingSelfRelationshipThrows() + { + (TestSqlQueryStructure structure, Mock metadata) = CreateStructure(EntitySourceType.Table, false); + metadata.SetupGet(x => x.RelationshipToFkDefinition).Returns(new Dictionary()); + + Assert.ThrowsException(() => structure.AddJoinPredicatesForRelationship( + new EntityRelationshipKey("Book", "related"), "Book", "table1", structure)); + } + + [TestMethod] + public void AddJoinPredicatesForRelatedEntity_MissingRelationshipThrows() + { + (TestSqlQueryStructure structure, Mock metadata) = CreateStructure(EntitySourceType.Table, false); + DatabaseTable relatedTable = new("dbo", "authors") { TableDefinition = new SourceDefinition() }; + metadata.SetupGet(x => x.EntityToDatabaseObject).Returns(new Dictionary + { + ["Book"] = structure.DatabaseObject, + ["Author"] = relatedTable + }); + metadata.Setup(x => x.GetSourceDefinition("Author")).Returns(relatedTable.TableDefinition); + + Assert.ThrowsException(() => + structure.AddJoinPredicatesForRelatedEntity("Author", "table1", structure)); + } + + [TestMethod] + public void ProcessOdataClause_NullPolicyStoresNullAndMissingOperationReturnsNull() + { + (TestSqlQueryStructure structure, _) = CreateStructure(EntitySourceType.Table, false); + + structure.ProcessOdataClause(null, EntityActionOperation.Read); + + Assert.IsTrue(structure.DbPolicyPredicatesForOperations.ContainsKey(EntityActionOperation.Read)); + Assert.IsNull(structure.GetDbPolicyForOperation(EntityActionOperation.Read)); + Assert.IsNull(structure.GetDbPolicyForOperation(EntityActionOperation.Create)); + } + + [DataTestMethod] + [DataRow(EntitySourceType.StoredProcedure, true, "stored procedure parameter")] + [DataRow(EntitySourceType.Table, true, "column")] + [DataRow(EntitySourceType.Table, false, "publicId")] + [DataRow(EntitySourceType.StoredProcedure, false, "id")] + public void GetParamAsSystemType_InvalidValueUsesSafeContextualMessage( + EntitySourceType sourceType, + bool isDevelopment, + string expectedMessagePart) + { + (TestSqlQueryStructure structure, Mock metadata) = CreateStructure(sourceType, isDevelopment); + metadata.Setup(x => x.TryGetExposedColumnName("Book", "id", out It.Ref.IsAny)) + .Returns((string _, string _, out string? name) => + { + name = "publicId"; + return true; + }); + + DataApiBuilderException exception = Assert.ThrowsException(() => + structure.ParseWithContext("not-an-int", "id", typeof(int))); + + StringAssert.Contains(exception.Message, expectedMessagePart); + } + + private static object InvokeParse(string value, Type targetType) + { + MethodInfo method = typeof(BaseSqlQueryStructure).GetMethod( + "ParseParamAsSystemType", + BindingFlags.Static | BindingFlags.NonPublic)!; + return method.Invoke(null, new object[] { value, targetType })!; + } + + private static (TestSqlQueryStructure Structure, Mock Metadata) CreateStructure( + EntitySourceType sourceType, + bool isDevelopment) + { + SourceDefinition sourceDefinition = new(); + sourceDefinition.Columns["id"] = new ColumnDefinition(typeof(int)); + StoredProcedureDefinition storedProcedureDefinition = new(); + storedProcedureDefinition.Columns["id"] = sourceDefinition.Columns["id"]; + DatabaseObject databaseObject = sourceType is EntitySourceType.StoredProcedure + ? new DatabaseStoredProcedure("dbo", "books") { StoredProcedureDefinition = storedProcedureDefinition } + : new DatabaseTable("dbo", "books") { TableDefinition = sourceDefinition }; + databaseObject.SourceType = sourceType; + + Mock metadata = new(); + metadata.SetupGet(x => x.EntityToDatabaseObject).Returns(new Dictionary + { + ["Book"] = databaseObject + }); + metadata.Setup(x => x.IsDevelopmentMode()).Returns(isDevelopment); + metadata.Setup(x => x.GetSourceDefinition("Book")).Returns(databaseObject.SourceDefinition); + + return (new TestSqlQueryStructure(metadata.Object), metadata); + } + + private sealed class TestSqlQueryStructure : BaseSqlQueryStructure + { + public TestSqlQueryStructure(ISqlMetadataProvider metadataProvider) + : base(metadataProvider, new Mock().Object, null!, entityName: "Book") + { + } + + public object ParseWithContext(string value, string fieldName, Type type) => + GetParamAsSystemType(value, fieldName, type); + } + } +} diff --git a/src/Service.Tests/UnitTests/ConfigObjectModelCoverageTests.cs b/src/Service.Tests/UnitTests/ConfigObjectModelCoverageTests.cs new file mode 100644 index 0000000000..fef4479199 --- /dev/null +++ b/src/Service.Tests/UnitTests/ConfigObjectModelCoverageTests.cs @@ -0,0 +1,223 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +using System; +using System.Collections.Generic; +using System.Linq; +using Azure.DataApiBuilder.Config.ObjectModel; +using Azure.DataApiBuilder.Config.ObjectModel.Embeddings; +using Azure.DataApiBuilder.Service.Exceptions; +using Microsoft.VisualStudio.TestTools.UnitTesting; + +namespace Azure.DataApiBuilder.Service.Tests.UnitTests +{ + [TestClass] + public class ConfigObjectModelCoverageTests + { + [TestMethod] + public void EntityActionPolicy_ProcessesFieldsAndRejectsNullDatabasePolicy() + { + Assert.ThrowsException(() => new EntityActionPolicy().ProcessedDatabaseFields()); + Assert.AreEqual( + "id eq 1 and snake_case eq 2", + new EntityActionPolicy(Database: "@item.id eq 1 and @item.snake_case eq 2").ProcessedDatabaseFields()); + } + + [TestMethod] + public void RuntimeOptions_DisabledFeaturesReturnFalse() + { + RuntimeOptions options = new( + Rest: new RestRuntimeOptions(Enabled: false), + GraphQL: new GraphQLRuntimeOptions(Enabled: false), + Mcp: new McpRuntimeOptions(Enabled: false, Path: "/mcp", DmlTools: null), + Host: null, + Health: new RuntimeHealthCheckConfig(enabled: false)); + + Assert.IsFalse(options.IsCachingEnabled); + Assert.IsFalse(options.IsRestEnabled); + Assert.IsFalse(options.IsGraphQLEnabled); + Assert.IsFalse(options.IsMcpEnabled); + Assert.IsFalse(options.IsHealthCheckEnabled); + Assert.IsFalse(options.IsEmbeddingsConfigured); + } + + [TestMethod] + public void ChildConfigMetadata_RecordPropertiesRoundTrip() + { + HashSet entities = new() { "Book" }; + HashSet autoentities = new() { "AutoBook" }; + ChildConfigMetadata metadata = new("child.json", entities, autoentities, HasDataSource: true); + + Assert.AreEqual("child.json", metadata.FileName); + Assert.AreSame(entities, metadata.EntityNames); + Assert.AreSame(autoentities, metadata.AutoentityDefinitionNames); + Assert.IsTrue(metadata.HasDataSource); + } + + [TestMethod] + public void EmbeddingsHealthCheck_DefaultConstructorUsesDefaults() + { + EmbeddingsHealthCheckConfig health = new(); + + Assert.AreEqual(EmbeddingsHealthCheckConfig.DEFAULT_THRESHOLD_MS, health.ThresholdMs); + Assert.AreEqual(EmbeddingsHealthCheckConfig.DEFAULT_TEST_TEXT, health.TestText); + Assert.IsFalse(health.UserProvidedThresholdMs); + Assert.IsFalse(health.UserProvidedTestText); + Assert.IsFalse(health.UserProvidedExpectedDimensions); + } + + [TestMethod] + public void DatasourceHealthCheck_ConstructorsTrackDefaultsAndUserValues() + { + DatasourceHealthCheckConfig defaults = new(); + DatasourceHealthCheckConfig configured = new(enabled: true, name: "primary", thresholdMs: 42); + + Assert.IsTrue(defaults.ThresholdMs > 0); + Assert.IsFalse(defaults.UserProvidedThresholdMs); + Assert.AreEqual("primary", configured.Name); + Assert.AreEqual(42, configured.ThresholdMs); + Assert.IsTrue(configured.UserProvidedThresholdMs); + } + + [TestMethod] + public void EmbeddingsOptions_NullOptionalFeaturesUseDocumentedFallbacks() + { + EmbeddingsOptions options = new(EmbeddingProviderType.OpenAI, "https://example.com", "key"); + + Assert.IsFalse(options.IsHealthCheckEnabled); + Assert.IsFalse(options.IsEndpointEnabled); + Assert.IsFalse(options.IsChunkingEnabled); + Assert.IsTrue(options.IsCachingEnabled); + Assert.IsFalse(options.IsLevel2CacheEnabled); + } + + [TestMethod] + public void EmbeddingsEndpointOptions_DefaultConstructorDisablesEndpoint() + { + EmbeddingsEndpointOptions options = new(); + + Assert.IsFalse(options.Enabled); + Assert.IsFalse(options.UserProvidedEnabled); + } + + [TestMethod] + public void Entity_ConfiguredHealthValuesAreReturned() + { + Entity entity = CreateEntity(new EntityHealthCheckConfig(enabled: true, first: 7, thresholdMs: 42)); + + Assert.AreEqual(7, entity.EntityFirst); + Assert.AreEqual(42, entity.EntityThresholdMs); + } + + [TestMethod] + public void EntityRelationshipKey_EqualityHandlesNullIdentityAndValues() + { + EntityRelationshipKey key = new("Book", "publisher"); + + Assert.IsFalse(key.Equals(null)); + Assert.IsTrue(key.Equals(key)); + Assert.IsTrue(key.Equals(new EntityRelationshipKey("Book", "publisher"))); + Assert.IsFalse(key.Equals(new EntityRelationshipKey("Book", "author"))); + Assert.IsFalse(key.Equals(new object())); + Assert.AreEqual(key.GetHashCode(), new EntityRelationshipKey("Book", "publisher").GetHashCode()); + } + + [TestMethod] + public void RuntimeAutoentities_GenericAndNonGenericEnumerationReturnEntries() + { + RuntimeAutoentities autoentities = new(new Dictionary + { + ["Book"] = new Autoentity(Patterns: null, Template: null, Permissions: null) + }); + + Assert.AreEqual("Book", autoentities.Single().Key); + Assert.IsTrue(((System.Collections.IEnumerable)autoentities).GetEnumerator().MoveNext()); + } + + [TestMethod] + public void RuntimeEntities_MissingIndexerThrowsConfigurationError() + { + RuntimeEntities entities = new(new Dictionary()); + + DataApiBuilderException exception = Assert.ThrowsException(() => _ = entities["Missing"]); + Assert.AreEqual(DataApiBuilderException.SubStatusCodes.ConfigValidationError, exception.SubStatusCode); + } + + [TestMethod] + public void DataSource_DefaultsAndTypedOptionsCoverFallbacks() + { + DataSource defaults = new(DatabaseType.MSSQL, string.Empty); + Assert.IsTrue(defaults.IsDatasourceHealthEnabled); + Assert.IsTrue(defaults.DatasourceThresholdMs > 0); + Assert.IsFalse(defaults.IsUserDelegatedAuthEnabled); + Assert.IsNull(defaults.GetTypedOptions()); + + DataSource wrongTypes = new( + DatabaseType.MSSQL, + string.Empty, + new Dictionary { ["set-session-context"] = "not-a-bool" }); + Assert.IsFalse(wrongTypes.GetTypedOptions()!.SetSessionContext); + Assert.ThrowsException(() => wrongTypes.GetTypedOptions()); + StringAssert.Contains(wrongTypes.DatabaseTypeNotSupportedMessage, DatabaseType.MSSQL.ToString()); + } + + [TestMethod] + public void RuntimeConfig_MissingEntityCacheLookupsThrow() + { + RuntimeConfig config = new( + Schema: string.Empty, + DataSource: new DataSource(DatabaseType.MSSQL, string.Empty), + Entities: new RuntimeEntities(new Dictionary())); + + Assert.ThrowsException(() => config.GetEntityCacheEntryTtl("Missing")); + Assert.ThrowsException(() => config.GetEntityCacheEntryLevel("Missing")); + Assert.ThrowsException(() => config.IsEntityCachingEnabled("Missing")); + } + + [TestMethod] + public void RuntimeConfig_EntityCacheTtlOverridesGlobalDefault() + { + Entity entity = CreateEntity() with { Cache = new EntityCacheOptions(Enabled: true, TtlSeconds: 42) }; + RuntimeConfig config = new( + Schema: string.Empty, + DataSource: new DataSource(DatabaseType.MSSQL, string.Empty), + Entities: new RuntimeEntities(new Dictionary { ["Book"] = entity })); + + Assert.AreEqual(42, config.GetEntityCacheEntryTtl("Book")); + } + + [TestMethod] + public void RuntimeConfig_MultipleCreateEnabledForSupportedDatabase() + { + RuntimeConfig config = new( + Schema: string.Empty, + DataSource: new DataSource(DatabaseType.MSSQL, string.Empty), + Runtime: new RuntimeOptions( + Rest: new RestRuntimeOptions(), + GraphQL: new GraphQLRuntimeOptions( + MultipleMutationOptions: new MultipleMutationOptions(new MultipleCreateOptions(true))), + Mcp: null, + Host: null), + Entities: new RuntimeEntities(new Dictionary())); + + Assert.IsTrue(config.IsMultipleCreateOperationEnabled()); + } + + private static Entity CreateEntity(EntityHealthCheckConfig? health = null) + { + return new Entity( + Source: new EntitySource("books", EntitySourceType.Table, null, null), + GraphQL: new EntityGraphQLOptions("Book", "Books"), + Fields: null, + Rest: new EntityRestOptions(), + Permissions: Array.Empty(), + Mappings: null, + Relationships: null, + Health: health); + } + + private sealed class UnsupportedOptions : IDataSourceOptions + { + } + } +} diff --git a/src/Service.Tests/UnitTests/ConfigureJwtBearerOptionsTests.cs b/src/Service.Tests/UnitTests/ConfigureJwtBearerOptionsTests.cs new file mode 100644 index 0000000000..a6ee19f470 --- /dev/null +++ b/src/Service.Tests/UnitTests/ConfigureJwtBearerOptionsTests.cs @@ -0,0 +1,88 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +using System.Collections.Generic; +using System.IO.Abstractions.TestingHelpers; +using Azure.DataApiBuilder.Config; +using Azure.DataApiBuilder.Config.ObjectModel; +using Azure.DataApiBuilder.Core.Configurations; +using Azure.DataApiBuilder.Service; +using Microsoft.AspNetCore.Authentication.JwtBearer; +using Microsoft.VisualStudio.TestTools.UnitTesting; + +namespace Azure.DataApiBuilder.Service.Tests.UnitTests +{ + [TestClass] + public class ConfigureJwtBearerOptionsTests + { + [TestMethod] + public void Configure_HotReloadDisabled_DoesNotChangeOptions() + { + RuntimeConfigProvider provider = CreateProvider( + CreateConfig(HostMode.Production, new AuthenticationOptions("AzureAD", new JwtOptions("aud", "https://issuer")))); + provider.IsLateConfigured = true; + JwtBearerOptions options = new() { MapInboundClaims = true, Audience = "original" }; + + new ConfigureJwtBearerOptions(provider).Configure("Bearer", options); + + Assert.IsTrue(options.MapInboundClaims); + Assert.AreEqual("original", options.Audience); + } + + [TestMethod] + public void Configure_MissingAuthentication_DoesNotChangeOptions() + { + RuntimeConfigProvider provider = CreateProvider(CreateConfig(HostMode.Development, authentication: null)); + JwtBearerOptions options = new() { MapInboundClaims = true }; + + new ConfigureJwtBearerOptions(provider).Configure("Bearer", options); + + Assert.IsTrue(options.MapInboundClaims); + Assert.IsNull(options.Audience); + } + + [DataTestMethod] + [DataRow("Custom")] + [DataRow("AzureAD")] + [DataRow("EntraID")] + public void Configure_JwtAuthentication_UpdatesAllTokenOptions(string providerName) + { + RuntimeConfigProvider provider = CreateProvider(CreateConfig( + HostMode.Development, + new AuthenticationOptions(providerName, new JwtOptions("api://audience", "https://issuer.example")))); + JwtBearerOptions options = new() { MapInboundClaims = true }; + ConfigureJwtBearerOptions configurator = new(provider); + + configurator.Configure(options); + + Assert.IsFalse(options.MapInboundClaims); + Assert.AreEqual("api://audience", options.Audience); + Assert.AreEqual("https://issuer.example", options.Authority); + Assert.AreEqual("api://audience", options.TokenValidationParameters.ValidAudience); + Assert.AreEqual("https://issuer.example", options.TokenValidationParameters.ValidIssuer); + Assert.AreEqual(AuthenticationOptions.ROLE_CLAIM_TYPE, options.TokenValidationParameters.RoleClaimType); + } + + private static RuntimeConfigProvider CreateProvider(RuntimeConfig config) + { + FileSystemRuntimeConfigLoader loader = new(new MockFileSystem()) + { + RuntimeConfig = config + }; + return new RuntimeConfigProvider(loader); + } + + private static RuntimeConfig CreateConfig(HostMode mode, AuthenticationOptions? authentication) + { + return new RuntimeConfig( + Schema: "test-schema", + DataSource: new DataSource(DatabaseType.MSSQL, "Server=test;", null), + Runtime: new( + Rest: new(), + GraphQL: new(), + Mcp: null, + Host: new(Cors: null, Authentication: authentication, Mode: mode)), + Entities: new(new Dictionary())); + } + } +} diff --git a/src/Service.Tests/UnitTests/CosmosEngineHelperTests.cs b/src/Service.Tests/UnitTests/CosmosEngineHelperTests.cs new file mode 100644 index 0000000000..f87ab0b986 --- /dev/null +++ b/src/Service.Tests/UnitTests/CosmosEngineHelperTests.cs @@ -0,0 +1,183 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +using System; +using System.Collections.Generic; +using System.Linq; +using System.Reflection; +using System.Threading.Tasks; +using Azure.DataApiBuilder.Auth; +using Azure.DataApiBuilder.Config.ObjectModel; +using Azure.DataApiBuilder.Core.Authorization; +using Azure.DataApiBuilder.Core.Models; +using Azure.DataApiBuilder.Core.Resolvers; +using Azure.DataApiBuilder.Core.Services; +using Azure.DataApiBuilder.Service.Exceptions; +using HotChocolate.Language; +using HotChocolate.Resolvers; +using Microsoft.Extensions.Primitives; +using Microsoft.VisualStudio.TestTools.UnitTesting; +using Moq; +using Newtonsoft.Json.Linq; + +namespace Azure.DataApiBuilder.Service.Tests.UnitTests +{ + [TestClass] + public class CosmosEngineHelperTests + { + [DataTestMethod] + [DataRow(EntityActionOperation.UpdateGraphQL, EntityActionOperation.Update)] + [DataRow(EntityActionOperation.Patch, EntityActionOperation.Update)] + [DataRow(EntityActionOperation.Create, EntityActionOperation.Create)] + public void AuthorizeMutation_DelegatesColumnAuthorization( + EntityActionOperation operation, + EntityActionOperation delegatedOperation) + { + Mock authorization = new(); + authorization.Setup(x => x.AreColumnsAllowedForOperation( + "Book", It.IsAny(), delegatedOperation, It.IsAny>())).Returns(true); + CosmosMutationEngine engine = new(null!, null!, authorization.Object); + IDictionary parameters = new Dictionary + { + ["item"] = new List { new("title", "DAB") } + }; + + engine.AuthorizeMutation(CreateContext(), parameters, "Book", operation); + + authorization.Verify(x => x.AreColumnsAllowedForOperation( + "Book", It.IsAny(), delegatedOperation, It.Is>(c => c.Contains("title"))), Times.Once); + } + + [TestMethod] + public void AuthorizeMutation_DeleteDoesNotPerformColumnAuthorization() + { + Mock authorization = new(); + CosmosMutationEngine engine = new(null!, null!, authorization.Object); + + engine.AuthorizeMutation(CreateContext(), new Dictionary { ["id"] = "1" }, "Book", EntityActionOperation.Delete); + + authorization.VerifyNoOtherCalls(); + } + + [TestMethod] + public void AuthorizeMutation_DeniedColumnsThrow() + { + Mock authorization = new(); + authorization.Setup(x => x.AreColumnsAllowedForOperation( + It.IsAny(), It.IsAny(), It.IsAny(), It.IsAny>())).Returns(false); + CosmosMutationEngine engine = new(null!, null!, authorization.Object); + + Assert.ThrowsException(() => engine.AuthorizeMutation( + CreateContext(), + new Dictionary { ["item"] = new List { new("title", "DAB") } }, + "Book", + EntityActionOperation.Create)); + } + + [TestMethod] + public void AuthorizeMutation_UnsupportedOperationThrows() + { + CosmosMutationEngine engine = new(null!, null!, new Mock().Object); + + Assert.ThrowsException(() => engine.AuthorizeMutation( + CreateContext(), new Dictionary(), "Book", EntityActionOperation.Read)); + } + + [TestMethod] + public void ParseVariableInputItem_MapsNonNullProperties() + { + object? result = InvokeMutation("ParseVariableInputItem", new Dictionary + { + ["title"] = "DAB", + ["ignored"] = null, + ["nested"] = new Dictionary { ["id"] = 7 } + }); + + JObject parsed = (JObject)result!; + Assert.AreEqual("DAB", parsed["title"]!.Value()); + Assert.IsNull(parsed["ignored"]); + Assert.AreEqual(7, parsed["nested"]!["id"]!.Value()); + Assert.AreEqual("value", InvokeMutation("ParseVariableInputItem", "value")); + } + + [TestMethod] + public void ParseInlineInputItem_HandlesObjectListArrayAndPrimitiveValues() + { + JObject node = (JObject)InvokeMutation("ParseInlineInputItem", new ObjectFieldNode("title", "DAB"))!; + JObject list = (JObject)InvokeMutation("ParseInlineInputItem", new List + { + new("title", "DAB"), + new("count", 7) + })!; + Mock nestedElement = new(); + nestedElement.SetupGet(x => x.Kind).Returns(SyntaxKind.ObjectValue); + nestedElement.SetupGet(x => x.Value).Returns(new List { new("id", 2) }); + JArray array = (JArray)InvokeMutation("ParseInlineInputItem", new List + { + new StringValueNode("one"), + nestedElement.Object + })!; + + Assert.AreEqual("DAB", node["title"]!.Value()); + Assert.AreEqual("DAB", list["title"]!.Value()); + Assert.AreEqual(7, list["count"]!.Value()); + Assert.AreEqual("one", array[0]!.Value()); + Assert.AreEqual(2, array[1]!["id"]!.Value()); + Assert.AreEqual(9, InvokeMutation("ParseInlineInputItem", 9)); + } + + [TestMethod] + public void GeneratePatchOperations_CreatesLeafAndArrayOperations() + { + JObject input = JObject.Parse(@"{ ""name"": ""DAB"", ""nested"": { ""id"": 7 }, ""tags"": [""a""] }"); + List operations = new(); + + InvokeMutation("GeneratePatchOperations", input, string.Empty, operations); + + Assert.AreEqual(3, operations.Count); + CollectionAssert.AreEquivalent(new[] { "/name", "/nested/id", "/tags" }, operations.Select(o => o.Path).ToArray()); + } + + [DataTestMethod] + [DataRow(null, null)] + [DataRow("continuation", "Y29udGludWF0aW9u")] + public void Base64Helpers_RoundTrip(string? plain, string? encoded) + { + Assert.AreEqual(encoded, InvokeQuery("Base64Encode", plain)); + Assert.AreEqual(plain, InvokeQuery("Base64Decode", encoded)); + } + + [TestMethod] + public async Task UnsupportedCosmosEngineEntryPointsThrow() + { + CosmosMutationEngine mutation = new(null!, null!, new Mock().Object); + CosmosQueryEngine query = (CosmosQueryEngine)System.Runtime.CompilerServices.RuntimeHelpers.GetUninitializedObject(typeof(CosmosQueryEngine)); + + await Assert.ThrowsExceptionAsync(() => mutation.ExecuteAsync((RestRequestContext)null!)); + await Assert.ThrowsExceptionAsync(() => query.ExecuteAsync((FindRequestContext)null!)); + await Assert.ThrowsExceptionAsync(() => query.ExecuteAsync((StoredProcedureRequestContext)null!, string.Empty)); + } + + private static IMiddlewareContext CreateContext() + { + Mock context = new(); + context.SetupGet(x => x.ContextData).Returns(new Dictionary + { + [AuthorizationResolver.CLIENT_ROLE_HEADER] = new StringValues(AuthorizationResolver.ROLE_ANONYMOUS) + }); + return context.Object; + } + + private static object? InvokeMutation(string methodName, params object?[] args) + { + MethodInfo method = typeof(CosmosMutationEngine).GetMethod(methodName, BindingFlags.Static | BindingFlags.NonPublic)!; + return method.Invoke(null, args); + } + + private static object? InvokeQuery(string methodName, params object?[] args) + { + MethodInfo method = typeof(CosmosQueryEngine).GetMethod(methodName, BindingFlags.Static | BindingFlags.NonPublic)!; + return method.Invoke(null, args); + } + } +} diff --git a/src/Service.Tests/UnitTests/CosmosSqlMetadataProviderHelperTests.cs b/src/Service.Tests/UnitTests/CosmosSqlMetadataProviderHelperTests.cs new file mode 100644 index 0000000000..f6547532aa --- /dev/null +++ b/src/Service.Tests/UnitTests/CosmosSqlMetadataProviderHelperTests.cs @@ -0,0 +1,238 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +using System; +using System.Collections.Concurrent; +using System.Collections.Generic; +using System.Reflection; +using System.Runtime.CompilerServices; +using Azure.DataApiBuilder.Config.DatabasePrimitives; +using Azure.DataApiBuilder.Config.ObjectModel; +using Azure.DataApiBuilder.Core.Models; +using Azure.DataApiBuilder.Core.Parsers; +using Azure.DataApiBuilder.Core.Services.MetadataProviders; +using Azure.DataApiBuilder.Service.Exceptions; +using HotChocolate.Language; +using Microsoft.VisualStudio.TestTools.UnitTesting; +using Moq; +using System.IO.Abstractions; + +namespace Azure.DataApiBuilder.Service.Tests.UnitTests +{ + [TestClass] + public class CosmosSqlMetadataProviderHelperTests + { + [TestMethod] + public void InterfaceMembers_ThatCosmosDoesNotSupport_Throw() + { + CosmosSqlMetadataProvider provider = CreateProvider(); + + Assert.ThrowsException(() => _ = provider.PairToFkDefinition); + Assert.ThrowsException(() => _ = provider.RelationshipToFkDefinition); + Assert.ThrowsException(() => provider.RelationshipToFkDefinition = new()); + Assert.ThrowsException(() => provider.GetQueryBuilder()); + Assert.ThrowsException(() => provider.VerifyForeignKeyExistsInDB(new(), new())); + Assert.ThrowsException(() => provider.ParseSchemaAndDbTableName("container")); + Assert.ThrowsException(() => provider.GetEntityNamesAndDbObjects()); + Assert.ThrowsException(() => provider.TryGetEntityNameFromPath("path", out _)); + Assert.ThrowsException(() => provider.TryGetExposedFieldToBackingFieldMap("Book", out _)); + Assert.ThrowsException(() => provider.TryGetBackingFieldToExposedFieldMap("Book", out _)); + Assert.ThrowsException(() => provider.InitializeAsync(new(), new())); + Assert.ThrowsException(() => provider.GetStoredProcedureDefinition("Book")); + } + + [TestMethod] + public void TrivialMetadataMembers_ReturnCosmosDefaults() + { + CosmosSqlMetadataProvider provider = CreateProvider(databaseType: DatabaseType.CosmosDB_NoSQL, isDevelopment: true); + + Assert.AreEqual(DatabaseType.CosmosDB_NoSQL, provider.GetDatabaseType()); + Assert.AreEqual(string.Empty, provider.GetDefaultSchemaName()); + Assert.IsTrue(provider.IsDevelopmentMode()); + Assert.AreEqual(0, provider.GetSourceDefinition("Book").Columns.Count); + Assert.AreSame(provider.GetODataParser(), provider.GetODataParser()); + Assert.IsTrue(provider.InitializeAsync().IsCompletedSuccessfully); + } + + [TestMethod] + public void FieldMappingMembers_ReturnInputFieldWithoutMapping() + { + CosmosSqlMetadataProvider provider = CreateProvider(); + + Assert.IsTrue(provider.TryGetExposedColumnName("Book", "id", out string? exposed)); + Assert.AreEqual("id", exposed); + Assert.IsTrue(provider.TryGetBackingColumn("Book", "id", out string? backing)); + Assert.AreEqual("id", backing); + Assert.IsFalse(provider.TryGetArrayElementSyntaxKind("Book", "id", out SyntaxKind kind)); + Assert.AreEqual(default, kind); + } + + [TestMethod] + public void PartitionKeyPath_CanBeAddedUpdatedAndRead() + { + CosmosSqlMetadataProvider provider = CreateProvider(); + + Assert.IsNull(provider.GetPartitionKeyPath("db", "container")); + provider.SetPartitionKeyPath("db", "container", "/tenantId"); + Assert.AreEqual("/tenantId", provider.GetPartitionKeyPath("db", "container")); + provider.SetPartitionKeyPath("db", "container", "/accountId"); + Assert.AreEqual("/accountId", provider.GetPartitionKeyPath("db", "container")); + } + + [DataTestMethod] + [DataRow(null, "container", "/id")] + [DataRow("db", null, "/id")] + [DataRow("db", "container", null)] + public void SetPartitionKeyPath_NullArgumentsThrow(string? database, string? container, string? path) + { + CosmosSqlMetadataProvider provider = CreateProvider(); + Assert.ThrowsException(() => provider.SetPartitionKeyPath(database!, container!, path!)); + } + + [DataTestMethod] + [DataRow(null, "container")] + [DataRow("db", null)] + public void GetPartitionKeyPath_NullArgumentsThrow(string? database, string? container) + { + CosmosSqlMetadataProvider provider = CreateProvider(); + Assert.ThrowsException(() => provider.GetPartitionKeyPath(database!, container!)); + } + + [DataTestMethod] + [DataRow("db.books", "configuredDb", "configuredContainer", "books")] + [DataRow("books", "configuredDb", "configuredContainer", "books")] + [DataRow("", "configuredDb", "configuredContainer", "configuredContainer")] + public void GetDatabaseObjectName_ResolvesSourceOrConfiguredContainer( + string source, + string configuredDatabase, + string configuredContainer, + string expected) + { + CosmosSqlMetadataProvider provider = CreateProvider( + entities: new() { ["Book"] = CreateEntity(source) }, + options: new(configuredDatabase, configuredContainer, null, null)); + + Assert.AreEqual(expected, provider.GetDatabaseObjectName("Book")); + } + + [DataTestMethod] + [DataRow("db.books", "configuredDb", "db")] + [DataRow("books", "configuredDb", "configuredDb")] + [DataRow("", "configuredDb", "configuredDb")] + public void GetSchemaName_ResolvesSourceOrConfiguredDatabase(string source, string configuredDatabase, string expected) + { + CosmosSqlMetadataProvider provider = CreateProvider( + entities: new() { ["Book"] = CreateEntity(source) }, + options: new(configuredDatabase, "container", null, null)); + + Assert.AreEqual(expected, provider.GetSchemaName("Book")); + } + + [TestMethod] + public void GetSchemaName_MissingDatabaseThrows() + { + CosmosSqlMetadataProvider provider = CreateProvider( + entities: new() { ["Book"] = CreateEntity(string.Empty) }, + options: new(null, "container", null, null)); + + Assert.ThrowsException(() => provider.GetSchemaName("Book")); + } + + [TestMethod] + public void GetEntityName_ResolvesDirectModelDirectiveAndSingularNames() + { + Dictionary entities = new() + { + ["Book"] = CreateEntity("db.books", singular: "Volume") + }; + DocumentNode schema = Utf8GraphQLParser.Parse("type BookAlias @model(name: \"Book\") { id: ID }"); + CosmosSqlMetadataProvider provider = CreateProvider(entities: entities, schema: schema); + + Assert.AreEqual("Book", provider.GetEntityName("Book")); + Assert.AreEqual("Book", provider.GetEntityName("BookAlias")); + Assert.AreEqual("Volume", provider.GetEntityName("Volume")); + Assert.ThrowsException(() => provider.GetEntityName("Missing")); + } + + [DataTestMethod] + [DataRow("")] + [DataRow("not valid graphql")] + public void ParseSchemaGraphQLDocument_InvalidSchemaThrows(string schema) + { + CosmosSqlMetadataProvider provider = CreateProvider(options: new("db", "container", null, schema)); + + Assert.ThrowsException(() => provider.ParseSchemaGraphQLDocument()); + } + + [TestMethod] + public void ParseSchemaGraphQLDocument_LoadsSchemaFromConfiguredFile() + { + Mock fileSystem = new(); + fileSystem.Setup(x => x.File.ReadAllText("schema.graphql")) + .Returns("type Book @model(name: \"Book\") { id: ID!, title: String }"); + CosmosSqlMetadataProvider provider = CreateProvider(options: new("db", "container", "schema.graphql", null)); + SetField(provider, "_fileSystem", fileSystem.Object); + + provider.ParseSchemaGraphQLDocument(); + InvokePrivate(provider, "ParseSchemaGraphQLFieldsForGraphQLType"); + + CollectionAssert.AreEquivalent(new[] { "id", "title" }, provider.GetSchemaGraphQLFieldNamesForEntityName("Book")); + Assert.AreEqual("ID!", provider.GetSchemaGraphQLFieldTypeFromFieldName("Book", "id")); + Assert.AreEqual("title", provider.GetSchemaGraphQLFieldFromFieldName("Book", "title")!.Name.Value); + Assert.AreEqual(0, provider.GetSchemaGraphQLFieldNamesForEntityName("Missing").Count); + Assert.IsNull(provider.GetSchemaGraphQLFieldTypeFromFieldName("Missing", "id")); + Assert.IsNull(provider.GetSchemaGraphQLFieldFromFieldName("Missing", "id")); + } + + [TestMethod] + public void AssertIfEntityIsAvailableInConfig_MissingEntityThrows() + { + CosmosSqlMetadataProvider provider = CreateProvider(); + + TargetInvocationException exception = Assert.ThrowsException( + () => InvokePrivate(provider, "AssertIfEntityIsAvailableInConfig", "Missing")); + + Assert.IsInstanceOfType(exception.InnerException); + } + + private static CosmosSqlMetadataProvider CreateProvider( + Dictionary? entities = null, + CosmosDbNoSQLDataSourceOptions? options = null, + DatabaseType databaseType = DatabaseType.CosmosDB_NoSQL, + bool isDevelopment = false, + DocumentNode? schema = null) + { + CosmosSqlMetadataProvider provider = (CosmosSqlMetadataProvider)RuntimeHelpers.GetUninitializedObject(typeof(CosmosSqlMetadataProvider)); + SetField(provider, "_runtimeConfigEntities", new RuntimeEntities(entities ?? new Dictionary())); + SetField(provider, "_cosmosDb", options ?? new CosmosDbNoSQLDataSourceOptions("db", "container", null, null)); + SetField(provider, "_databaseType", databaseType); + SetField(provider, "_isDevelopmentMode", isDevelopment); + SetField(provider, "_partitionKeyPaths", new ConcurrentDictionary()); + SetField(provider, "_oDataParser", new ODataParser()); + SetField(provider, "_graphQLTypeToFieldsMap", new Dictionary>()); + provider.GraphQLSchemaRoot = schema ?? new DocumentNode(Array.Empty()); + return provider; + } + + private static Entity CreateEntity(string source, string singular = "Book") => + new( + Source: new EntitySource(source, EntitySourceType.Table, null, null), + GraphQL: new EntityGraphQLOptions(singular, "Books"), + Fields: null, + Rest: new EntityRestOptions(Enabled: true), + Permissions: Array.Empty(), + Mappings: null, + Relationships: null); + + private static void SetField(CosmosSqlMetadataProvider provider, string name, object value) + { + typeof(CosmosSqlMetadataProvider).GetField(name, BindingFlags.Instance | BindingFlags.NonPublic)!.SetValue(provider, value); + } + + private static object? InvokePrivate(CosmosSqlMetadataProvider provider, string name, params object[] arguments) + { + return typeof(CosmosSqlMetadataProvider).GetMethod(name, BindingFlags.Instance | BindingFlags.NonPublic)! + .Invoke(provider, arguments); + } + } +} diff --git a/src/Service.Tests/UnitTests/DatabaseObjectUnitTests.cs b/src/Service.Tests/UnitTests/DatabaseObjectUnitTests.cs new file mode 100644 index 0000000000..b51087f010 --- /dev/null +++ b/src/Service.Tests/UnitTests/DatabaseObjectUnitTests.cs @@ -0,0 +1,167 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +using System; +using System.Collections.Generic; +using System.Data; +using Azure.DataApiBuilder.Config.DatabasePrimitives; +using Azure.DataApiBuilder.Config.ObjectModel; +using Microsoft.VisualStudio.TestTools.UnitTesting; + +namespace Azure.DataApiBuilder.Service.Tests.UnitTests +{ + [TestClass] + public class DatabaseObjectUnitTests + { + [TestMethod] + public void SourceDefinition_ReturnsDefinitionForEverySupportedSourceType() + { + SourceDefinition tableDefinition = new(); + ViewDefinition viewDefinition = new(); + StoredProcedureDefinition storedProcedureDefinition = new(); + + Assert.AreSame(tableDefinition, new DatabaseTable("dbo", "books") + { + SourceType = EntitySourceType.Table, + TableDefinition = tableDefinition + }.SourceDefinition); + Assert.AreSame(viewDefinition, new DatabaseView("dbo", "books_view") + { + SourceType = EntitySourceType.View, + ViewDefinition = viewDefinition + }.SourceDefinition); + Assert.AreSame(storedProcedureDefinition, new DatabaseStoredProcedure("dbo", "get_books") + { + SourceType = EntitySourceType.StoredProcedure, + StoredProcedureDefinition = storedProcedureDefinition + }.SourceDefinition); + } + + [TestMethod] + public void SourceDefinition_UnsupportedSourceType_Throws() + { + DatabaseTable databaseObject = new("dbo", "books") + { + SourceType = (EntitySourceType)int.MaxValue, + TableDefinition = new() + }; + + Exception exception = Assert.ThrowsException(() => _ = databaseObject.SourceDefinition); + StringAssert.Contains(exception.Message, "Unsupported EntitySourceType"); + } + + [TestMethod] + public void StoredProcedureDefinition_UnknownParameter_ReturnsNull() + { + StoredProcedureDefinition definition = new() + { + Parameters = new Dictionary + { + ["id"] = new() { DbType = DbType.Int32 } + } + }; + + Assert.AreEqual(DbType.Int32, definition.GetDbTypeForParam("id")); + Assert.IsNull(definition.GetDbTypeForParam("missing")); + } + + [DataTestMethod] + [DataRow(RelationshipRole.Target, RelationshipRole.None, true)] + [DataRow(RelationshipRole.None, RelationshipRole.Target, false)] + public void ForeignKeyDefinition_ResolveTargetColumns_ReturnsRoleColumns( + RelationshipRole referencingRole, + RelationshipRole referencedRole, + bool expectReferencing) + { + ForeignKeyDefinition definition = CreateForeignKeyDefinition(referencingRole, referencedRole); + + CollectionAssert.AreEqual( + expectReferencing ? definition.ReferencingColumns : definition.ReferencedColumns, + definition.ResolveTargetColumns()); + } + + [TestMethod] + public void ForeignKeyDefinition_ResolveTargetColumns_WithoutTargetRole_Throws() + { + ForeignKeyDefinition definition = CreateForeignKeyDefinition(RelationshipRole.Source, RelationshipRole.Linking); + + StringAssert.Contains( + Assert.ThrowsException(() => definition.ResolveTargetColumns()).Message, + "Unable to resolve target columns"); + } + + [DataTestMethod] + [DataRow(RelationshipRole.Source, RelationshipRole.None, true)] + [DataRow(RelationshipRole.None, RelationshipRole.Source, false)] + public void ForeignKeyDefinition_ResolveSourceColumns_ReturnsRoleColumns( + RelationshipRole referencingRole, + RelationshipRole referencedRole, + bool expectReferencing) + { + ForeignKeyDefinition definition = CreateForeignKeyDefinition(referencingRole, referencedRole); + + CollectionAssert.AreEqual( + expectReferencing ? definition.ReferencingColumns : definition.ReferencedColumns, + definition.ResolveSourceColumns()); + } + + [TestMethod] + public void ForeignKeyDefinition_ResolveSourceColumns_WithoutSourceRole_Throws() + { + ForeignKeyDefinition definition = CreateForeignKeyDefinition(RelationshipRole.Target, RelationshipRole.Linking); + + StringAssert.Contains( + Assert.ThrowsException(() => definition.ResolveSourceColumns()).Message, + "Unable to resolve source columns"); + } + + [TestMethod] + public void ForeignKeyDefinition_Equality_UsesPairAndOrderedColumns() + { + ForeignKeyDefinition first = CreateForeignKeyDefinition(RelationshipRole.Source, RelationshipRole.Target); + ForeignKeyDefinition equal = CreateForeignKeyDefinition(RelationshipRole.Source, RelationshipRole.Target); + ForeignKeyDefinition different = CreateForeignKeyDefinition(RelationshipRole.Source, RelationshipRole.Target); + different.ReferencedColumns = new() { "different" }; + + Assert.IsTrue(first.Equals((object)equal)); + Assert.IsTrue(first.Equals(equal)); + Assert.IsFalse(first.Equals((ForeignKeyDefinition?)null)); + Assert.IsFalse(first.Equals(different)); + _ = first.GetHashCode(); + _ = equal.GetHashCode(); + } + + [TestMethod] + public void RelationshipPair_ConstructorsAndEquality_UseDatabaseObjects() + { + DatabaseTable referencing = new("dbo", "books"); + DatabaseTable referenced = new("dbo", "publishers"); + RelationShipPair unnamed = new(referencing, referenced); + RelationShipPair named = new("book_publisher", referencing, referenced); + RelationShipPair equal = new("other_name", new DatabaseTable("DBO", "BOOKS"), new DatabaseTable("dbo", "Publishers")); + + Assert.AreEqual(string.Empty, unnamed.RelationshipName); + Assert.AreEqual("book_publisher", named.RelationshipName); + Assert.IsTrue(named.Equals((object)equal)); + Assert.IsTrue(named.Equals(equal)); + Assert.IsFalse(named.Equals((RelationShipPair?)null)); + Assert.AreEqual(named.GetHashCode(), equal.GetHashCode()); + } + + private static ForeignKeyDefinition CreateForeignKeyDefinition( + RelationshipRole referencingRole, + RelationshipRole referencedRole) + { + return new ForeignKeyDefinition + { + ReferencingEntityRole = referencingRole, + ReferencedEntityRole = referencedRole, + Pair = new RelationShipPair( + new DatabaseTable("dbo", "books"), + new DatabaseTable("dbo", "publishers")), + ReferencingColumns = new() { "publisher_id" }, + ReferencedColumns = new() { "id" } + }; + } + } +} diff --git a/src/Service.Tests/UnitTests/DeserializationVariableReplacementSettingsTests.cs b/src/Service.Tests/UnitTests/DeserializationVariableReplacementSettingsTests.cs index 6548b5d84a..35bb955cbb 100644 --- a/src/Service.Tests/UnitTests/DeserializationVariableReplacementSettingsTests.cs +++ b/src/Service.Tests/UnitTests/DeserializationVariableReplacementSettingsTests.cs @@ -3,6 +3,7 @@ using System; using System.IO; +using System.Reflection; using System.Text.RegularExpressions; using Azure.DataApiBuilder.Config; using Azure.DataApiBuilder.Config.Converters; @@ -211,6 +212,37 @@ public void Constructor_RemoteAkvEndpointNoRetryPolicy_RegistersStrategy() Assert.AreEqual(1, settings.ReplacementStrategies.Count); } + [DataTestMethod] + [DataRow("", false, "null or empty")] + [DataRow("-leading", false, "start and end")] + [DataRow("trailing-", false, "start and end")] + [DataRow("bad_name", false, "Invalid character")] + [DataRow("valid-name-123", true, "")] + public void IsValidAkvSecretName_ValidatesShape(string name, bool expected, string errorFragment) + { + MethodInfo method = typeof(DeserializationVariableReplacementSettings).GetMethod( + "IsValidAkvSecretName", + BindingFlags.Static | BindingFlags.NonPublic)!; + object?[] arguments = { name, null }; + + bool result = (bool)method.Invoke(null, arguments)!; + + Assert.AreEqual(expected, result); + StringAssert.Contains((string)arguments[1]!, errorFragment); + } + + [TestMethod] + public void IsValidAkvSecretName_RejectsNamesLongerThanLimit() + { + MethodInfo method = typeof(DeserializationVariableReplacementSettings).GetMethod( + "IsValidAkvSecretName", + BindingFlags.Static | BindingFlags.NonPublic)!; + object?[] arguments = { new string('a', 128), null }; + + Assert.IsFalse((bool)method.Invoke(null, arguments)!); + StringAssert.Contains((string)arguments[1]!, "outside allowed range"); + } + private static string WriteAkvFile(params string[] lines) { string path = Path.Combine(Path.GetTempPath(), $"dab-test-{Guid.NewGuid():N}.akv"); diff --git a/src/Service.Tests/UnitTests/DmlToolsConfigConverterTests.cs b/src/Service.Tests/UnitTests/DmlToolsConfigConverterTests.cs index 75020fe951..7887dca716 100644 --- a/src/Service.Tests/UnitTests/DmlToolsConfigConverterTests.cs +++ b/src/Service.Tests/UnitTests/DmlToolsConfigConverterTests.cs @@ -74,6 +74,17 @@ public void Deserialize_BooleanFalse_DisablesAllTools() Assert.IsFalse(config.AggregateRecords); } + [TestMethod] + public void Deserialize_NullBypassesConverter_AndUnexpectedTokenReturnsDefaultConfiguration() + { + DmlToolsConfig fromNull = JsonSerializer.Deserialize("null", GetOptions()); + DmlToolsConfig fromNumber = JsonSerializer.Deserialize("42", GetOptions()); + + Assert.IsNull(fromNull); + Assert.IsNotNull(fromNumber); + Assert.IsTrue(fromNumber.DescribeEntities); + } + [TestMethod] public void Deserialize_ObjectWithIndividualSettings_AppliesOverrides() { @@ -239,5 +250,6 @@ public void Serialize_DefaultConfig_WritesNothing() Assert.IsFalse(jObject.ContainsKey("dml-tools"), "Default config should not be written."); } + } } diff --git a/src/Service.Tests/UnitTests/DwSqlQueryBuilderHelperTests.cs b/src/Service.Tests/UnitTests/DwSqlQueryBuilderHelperTests.cs new file mode 100644 index 0000000000..2152cab39a --- /dev/null +++ b/src/Service.Tests/UnitTests/DwSqlQueryBuilderHelperTests.cs @@ -0,0 +1,172 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +using System.Collections.Generic; +using System.Reflection; +using System.Runtime.CompilerServices; +using Azure.DataApiBuilder.Config.DatabasePrimitives; +using Azure.DataApiBuilder.Config.ObjectModel; +using Azure.DataApiBuilder.Core.Models; +using Azure.DataApiBuilder.Core.Resolvers; +using Azure.DataApiBuilder.Core.Services; +using Microsoft.VisualStudio.TestTools.UnitTesting; +using Moq; + +namespace Azure.DataApiBuilder.Service.Tests.UnitTests +{ + [TestClass] + public class DwSqlQueryBuilderHelperTests + { + [TestMethod] + public void HasToOneOrNoRelation_NullStructureReturnsTrue() + { + Assert.IsTrue(InvokeHasToOneOrNoRelation(null, false)); + } + + [TestMethod] + public void HasToOneOrNoRelation_EmptyAndNestedPointQueriesReturnTrue() + { + SqlQueryStructure child = CreateStructure(isList: false); + SqlQueryStructure parent = CreateStructure(isList: true, ("child", child)); + + Assert.IsTrue(InvokeHasToOneOrNoRelation(parent, false)); + } + + [TestMethod] + public void HasToOneOrNoRelation_NestedListQueryReturnsFalse() + { + SqlQueryStructure child = CreateStructure(isList: true); + SqlQueryStructure parent = CreateStructure(isList: true, ("children", child)); + + Assert.IsFalse(InvokeHasToOneOrNoRelation(parent, false)); + Assert.IsFalse(InvokeHasToOneOrNoRelation(child, true)); + } + + [TestMethod] + public void GenerateColumnsAsJsonObject_HandlesSingleAndMultipleColumns() + { + SqlQueryStructure single = CreateStructure(false); + SetBaseProperty(single, "Columns", new List + { + new("dbo", "books", "id", "id", "table0") + }); + SqlQueryStructure multiple = CreateStructure(false); + SetBaseProperty(multiple, "Columns", new List + { + new("dbo", "books", "id", "id", "table0"), + new("dbo", "books", "title", "bookTitle", "table0") + }); + + Assert.AreEqual("JSON_OBJECT('id': [id])", InvokeStatic("GenerateColumnsAsJsonObject", single)); + Assert.AreEqual("JSON_OBJECT('id': [id],'bookTitle': [bookTitle])", InvokeStatic("GenerateColumnsAsJsonObject", multiple)); + } + + [TestMethod] + public void BuildProcedureParameterList_FormatsValuesAndHandlesEmptyInput() + { + Assert.AreEqual(string.Empty, InvokeStatic("BuildProcedureParameterList", new Dictionary())); + Assert.AreEqual("@id = @param0, @name = @param1", InvokeStatic( + "BuildProcedureParameterList", + new Dictionary { ["id"] = "@param0", ["name"] = "@param1" })); + } + + [TestMethod] + public void Build_OptimizedSimpleQueryUsesJsonFunctions() + { + SqlQueryStructure structure = CreateBuildableStructure(isList: true); + + string query = new DwSqlQueryBuilder(enableNto1JoinOpt: true).Build(structure); + + StringAssert.Contains(query, "SELECT TOP 100 [table0].[id] AS [id]"); + StringAssert.Contains(query, "FROM [dbo].[books] AS [table0]"); + StringAssert.Contains(query, "FOR JSON PATH, INCLUDE_NULL_VALUES"); + Assert.IsFalse(query.Contains("STRING_AGG")); + } + + [TestMethod] + public void Build_UnoptimizedSimpleQueryUsesStringAggregation() + { + SqlQueryStructure structure = CreateBuildableStructure(isList: true); + + string query = new DwSqlQueryBuilder(enableNto1JoinOpt: false).Build(structure); + + StringAssert.Contains(query, "STRING_AGG"); + StringAssert.Contains(query, "FROM [dbo].[books] AS [table0]"); + Assert.IsFalse(query.Contains("FOR JSON PATH")); + } + + private static bool InvokeHasToOneOrNoRelation(SqlQueryStructure? structure, bool isSubQuery) => + InvokeStatic("HasToOneOrNoRelation", structure, isSubQuery); + + private static T InvokeStatic(string methodName, params object?[] arguments) + { + MethodInfo method = typeof(DwSqlQueryBuilder).GetMethod(methodName, BindingFlags.Static | BindingFlags.NonPublic)!; + return (T)method.Invoke(null, arguments)!; + } + + private static SqlQueryStructure CreateStructure(bool isList, params (string Alias, SqlQueryStructure Query)[] joins) + { + SqlQueryStructure structure = (SqlQueryStructure)RuntimeHelpers.GetUninitializedObject(typeof(SqlQueryStructure)); + structure.IsListQuery = isList; + Dictionary joinQueries = new(); + foreach ((string alias, SqlQueryStructure query) in joins) + { + joinQueries.Add(alias, query); + } + + typeof(SqlQueryStructure).GetField("k__BackingField", BindingFlags.Instance | BindingFlags.NonPublic)! + .SetValue(structure, joinQueries); + return structure; + } + + private static SqlQueryStructure CreateBuildableStructure(bool isList) + { + SqlQueryStructure structure = CreateStructure(isList); + SourceDefinition sourceDefinition = new(); + sourceDefinition.Columns.Add("id", new ColumnDefinition { SystemType = typeof(int) }); + DatabaseTable databaseTable = new("dbo", "books") { TableDefinition = sourceDefinition }; + Mock metadataProvider = new(); + metadataProvider.Setup(x => x.GetSourceDefinition("Book")).Returns(sourceDefinition); + + SetField(structure, "EntityName", "Book"); + SetField(structure, "MetadataProvider", metadataProvider.Object); + SetField(structure, "DatabaseObject", databaseTable); + SetField(structure, "SourceAlias", "table0"); + SetField(structure, "Columns", new List + { + new("dbo", "books", "id", "id", "table0") + }); + SetField(structure, "Predicates", new List()); + SetField(structure, "DbPolicyPredicatesForOperations", new Dictionary()); + SetField(structure, "Joins", new List()); + SetField(structure, "FilterPredicates", string.Empty); + SetField(structure, "OrderByColumns", new List()); + SetField(structure, "PaginationMetadata", new PaginationMetadata(structure)); + SetField(structure, "GroupByMetadata", new GroupByMetadata()); + typeof(SqlQueryStructure).GetField("_limit", BindingFlags.Instance | BindingFlags.NonPublic)! + .SetValue(structure, (uint?)100); + return structure; + } + + private static void SetField(SqlQueryStructure structure, string name, T value) + { + for (System.Type? type = typeof(SqlQueryStructure); type is not null; type = type.BaseType) + { + FieldInfo? field = type.GetField($"<{name}>k__BackingField", BindingFlags.Instance | BindingFlags.NonPublic); + if (field is not null) + { + field.SetValue(structure, value); + return; + } + } + + Assert.Fail($"Could not find backing field for {name}."); + } + + private static void SetBaseProperty(SqlQueryStructure structure, string name, T value) + { + typeof(BaseQueryStructure).GetField($"<{name}>k__BackingField", BindingFlags.Instance | BindingFlags.NonPublic)! + .SetValue(structure, value); + } + } +} diff --git a/src/Service.Tests/UnitTests/EmbeddingTelemetryHelperTests.cs b/src/Service.Tests/UnitTests/EmbeddingTelemetryHelperTests.cs new file mode 100644 index 0000000000..64ce153d3a --- /dev/null +++ b/src/Service.Tests/UnitTests/EmbeddingTelemetryHelperTests.cs @@ -0,0 +1,119 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +using System; +using System.Collections.Generic; +using System.Diagnostics; +using System.Diagnostics.Metrics; +using Azure.DataApiBuilder.Core.Services.Embeddings; +using Azure.DataApiBuilder.Core.Telemetry; +using Microsoft.VisualStudio.TestTools.UnitTesting; + +namespace Azure.DataApiBuilder.Service.Tests.UnitTests +{ + [TestClass] + public class EmbeddingTelemetryHelperTests + { + [TestMethod] + public void MetricHelpers_RecordAllEmbeddingMeasurements() + { + List measurements = new(); + using MeterListener listener = new(); + listener.InstrumentPublished = (instrument, meterListener) => + { + if (instrument.Meter.Name == EmbeddingTelemetryHelper.MeterName) + { + meterListener.EnableMeasurementEvents(instrument); + } + }; + listener.SetMeasurementEventCallback((instrument, _, _, _) => measurements.Add(instrument.Name)); + listener.SetMeasurementEventCallback((instrument, _, _, _) => measurements.Add(instrument.Name)); + listener.SetMeasurementEventCallback((instrument, _, _, _) => measurements.Add(instrument.Name)); + listener.Start(); + + EmbeddingTelemetryHelper.TrackEmbeddingRequest("provider", 2); + EmbeddingTelemetryHelper.TrackApiCall("provider", 2); + EmbeddingTelemetryHelper.TrackCacheHit("provider"); + EmbeddingTelemetryHelper.TrackCacheMiss("provider"); + EmbeddingTelemetryHelper.TrackError("provider", "failure"); + EmbeddingTelemetryHelper.TrackApiDuration("provider", TimeSpan.FromMilliseconds(12), 2); + EmbeddingTelemetryHelper.TrackTotalDuration("provider", TimeSpan.FromMilliseconds(20), fromCache: true); + EmbeddingTelemetryHelper.TrackTokenUsage("provider", 30); + EmbeddingTelemetryHelper.TrackDimensions("provider", 1536); + + CollectionAssert.Contains(measurements, "embedding_requests_total"); + CollectionAssert.Contains(measurements, "embedding_api_calls_total"); + CollectionAssert.Contains(measurements, "embedding_cache_hits_total"); + CollectionAssert.Contains(measurements, "embedding_cache_misses_total"); + CollectionAssert.Contains(measurements, "embedding_errors_total"); + CollectionAssert.Contains(measurements, "embedding_texts_processed_total"); + CollectionAssert.Contains(measurements, "embedding_api_duration_ms"); + CollectionAssert.Contains(measurements, "embedding_total_duration_ms"); + CollectionAssert.Contains(measurements, "embedding_tokens_total"); + CollectionAssert.Contains(measurements, "embedding_dimensions"); + } + + [TestMethod] + public void ActivityHelpers_SetSuccessAndCacheTags() + { + using ActivityListener listener = CreateListener(); + using Activity? activity = EmbeddingTelemetryHelper.StartEmbeddingActivity("EmbedAsync"); + Assert.IsNotNull(activity); + + activity.SetEmbeddingActivityTags("azure-openai", "model", 3); + activity.SetCacheActivityTags(2, 1); + activity.SetEmbeddingActivitySuccess(12.5, 1536); + + Assert.AreEqual("azure-openai", activity.GetTagItem("embedding.provider")); + Assert.AreEqual("model", activity.GetTagItem("embedding.model")); + Assert.AreEqual(3, activity.GetTagItem("embedding.text_count")); + Assert.AreEqual(2, activity.GetTagItem("embedding.cache_hits")); + Assert.AreEqual(1, activity.GetTagItem("embedding.cache_misses")); + Assert.AreEqual(12.5, activity.GetTagItem("embedding.duration_ms")); + Assert.AreEqual(1536, activity.GetTagItem("embedding.dimensions")); + Assert.AreEqual(ActivityStatusCode.Ok, activity.Status); + } + + [TestMethod] + public void ActivityHelpers_OmitOptionalModelAndDimensions() + { + using ActivityListener listener = CreateListener(); + using Activity? activity = EmbeddingTelemetryHelper.StartEmbeddingActivity("EmbedBatchAsync"); + Assert.IsNotNull(activity); + + activity.SetEmbeddingActivityTags("openai", null, 1); + activity.SetEmbeddingActivitySuccess(1.5); + + Assert.IsNull(activity.GetTagItem("embedding.model")); + Assert.IsNull(activity.GetTagItem("embedding.dimensions")); + } + + [TestMethod] + public void SetEmbeddingActivityError_RecordsExceptionDetails() + { + using ActivityListener listener = CreateListener(); + using Activity? activity = EmbeddingTelemetryHelper.StartEmbeddingActivity("EmbedAsync"); + Assert.IsNotNull(activity); + InvalidOperationException error = new("boom"); + + activity.SetEmbeddingActivityError(error); + + Assert.AreEqual(ActivityStatusCode.Error, activity.Status); + Assert.AreEqual("boom", activity.StatusDescription); + Assert.AreEqual(nameof(InvalidOperationException), activity.GetTagItem("error.type")); + Assert.AreEqual("boom", activity.GetTagItem("error.message")); + } + + private static ActivityListener CreateListener() + { + string sourceName = TelemetryTracesHelper.DABActivitySource.Name; + ActivityListener listener = new() + { + ShouldListenTo = source => source.Name == sourceName, + Sample = (ref ActivityCreationOptions _) => ActivitySamplingResult.AllDataAndRecorded + }; + ActivitySource.AddActivityListener(listener); + return listener; + } + } +} diff --git a/src/Service.Tests/UnitTests/EmbeddingsOptionsConverterTests.cs b/src/Service.Tests/UnitTests/EmbeddingsOptionsConverterTests.cs index 4e99a7dfb1..cccd44dcc4 100644 --- a/src/Service.Tests/UnitTests/EmbeddingsOptionsConverterTests.cs +++ b/src/Service.Tests/UnitTests/EmbeddingsOptionsConverterTests.cs @@ -205,5 +205,56 @@ public void RoundTrip_PreservesScalarAndNestedValues() Assert.AreEqual(original.Endpoint.Enabled, result.Endpoint.Enabled); CollectionAssert.AreEqual(original.Endpoint.Roles, result.Endpoint.Roles); } + + [DataTestMethod] + [DataRow(@"{ ""provider"":""openai"", ""base-url"":""u"", ""api-key"":""k"", ""endpoint"": true }")] + [DataRow(@"{ ""provider"":""openai"", ""base-url"":""u"", ""api-key"":""k"", ""health"": [] }")] + [DataRow(@"{ ""provider"":""openai"", ""base-url"":""u"", ""api-key"":""k"", ""chunking"": ""bad"" }")] + public void Deserialize_InvalidNestedObjectType_Throws(string json) + { + Assert.ThrowsException(() => + JsonSerializer.Deserialize(json, GetOptions())); + } + + [TestMethod] + public void Serialize_AllOptionalObjects_WritesEachSection() + { + string json = @"{ + ""provider"": ""openai"", + ""base-url"": ""https://api.openai.com"", + ""api-key"": ""key"", + ""model"": ""model"", + ""api-version"": ""v1"", + ""dimensions"": 256, + ""timeout-ms"": 1000, + ""endpoint"": { ""enabled"": true, ""roles"": [""reader""] }, + ""health"": { ""enabled"": true, ""threshold-ms"": 50 }, + ""chunking"": { ""enabled"": true, ""size-chars"": 100, ""overlap-chars"": 5 }, + ""cache"": { ""enabled"": true, ""ttl-hours"": 1 } + }"; + EmbeddingsOptions options = JsonSerializer.Deserialize(json, GetOptions()); + + JObject serialized = JObject.Parse(JsonSerializer.Serialize(options, GetOptions())); + + Assert.IsNotNull(serialized["endpoint"]); + Assert.IsNotNull(serialized["health"]); + Assert.IsNotNull(serialized["chunking"]); + Assert.IsNotNull(serialized["cache"]); + Assert.AreEqual("model", serialized["model"]!.Value()); + Assert.AreEqual("v1", serialized["api-version"]!.Value()); + Assert.AreEqual(256, serialized["dimensions"]!.Value()); + Assert.AreEqual(1000, serialized["timeout-ms"]!.Value()); + } + + [TestMethod] + public void Serialize_UnknownProvider_Throws() + { + EmbeddingsOptions options = new( + Provider: (EmbeddingProviderType)999, + BaseUrl: "https://example.test", + ApiKey: "key"); + + Assert.ThrowsException(() => JsonSerializer.Serialize(options, GetOptions())); + } } } diff --git a/src/Service.Tests/UnitTests/EntityApiOptionsConverterCoverageTests.cs b/src/Service.Tests/UnitTests/EntityApiOptionsConverterCoverageTests.cs new file mode 100644 index 0000000000..73ddea7d06 --- /dev/null +++ b/src/Service.Tests/UnitTests/EntityApiOptionsConverterCoverageTests.cs @@ -0,0 +1,137 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +using System; +using System.Text.Json; +using System.Text.Json.Serialization; +using Azure.DataApiBuilder.Config; +using Azure.DataApiBuilder.Config.ObjectModel; +using Microsoft.VisualStudio.TestTools.UnitTesting; + +namespace Azure.DataApiBuilder.Service.Tests.UnitTests +{ + [TestClass] + public class EntityApiOptionsConverterCoverageTests + { + private static JsonSerializerOptions Options => RuntimeConfigLoader.GetSerializationOptions(); + + [TestMethod] + public void RestOptions_ObjectReadsAllPropertiesAndWritesThem() + { + const string Json = """ + { "path": "/books", "methods": ["get", "post"], "enabled": false } + """; + + EntityRestOptions? options = JsonSerializer.Deserialize(Json, Options); + string serialized = JsonSerializer.Serialize(options, Options); + + Assert.IsNotNull(options); + Assert.AreEqual("/books", options.Path); + CollectionAssert.AreEqual(new[] { SupportedHttpVerb.Get, SupportedHttpVerb.Post }, options.Methods); + Assert.IsFalse(options.Enabled); + StringAssert.Contains(serialized, "\"methods\""); + } + + [DataTestMethod] + [DataRow("\"/books\"", "/books", true)] + [DataRow("true", null, true)] + [DataRow("false", null, false)] + public void RestOptions_ShorthandFormsDeserialize(string json, string? expectedPath, bool expectedEnabled) + { + EntityRestOptions? options = JsonSerializer.Deserialize(json, Options); + + Assert.IsNotNull(options); + Assert.AreEqual(expectedPath, options.Path); + Assert.AreEqual(expectedEnabled, options.Enabled); + } + + [DataTestMethod] + [DataRow("{\"path\":42}")] + [DataRow("{\"unexpected\":true}")] + [DataRow("42")] + public void RestOptions_InvalidFormsThrow(string json) + { + Assert.ThrowsException(() => JsonSerializer.Deserialize(json, Options)); + } + + [TestMethod] + public void RestOptions_WriteNullPath_WhenNullsAreNotIgnored() + { + JsonSerializerOptions options = new(Options) { DefaultIgnoreCondition = JsonIgnoreCondition.Never }; + + string json = JsonSerializer.Serialize(new EntityRestOptions(Array.Empty(), null, true), options); + using JsonDocument document = JsonDocument.Parse(json); + + Assert.AreEqual(JsonValueKind.Null, document.RootElement.GetProperty("path").ValueKind); + } + + [TestMethod] + public void GraphQLOptions_ObjectReadsNestedTypeAndOperationAndWritesThem() + { + const string Json = """ + { + "enabled": true, + "type": { "singular": "book", "ignored": "value", "plural": "books" }, + "operation": "mutation" + } + """; + + EntityGraphQLOptions? options = JsonSerializer.Deserialize(Json, Options); + string serialized = JsonSerializer.Serialize(options, Options); + + Assert.IsNotNull(options); + Assert.AreEqual("book", options.Singular); + Assert.AreEqual("books", options.Plural); + Assert.AreEqual(GraphQLOperation.Mutation, options.Operation); + StringAssert.Contains(serialized, "\"operation\""); + } + + [DataTestMethod] + [DataRow("true", "", true)] + [DataRow("false", "", false)] + [DataRow("\"book\"", "book", true)] + public void GraphQLOptions_ShorthandFormsDeserialize(string json, string expectedSingular, bool expectedEnabled) + { + EntityGraphQLOptions? options = JsonSerializer.Deserialize(json, Options); + + Assert.IsNotNull(options); + Assert.AreEqual(expectedSingular, options.Singular); + Assert.AreEqual(expectedEnabled, options.Enabled); + } + + [DataTestMethod] + [DataRow("{\"type\":[]}")] + [DataRow("42")] + public void GraphQLOptions_InvalidFormsThrow(string json) + { + Assert.ThrowsException(() => JsonSerializer.Deserialize(json, Options)); + } + + [TestMethod] + public void GraphQLOptions_WriteNullOperation_WhenNullsAreNotIgnored() + { + JsonSerializerOptions options = new(Options) { DefaultIgnoreCondition = JsonIgnoreCondition.Never }; + + string json = JsonSerializer.Serialize(new EntityGraphQLOptions("book", "books", true), options); + using JsonDocument document = JsonDocument.Parse(json); + + Assert.AreEqual(JsonValueKind.Null, document.RootElement.GetProperty("operation").ValueKind); + } + + [TestMethod] + public void EntityAction_ObjectWithoutExcludeNormalizesToEmptyCollection() + { + EntityAction? action = JsonSerializer.Deserialize( + "{\"action\":\"read\",\"fields\":{\"include\":[\"id\"]}}", Options); + + Assert.IsNotNull(action?.Fields?.Exclude); + Assert.AreEqual(0, action.Fields.Exclude.Count); + } + + [TestMethod] + public void EntityCacheOptions_NonObjectThrows() + { + Assert.ThrowsException(() => JsonSerializer.Deserialize("true", Options)); + } + } +} diff --git a/src/Service.Tests/UnitTests/EntityHealthOptionsConverterTests.cs b/src/Service.Tests/UnitTests/EntityHealthOptionsConverterTests.cs new file mode 100644 index 0000000000..0cc2e48920 --- /dev/null +++ b/src/Service.Tests/UnitTests/EntityHealthOptionsConverterTests.cs @@ -0,0 +1,132 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +using System.Text.Json; +using System.Text.Json.Serialization; +using Azure.DataApiBuilder.Config; +using Azure.DataApiBuilder.Config.Converters; +using Azure.DataApiBuilder.Config.HealthCheck; +using Azure.DataApiBuilder.Config.ObjectModel; +using Microsoft.VisualStudio.TestTools.UnitTesting; + +namespace Azure.DataApiBuilder.Service.Tests.UnitTests +{ + [TestClass] + public class EntityHealthOptionsConverterTests + { + private static JsonSerializerOptions Options => RuntimeConfigLoader.GetSerializationOptions(); + + [TestMethod] + public void Deserialize_Null_UsesDefaults() + { + EntityHealthOptionsConvertorFactory factory = new(); + JsonConverter converter = + (JsonConverter)factory.CreateConverter(typeof(EntityHealthCheckConfig), Options)!; + Utf8JsonReader reader = new("null"u8); + Assert.IsTrue(reader.Read()); + + EntityHealthCheckConfig? result = converter.Read(ref reader, typeof(EntityHealthCheckConfig), Options); + + Assert.IsNotNull(result); + Assert.IsTrue(result.Enabled); + Assert.AreEqual(HealthCheckConstants.DEFAULT_FIRST_VALUE, result.First); + Assert.AreEqual(HealthCheckConstants.DEFAULT_THRESHOLD_RESPONSE_TIME_MS, result.ThresholdMs); + } + + [TestMethod] + public void Deserialize_AllProperties_PreservesValuesAndPresence() + { + const string json = """ + { + "enabled": false, + "first": 12, + "threshold-ms": 345 + } + """; + + EntityHealthCheckConfig? result = JsonSerializer.Deserialize(json, Options); + + Assert.IsNotNull(result); + Assert.IsFalse(result.Enabled); + Assert.AreEqual(12, result.First); + Assert.AreEqual(345, result.ThresholdMs); + Assert.IsTrue(result.UserProvidedEnabled); + Assert.IsTrue(result.UserProvidedFirst); + Assert.IsTrue(result.UserProvidedThresholdMs); + } + + [TestMethod] + public void Deserialize_NullProperties_UsesDefaultsWithoutPresenceFlags() + { + const string json = """ + { + "enabled": null, + "first": null, + "threshold-ms": null + } + """; + + EntityHealthCheckConfig? result = JsonSerializer.Deserialize(json, Options); + + Assert.IsNotNull(result); + Assert.IsTrue(result.Enabled); + Assert.AreEqual(HealthCheckConstants.DEFAULT_FIRST_VALUE, result.First); + Assert.AreEqual(HealthCheckConstants.DEFAULT_THRESHOLD_RESPONSE_TIME_MS, result.ThresholdMs); + Assert.IsFalse(result.UserProvidedEnabled); + Assert.IsFalse(result.UserProvidedFirst); + Assert.IsFalse(result.UserProvidedThresholdMs); + } + + [DataTestMethod] + [DataRow("{\"first\":0}", "first")] + [DataRow("{\"first\":-1}", "first")] + [DataRow("{\"threshold-ms\":0}", "ttl-seconds")] + [DataRow("{\"threshold-ms\":-1}", "ttl-seconds")] + [DataRow("{\"unexpected\":1}", "Unexpected property")] + public void Deserialize_InvalidValue_ThrowsJsonException(string json, string expectedMessage) + { + JsonException exception = Assert.ThrowsException( + () => JsonSerializer.Deserialize(json, Options)); + + StringAssert.Contains(exception.Message, expectedMessage); + } + + [DataTestMethod] + [DataRow("\"health\"")] + [DataRow("42")] + [DataRow("true")] + [DataRow("[]")] + [DataRow("{")] + public void Deserialize_InvalidShape_ThrowsJsonException(string json) + { + Assert.ThrowsException( + () => JsonSerializer.Deserialize(json, Options)); + } + + [TestMethod] + public void Serialize_UserProvidedValues_WritesAllProperties() + { + EntityHealthCheckConfig value = new(enabled: false, first: 7, thresholdMs: 250); + + string json = JsonSerializer.Serialize(value, Options); + using JsonDocument document = JsonDocument.Parse(json); + + Assert.IsFalse(document.RootElement.GetProperty("enabled").GetBoolean()); + Assert.AreEqual(7, document.RootElement.GetProperty("first").GetInt32()); + Assert.AreEqual(250, document.RootElement.GetProperty("threshold-ms").GetInt32()); + } + + [TestMethod] + public void Serialize_OnlyEnabledProvided_OmitsDefaultOptionalValues() + { + EntityHealthCheckConfig value = new(enabled: true); + + string json = JsonSerializer.Serialize(value, Options); + using JsonDocument document = JsonDocument.Parse(json); + + Assert.IsTrue(document.RootElement.GetProperty("enabled").GetBoolean()); + Assert.IsFalse(document.RootElement.TryGetProperty("first", out _)); + Assert.IsFalse(document.RootElement.TryGetProperty("threshold-ms", out _)); + } + } +} diff --git a/src/Service.Tests/UnitTests/EntitySourceConverterTests.cs b/src/Service.Tests/UnitTests/EntitySourceConverterTests.cs new file mode 100644 index 0000000000..bc78d78ce3 --- /dev/null +++ b/src/Service.Tests/UnitTests/EntitySourceConverterTests.cs @@ -0,0 +1,108 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +using System.Text.Json; +using Azure.DataApiBuilder.Config; +using Azure.DataApiBuilder.Config.ObjectModel; +using Microsoft.VisualStudio.TestTools.UnitTesting; + +namespace Azure.DataApiBuilder.Service.Tests.UnitTests +{ + [TestClass] + public class EntitySourceConverterTests + { + private static JsonSerializerOptions Options => RuntimeConfigLoader.GetSerializationOptions(); + + [TestMethod] + public void Deserialize_StringSource_CreatesTableSource() + { + EntitySource? source = JsonSerializer.Deserialize("\"dbo.books\"", Options); + + Assert.IsNotNull(source); + Assert.AreEqual("dbo.books", source.Object); + Assert.AreEqual(EntitySourceType.Table, source.Type); + Assert.AreEqual(0, source.Parameters?.Count); + Assert.AreEqual(0, source.KeyFields?.Length); + } + + [TestMethod] + public void Deserialize_LegacyParameters_ConvertsClrValuesToStrings() + { + const string json = """ + { + "object": "dbo.run_report", + "type": "stored-procedure", + "parameters": { + "text": "value", + "integer": 42, + "decimal": 1.25, + "huge": 1e400, + "truth": true, + "falsehood": false, + "nothing": null, + "complex": { "x": 1 } + } + } + """; + + EntitySource? source = JsonSerializer.Deserialize(json, Options); + + Assert.IsNotNull(source); + Assert.IsNotNull(source.Parameters); + Assert.AreEqual(8, source.Parameters.Count); + Assert.AreEqual("value", FindDefault(source, "text")); + Assert.AreEqual("42", FindDefault(source, "integer")); + Assert.AreEqual("1.25", FindDefault(source, "decimal")); + Assert.AreEqual(double.PositiveInfinity.ToString(), FindDefault(source, "huge")); + Assert.AreEqual("True", FindDefault(source, "truth")); + Assert.AreEqual("False", FindDefault(source, "falsehood")); + Assert.AreEqual(string.Empty, FindDefault(source, "nothing")); + Assert.AreEqual("{ \"x\": 1 }", FindDefault(source, "complex")); + } + + [TestMethod] + public void Deserialize_ModernParameters_PreservesList() + { + const string json = """ + { + "object": "dbo.run_report", + "type": "stored-procedure", + "parameters": [ + { "name": "limit", "default": "10", "required": false } + ] + } + """; + + EntitySource? source = JsonSerializer.Deserialize(json, Options); + + Assert.IsNotNull(source); + Assert.IsNotNull(source.Parameters); + Assert.AreEqual(1, source.Parameters.Count); + Assert.AreEqual("limit", source.Parameters[0].Name); + Assert.AreEqual("10", source.Parameters[0].Default); + } + + [TestMethod] + public void Serialize_RoundTripsObjectSource() + { + EntitySource source = new( + "dbo.books", + EntitySourceType.Table, + Parameters: new(), + KeyFields: new[] { "id" }); + + string json = JsonSerializer.Serialize(source, Options); + EntitySource? roundTripped = JsonSerializer.Deserialize(json, Options); + + Assert.IsNotNull(roundTripped); + Assert.AreEqual(source.Object, roundTripped.Object); + Assert.AreEqual(source.Type, roundTripped.Type); + CollectionAssert.AreEqual(source.KeyFields, roundTripped.KeyFields); + } + + private static string? FindDefault(EntitySource source, string name) + { + return source.Parameters?.Find(parameter => parameter.Name == name)?.Default; + } + } +} diff --git a/src/Service.Tests/UnitTests/ExecutionHelperScalarTests.cs b/src/Service.Tests/UnitTests/ExecutionHelperScalarTests.cs new file mode 100644 index 0000000000..499892178f --- /dev/null +++ b/src/Service.Tests/UnitTests/ExecutionHelperScalarTests.cs @@ -0,0 +1,80 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +using System; +using System.Reflection; +using System.Text.Json; +using Azure.DataApiBuilder.Service.Exceptions; +using Azure.DataApiBuilder.Service.GraphQLBuilder.CustomScalars; +using Azure.DataApiBuilder.Service.Services; +using HotChocolate.Types; +using Microsoft.VisualStudio.TestTools.UnitTesting; +using NodaTime; + +namespace Azure.DataApiBuilder.Service.Tests.UnitTests +{ + [TestClass] + public class ExecutionHelperScalarTests + { + [TestMethod] + public void CoerceJsonLeafValueToRuntimeType_ConvertsNumericTypes() + { + Assert.AreEqual((byte)7, Coerce("7", new UnsignedByteType())); + Assert.AreEqual((short)-12, Coerce("-12", new ShortType())); + Assert.AreEqual(42, Coerce("42", new IntType())); + Assert.AreEqual(9007199254740991L, Coerce("9007199254740991", new LongType())); + Assert.AreEqual(1.25d, Coerce("1.25", new FloatType())); + Assert.AreEqual(1.25f, Coerce("1.25", new SingleType())); + Assert.AreEqual(1.25m, Coerce("1.25", new DecimalType())); + } + + [TestMethod] + public void CoerceJsonLeafValueToRuntimeType_ConvertsTextualTypes() + { + Assert.AreEqual("text", Coerce("\"text\"", new StringType())); + Assert.AreEqual(new Uri("https://example.test/path"), Coerce("\"https://example.test/path\"", new UrlType())); + Guid id = Guid.NewGuid(); + Assert.AreEqual(id, Coerce($"\"{id}\"", new UuidType())); + Assert.AreEqual(TimeSpan.FromMinutes(90), Coerce("\"PT1H30M\"", new DurationType())); + CollectionAssert.AreEqual(new byte[] { 1, 2, 3 }, (byte[])Coerce("\"AQID\"", new Base64StringType())!); + Assert.AreEqual("text", Coerce("\"text\"", new AnyType())); + } + + [TestMethod] + public void CoerceJsonLeafValueToRuntimeType_ConvertsTemporalAndBooleanTypes() + { + Assert.AreEqual(true, Coerce("true", new BooleanType())); + Assert.AreEqual(false, Coerce("false", new BooleanType())); + Assert.AreEqual(DateTimeOffset.Parse("2026-08-28T12:30:00Z"), Coerce("\"2026-08-28T12:30:00Z\"", new DateTimeType())); + Assert.AreEqual(DateTimeOffset.Parse("2026-08-28"), Coerce("\"2026-08-28\"", new DateType())); + Assert.AreEqual(new LocalTime(12, 30, 15), Coerce("\"12:30:15\"", new HotChocolate.Types.NodaTime.LocalTimeType())); + Assert.IsNull(Coerce("\"null\"", new HotChocolate.Types.NodaTime.LocalTimeType())); + } + + [TestMethod] + public void CoerceJsonLeafValueToRuntimeType_NullAndInvalidTemporalValuesReturnNull() + { + Assert.IsNull(Coerce("null", new StringType())); + Assert.IsNull(Coerce("\"not-a-date\"", new DateTimeType())); + Assert.IsNull(Coerce("\"not-a-date\"", new DateType())); + } + + [TestMethod] + public void CoerceJsonLeafValueToRuntimeType_InvalidRepresentationThrowsMappedException() + { + TargetInvocationException exception = Assert.ThrowsException(() => + Coerce("\"not-an-integer\"", new IntType())); + + Assert.IsInstanceOfType(exception.InnerException); + } + + private static object? Coerce(string json, ITypeDefinition type) + { + using JsonDocument document = JsonDocument.Parse(json); + MethodInfo method = typeof(ExecutionHelper).GetMethod( + "CoerceJsonLeafValueToRuntimeType", + BindingFlags.Static | BindingFlags.NonPublic)!; + return method.Invoke(null, new object[] { document.RootElement, type, "field" }); + } + } +} diff --git a/src/Service.Tests/UnitTests/GraphQLFilterParserUnitTests.cs b/src/Service.Tests/UnitTests/GraphQLFilterParserUnitTests.cs index b2a1e56672..8a051af4eb 100644 --- a/src/Service.Tests/UnitTests/GraphQLFilterParserUnitTests.cs +++ b/src/Service.Tests/UnitTests/GraphQLFilterParserUnitTests.cs @@ -2,11 +2,15 @@ // Licensed under the MIT License. using System.Collections.Generic; +using System.Reflection; using Azure.DataApiBuilder.Config.ObjectModel; using Azure.DataApiBuilder.Core.Configurations; using Azure.DataApiBuilder.Core.Models; +using Azure.DataApiBuilder.Core.Resolvers; using Azure.DataApiBuilder.Core.Services.MetadataProviders; using Azure.DataApiBuilder.Service.Exceptions; +using HotChocolate.Language; +using Microsoft.AspNetCore.Http; using Microsoft.VisualStudio.TestTools.UnitTesting; using Moq; @@ -62,6 +66,88 @@ public void GetMaxNestedFilterDepth_ResolvesEffectiveLimit(int? depthLimit, int Assert.AreEqual(expected, parser.GetMaxNestedFilterDepth()); } + [TestMethod] + public void GetHttpContextFromMiddlewareContext_ReturnsStoredContext() + { + GQLFilterParser parser = CreateParserWithDepthLimit(null); + DefaultHttpContext httpContext = new(); + Mock middleware = new(); + middleware.SetupGet(x => x.ContextData).Returns(new Dictionary + { + [nameof(HttpContext)] = httpContext + }); + + Assert.AreSame(httpContext, parser.GetHttpContextFromMiddlewareContext(middleware.Object)); + } + + [TestMethod] + public void GetHttpContextFromMiddlewareContext_MissingContextThrows() + { + GQLFilterParser parser = CreateParserWithDepthLimit(null); + Mock middleware = new(); + middleware.SetupGet(x => x.ContextData).Returns(new Dictionary()); + + Assert.ThrowsException(() => + parser.GetHttpContextFromMiddlewareContext(middleware.Object)); + } + + [TestMethod] + public void MakeChainPredicate_EmptyOperandsReturnsFalsePredicate() + { + Predicate predicate = GQLFilterParser.MakeChainPredicate(new(), PredicateOperation.AND); + + Assert.IsNotNull(predicate); + } + + [TestMethod] + public void MakeChainPredicate_MultipleOperandsBuildsRecursiveChain() + { + Predicate first = Predicate.MakeFalsePredicate(); + Predicate second = Predicate.MakeFalsePredicate(); + List operands = new() { new(first), new(second) }; + + Predicate result = GQLFilterParser.MakeChainPredicate(operands, PredicateOperation.OR); + + Assert.AreEqual(PredicateOperation.OR, result.Op); + } + + [TestMethod] + public void PreprocessInOperatorValues_RejectsNonListValue() + { + TargetInvocationException exception = Assert.ThrowsException(() => + InvokePreprocessInOperatorValues("not-a-list")); + + Assert.IsInstanceOfType(exception.InnerException); + } + + [TestMethod] + public void PreprocessInOperatorValues_RejectsMoreThanOneHundredValues() + { + List values = new(); + for (int index = 0; index < 101; index++) + { + values.Add(new IntValueNode(index)); + } + + TargetInvocationException exception = Assert.ThrowsException(() => + InvokePreprocessInOperatorValues(values)); + + Assert.IsInstanceOfType(exception.InnerException); + } + + [TestMethod] + public void PreprocessInOperatorValues_FiltersNullsAndReturnsNullForEmptyValues() + { + List mixed = new() { NullValueNode.Default, new IntValueNode(7) }; + + List filtered = (List)InvokePreprocessInOperatorValues(mixed)!; + + Assert.AreEqual(1, filtered.Count); + Assert.AreEqual("7", filtered[0].Value); + Assert.IsNull(InvokePreprocessInOperatorValues(new List { NullValueNode.Default })); + Assert.IsNull(InvokePreprocessInOperatorValues(new List())); + } + private static GQLFilterParser CreateParserWithDepthLimit(int? depthLimit) { RuntimeConfig config = new( @@ -78,5 +164,12 @@ private static GQLFilterParser CreateParserWithDepthLimit(int? depthLimit) Mock metadataProviderFactory = new(); return new GQLFilterParser(provider, metadataProviderFactory.Object); } + + private static object? InvokePreprocessInOperatorValues(object value) + { + return typeof(FieldFilterParser).GetMethod( + "PreprocessInOperatorValues", + BindingFlags.Static | BindingFlags.NonPublic)!.Invoke(null, new[] { value }); + } } } diff --git a/src/Service.Tests/UnitTests/McpLogNotificationTests.cs b/src/Service.Tests/UnitTests/McpLogNotificationTests.cs index 7f38def9c3..87ef01c608 100644 --- a/src/Service.Tests/UnitTests/McpLogNotificationTests.cs +++ b/src/Service.Tests/UnitTests/McpLogNotificationTests.cs @@ -3,6 +3,7 @@ #nullable enable +using System; using System.IO; using System.Text; using System.Text.Json; @@ -10,6 +11,7 @@ using Azure.DataApiBuilder.Mcp.Telemetry; using Microsoft.Extensions.Logging; using Microsoft.VisualStudio.TestTools.UnitTesting; +using Moq; namespace Azure.DataApiBuilder.Service.Tests.UnitTests { @@ -228,5 +230,89 @@ public void McpLogger_RespectsRuntimeIsEnabledFlip() writer.IsEnabled = false; Assert.IsFalse(logger.IsEnabled(LogLevel.Information)); } + + [TestMethod] + public void McpLogger_BeginScope_ReturnsReusableDisposableScope() + { + McpLogger logger = new("Category", Mock.Of()); + + IDisposable? first = logger.BeginScope("first"); + IDisposable? second = logger.BeginScope("second"); + + Assert.IsNotNull(first); + Assert.AreSame(first, second); + first.Dispose(); + } + + [TestMethod] + public void McpLogger_Disabled_DoesNotInvokeFormatterOrWriter() + { + Mock writer = new(); + writer.SetupGet(x => x.IsEnabled).Returns(false); + McpLogger logger = new("Category", writer.Object); + bool formatterCalled = false; + + logger.Log(LogLevel.Information, default, "state", null, (state, exception) => + { + formatterCalled = true; + return state; + }); + + Assert.IsFalse(formatterCalled); + writer.Verify(x => x.WriteNotification(It.IsAny(), It.IsAny(), It.IsAny()), Times.Never); + } + + [TestMethod] + public void McpLogger_EnabledWithNullFormatter_Throws() + { + Mock writer = new(); + writer.SetupGet(x => x.IsEnabled).Returns(true); + McpLogger logger = new("Category", writer.Object); + + Assert.ThrowsException(() => + logger.Log(LogLevel.Information, default, "state", null, null!)); + } + + [TestMethod] + public void McpLogger_EmptyMessageWithoutException_DoesNotWrite() + { + Mock writer = new(); + writer.SetupGet(x => x.IsEnabled).Returns(true); + McpLogger logger = new("Category", writer.Object); + + logger.Log(LogLevel.Information, default, "state", null, (_, _) => string.Empty); + + writer.Verify(x => x.WriteNotification(It.IsAny(), It.IsAny(), It.IsAny()), Times.Never); + } + + [DataTestMethod] + [DataRow("")] + [DataRow("message")] + public void McpLogger_Exception_AppendsDetailsAndWrites(string formattedMessage) + { + Mock writer = new(); + writer.SetupGet(x => x.IsEnabled).Returns(true); + McpLogger logger = new("Category", writer.Object); + InvalidOperationException exception = new("failure"); + + logger.Log(LogLevel.Error, default, "state", exception, (_, _) => formattedMessage); + + writer.Verify(x => x.WriteNotification( + LogLevel.Error, + "Category", + It.Is(message => message.Contains("InvalidOperationException") && message.Contains("failure"))), Times.Once); + } + + [TestMethod] + public void McpLoggerProvider_DisposeClearsLoggersAndPreventsFurtherCreation() + { + McpLoggerProvider provider = new(Mock.Of()); + _ = provider.CreateLogger("Category"); + + provider.Dispose(); + provider.Dispose(); + + Assert.ThrowsException(() => provider.CreateLogger("Other")); + } } } diff --git a/src/Service.Tests/UnitTests/McpStdioServerContentBlockTests.cs b/src/Service.Tests/UnitTests/McpStdioServerContentBlockTests.cs index 7f361bdbde..7c760dd3f6 100644 --- a/src/Service.Tests/UnitTests/McpStdioServerContentBlockTests.cs +++ b/src/Service.Tests/UnitTests/McpStdioServerContentBlockTests.cs @@ -6,6 +6,7 @@ using System; using System.Collections.Generic; using System.IO; +using System.Linq; using System.Reflection; using System.Text; using System.Text.Json; @@ -170,6 +171,138 @@ public void HandleCallTool_SuccessResult_OmitsIsErrorFromWire() "isError must be absent from the wire for successful tool results."); } + [TestMethod] + public void CoerceToMcpContentBlocks_Null_ReturnsEmptyArray() + { + Assert.AreEqual(0, InvokeCoerceToMcpContentBlocks(null).Length); + } + + [TestMethod] + public void CoerceToMcpContentBlocks_MixedEnumerable_NormalizesStringsAndJson() + { + using JsonDocument json = JsonDocument.Parse("{\"answer\":42}"); + TextContentBlock existingBlock = new() { Text = "already normalized" }; + object input = new + { + Content = new object[] { "plain text", json.RootElement.Clone(), existingBlock } + }; + + object[] result = InvokeCoerceToMcpContentBlocks(input); + + Assert.AreEqual(3, result.Length); + JsonElement textBlock = SerializeToElement(result[0]); + Assert.AreEqual("text", textBlock.GetProperty("type").GetString()); + Assert.AreEqual("plain text", textBlock.GetProperty("text").GetString()); + + JsonElement jsonBlock = SerializeToElement(result[1]); + Assert.AreEqual("application/json", jsonBlock.GetProperty("type").GetString()); + Assert.AreEqual(42, jsonBlock.GetProperty("data").GetProperty("answer").GetInt32()); + Assert.AreSame(existingBlock, result[2]); + } + + [TestMethod] + public void CoerceToMcpContentBlocks_StringContent_ReturnsTextBlock() + { + object[] result = InvokeCoerceToMcpContentBlocks(new { Content = "hello" }); + + JsonElement block = SerializeToElement(result.Single()); + Assert.AreEqual("text", block.GetProperty("type").GetString()); + Assert.AreEqual("hello", block.GetProperty("text").GetString()); + } + + [TestMethod] + public void CoerceToMcpContentBlocks_JsonContent_ReturnsApplicationJsonBlock() + { + using JsonDocument json = JsonDocument.Parse("[1,2,3]"); + + object[] result = InvokeCoerceToMcpContentBlocks(new { Content = json.RootElement.Clone() }); + + JsonElement block = SerializeToElement(result.Single()); + Assert.AreEqual("application/json", block.GetProperty("type").GetString()); + Assert.AreEqual(3, block.GetProperty("data").GetArrayLength()); + } + + [TestMethod] + public void CoerceToMcpContentBlocks_RawJsonElement_ReturnsApplicationJsonBlock() + { + using JsonDocument json = JsonDocument.Parse("true"); + + object[] result = InvokeCoerceToMcpContentBlocks(json.RootElement.Clone()); + + JsonElement block = SerializeToElement(result.Single()); + Assert.AreEqual("application/json", block.GetProperty("type").GetString()); + Assert.IsTrue(block.GetProperty("data").GetBoolean()); + } + + [TestMethod] + public void CoerceToMcpContentBlocks_ObjectWithoutContent_SerializesAsText() + { + object[] result = InvokeCoerceToMcpContentBlocks(new { Status = "ok", Count = 2 }); + + JsonElement block = SerializeToElement(result.Single()); + Assert.AreEqual("text", block.GetProperty("type").GetString()); + StringAssert.Contains(block.GetProperty("text").GetString(), "\"Status\":\"ok\""); + } + + [TestMethod] + public void SafeToString_LargeJson_TruncatesPreview() + { + string result = InvokeSafeToString(new { Value = new string('x', (32 * 1024) + 500) }); + + Assert.IsTrue(result.StartsWith("{\"Value\":\"", StringComparison.Ordinal)); + StringAssert.Contains(result, "... [truncated, total length="); + Assert.IsTrue(result.Length < (33 * 1024)); + } + + [TestMethod] + public void SafeToString_SerializationFailure_UsesToStringFallback() + { + Assert.AreEqual("fallback", InvokeSafeToString(new SelfReferencingObject("fallback"))); + } + + [TestMethod] + public void SafeToString_NullToStringFallback_ReturnsEmptyString() + { + Assert.AreEqual(string.Empty, InvokeSafeToString(new SelfReferencingObject(null))); + } + + [DataTestMethod] + [DataRow("\"abc\"", "abc", DisplayName = "String id")] + [DataRow("9223372036854775807", long.MaxValue, DisplayName = "Int64 id")] + [DataRow("true", null, DisplayName = "Boolean id")] + [DataRow("[]", null, DisplayName = "Array id")] + [DataRow("{}", null, DisplayName = "Object id")] + public void GetIdValue_ConvertsSupportedPrimitiveTypes(string json, object? expected) + { + using JsonDocument document = JsonDocument.Parse(json); + + object? actual = InvokeGetIdValue(document.RootElement); + + Assert.AreEqual(expected, actual); + } + + [TestMethod] + public void GetIdValue_FloatingPointId_ReturnsDouble() + { + using JsonDocument document = JsonDocument.Parse("3.25"); + + object? actual = InvokeGetIdValue(document.RootElement); + + Assert.IsInstanceOfType(actual); + Assert.AreEqual(3.25, (double)actual, 0.0001); + } + + [TestMethod] + public void GetIdValue_NumberOutsideFiniteRange_ReturnsPositiveInfinity() + { + using JsonDocument document = JsonDocument.Parse("1e400"); + + object? actual = InvokeGetIdValue(document.RootElement); + + Assert.IsInstanceOfType(actual); + Assert.IsTrue(double.IsPositiveInfinity((double)actual)); + } + private static (McpStdioServer server, MemoryStream memoryStream, McpStdoutWriter stdoutWriter) CreateServerWithCapturedOutput() { MemoryStream memoryStream = new(); @@ -203,7 +336,7 @@ private static string ReadCapturedOutput(McpStdoutWriter stdoutWriter, MemoryStr return reader.ReadToEnd().TrimEnd(); } - private static object[] InvokeCoerceToMcpContentBlocks(object callResult) + private static object[] InvokeCoerceToMcpContentBlocks(object? callResult) { MethodInfo? coerceMethod = typeof(McpStdioServer).GetMethod( "CoerceToMcpContentBlocks", @@ -211,10 +344,36 @@ private static object[] InvokeCoerceToMcpContentBlocks(object callResult) Assert.IsNotNull(coerceMethod, "Failed to resolve CoerceToMcpContentBlocks via reflection."); - object? result = coerceMethod!.Invoke(obj: null, parameters: new object[] { callResult }); + object? result = coerceMethod!.Invoke(obj: null, parameters: new object?[] { callResult }); return (object[])result!; } + private static string InvokeSafeToString(object value) + { + MethodInfo? safeToStringMethod = typeof(McpStdioServer).GetMethod( + "SafeToString", + BindingFlags.NonPublic | BindingFlags.Static); + + Assert.IsNotNull(safeToStringMethod, "Failed to resolve SafeToString via reflection."); + return (string)safeToStringMethod.Invoke(obj: null, parameters: new[] { value })!; + } + + private static object? InvokeGetIdValue(JsonElement id) + { + MethodInfo? getIdValueMethod = typeof(McpStdioServer).GetMethod( + "GetIdValue", + BindingFlags.NonPublic | BindingFlags.Static); + + Assert.IsNotNull(getIdValueMethod, "Failed to resolve GetIdValue via reflection."); + return getIdValueMethod.Invoke(obj: null, parameters: new object[] { id }); + } + + private static JsonElement SerializeToElement(object value) + { + using JsonDocument document = JsonDocument.Parse(JsonSerializer.Serialize(value)); + return document.RootElement.Clone(); + } + private static void InvokeWriteResult(McpStdioServer server, JsonElement id, object resultObject) { MethodInfo? writeResultMethod = typeof(McpStdioServer).GetMethod( @@ -227,5 +386,19 @@ private static void InvokeWriteResult(McpStdioServer server, JsonElement id, obj // We pass a non-nullable JsonElement, so wrap it as JsonElement? writeResultMethod!.Invoke(server, new object?[] { (JsonElement?)id, resultObject }); } + + private sealed class SelfReferencingObject + { + private readonly string? _text; + + public SelfReferencingObject(string? text) + { + _text = text; + } + + public SelfReferencingObject Self => this; + + public override string? ToString() => _text; + } } } diff --git a/src/Service.Tests/UnitTests/McpStdioServerProtocolTests.cs b/src/Service.Tests/UnitTests/McpStdioServerProtocolTests.cs new file mode 100644 index 0000000000..e20b9defed --- /dev/null +++ b/src/Service.Tests/UnitTests/McpStdioServerProtocolTests.cs @@ -0,0 +1,611 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +#nullable enable + +using System; +using System.Collections.Generic; +using System.Diagnostics.CodeAnalysis; +using System.IO; +using System.Linq; +using System.Security.Claims; +using System.Text.Json; +using System.Threading; +using System.Threading.Tasks; +using Azure.DataApiBuilder.Config; +using Azure.DataApiBuilder.Config.ObjectModel; +using Azure.DataApiBuilder.Core.Configurations; +using Azure.DataApiBuilder.Core.Telemetry; +using Azure.DataApiBuilder.Mcp.Core; +using Azure.DataApiBuilder.Mcp.Model; +using Azure.DataApiBuilder.Mcp.Telemetry; +using Microsoft.AspNetCore.Http; +using Microsoft.Extensions.Configuration; +using Microsoft.Extensions.DependencyInjection; +using Microsoft.Extensions.Logging; +using Microsoft.VisualStudio.TestTools.UnitTesting; +using ModelContextProtocol.Protocol; +using static Azure.DataApiBuilder.Mcp.Model.McpEnums; + +namespace Azure.DataApiBuilder.Service.Tests.UnitTests +{ + [TestClass] + public class McpStdioServerProtocolTests + { + [TestMethod] + public async Task RunAsync_OversizedRequest_ReturnsInvalidRequestAndContinues() + { + string input = new('x', (1024 * 1024) + 1); + input += Environment.NewLine + Request(id: 2, method: "ping") + Environment.NewLine; + (McpStdioServer server, StringWriter output, _) = CreateServer(input); + + await server.RunAsync(CancellationToken.None); + + JsonElement[] responses = ParseResponses(output); + Assert.AreEqual(2, responses.Length); + AssertError(responses[0], expectedId: null, McpStdioJsonRpcErrorCodes.INVALID_REQUEST, "Request too large"); + Assert.IsTrue(responses[1].GetProperty("result").GetProperty("ok").GetBoolean()); + } + + [TestMethod] + public async Task RunAsync_MalformedJson_ReturnsParseErrorAndContinues() + { + string input = "{not-json" + Environment.NewLine + Request(id: 2, method: "ping") + Environment.NewLine; + (McpStdioServer server, StringWriter output, _) = CreateServer(input); + + await server.RunAsync(CancellationToken.None); + + JsonElement[] responses = ParseResponses(output); + Assert.AreEqual(2, responses.Length); + AssertError(responses[0], expectedId: null, McpStdioJsonRpcErrorCodes.PARSE_ERROR, "Parse error"); + Assert.AreEqual(2, responses[1].GetProperty("id").GetInt32()); + } + + [TestMethod] + public async Task RunAsync_MissingMethod_PreservesStringIdInInvalidRequest() + { + (McpStdioServer server, StringWriter output, _) = CreateServer( + "{\"jsonrpc\":\"2.0\",\"id\":\"request-id\"}" + Environment.NewLine); + + await server.RunAsync(CancellationToken.None); + + AssertError(ParseResponses(output).Single(), "request-id", McpStdioJsonRpcErrorCodes.INVALID_REQUEST, "Invalid Request"); + } + + [TestMethod] + public async Task RunAsync_InitializedNotificationProducesNoResponse() + { + string input = Request(id: null, method: "notifications/initialized") + Environment.NewLine + + Request(id: 7, method: "shutdown") + Environment.NewLine + + Request(id: 8, method: "ping") + Environment.NewLine; + (McpStdioServer server, StringWriter output, _) = CreateServer(input); + + await server.RunAsync(CancellationToken.None); + + JsonElement[] responses = ParseResponses(output); + Assert.AreEqual(1, responses.Length); + Assert.AreEqual(7, responses[0].GetProperty("id").GetInt32()); + Assert.IsTrue(responses[0].GetProperty("result").GetProperty("ok").GetBoolean()); + } + + [TestMethod] + public async Task RunAsync_UnknownMethodReturnsMethodNotFoundThenProcessesPing() + { + string input = Request(id: 1, method: "unknown/method") + Environment.NewLine + + Request(id: 2, method: "ping") + Environment.NewLine; + (McpStdioServer server, StringWriter output, _) = CreateServer(input); + + await server.RunAsync(CancellationToken.None); + + JsonElement[] responses = ParseResponses(output); + Assert.AreEqual(2, responses.Length); + AssertError(responses[0], 1L, McpStdioJsonRpcErrorCodes.METHOD_NOT_FOUND, "Method not found: unknown/method"); + Assert.IsTrue(responses[1].GetProperty("result").GetProperty("ok").GetBoolean()); + } + + [TestMethod] + public async Task RunAsync_InitializeConfigurationFailureReturnsInternalError() + { + RuntimeConfigProvider provider = new ThrowingRuntimeConfigProvider(); + string input = Request(id: 1, method: "initialize", @params: "{}") + Environment.NewLine; + (McpStdioServer server, StringWriter output, _) = CreateServer(input, runtimeConfigProvider: provider); + + await server.RunAsync(CancellationToken.None); + + AssertError(ParseResponses(output).Single(), 1L, McpStdioJsonRpcErrorCodes.INTERNAL_ERROR, "Internal error"); + } + + [TestMethod] + public async Task RunAsync_ListToolsReturnsOnlyEnabledToolMetadata() + { + McpToolRegistry registry = new(); + registry.RegisterTool(new RecordingTool("enabled", isEnabled: true)); + registry.RegisterTool(new RecordingTool("disabled", isEnabled: false)); + string input = Request(id: 4, method: "tools/list") + Environment.NewLine; + (McpStdioServer server, StringWriter output, _) = CreateServer(input, registry: registry); + + await server.RunAsync(CancellationToken.None); + + JsonElement tools = ParseResponses(output).Single().GetProperty("result").GetProperty("tools"); + Assert.AreEqual(1, tools.GetArrayLength()); + Assert.AreEqual("enabled", tools[0].GetProperty("name").GetString()); + Assert.AreEqual("Test tool enabled", tools[0].GetProperty("description").GetString()); + Assert.AreEqual("object", tools[0].GetProperty("inputSchema").GetProperty("type").GetString()); + } + + [TestMethod] + public async Task RunAsync_ListToolsConfigurationFailureReturnsInternalError() + { + string input = Request(id: 3, method: "tools/list") + Environment.NewLine; + (McpStdioServer server, StringWriter output, _) = CreateServer( + input, + runtimeConfigProvider: new ThrowingRuntimeConfigProvider()); + + await server.RunAsync(CancellationToken.None); + + AssertError(ParseResponses(output).Single(), 3L, McpStdioJsonRpcErrorCodes.INTERNAL_ERROR, "Internal error"); + } + + [DataTestMethod] + [DataRow(null, "Missing params", DisplayName = "Missing params")] + [DataRow("[]", "Missing params", DisplayName = "Params is an array")] + [DataRow("{}", "Missing tool name", DisplayName = "Missing name")] + [DataRow("{\"name\":null}", "Missing tool name", DisplayName = "Null name")] + [DataRow("{\"name\":\" \"}", "Missing tool name", DisplayName = "Whitespace name")] + [DataRow("{\"name\":\"missing\"}", "Tool not found: missing", DisplayName = "Unknown tool")] + public async Task RunAsync_CallToolRejectsInvalidParameters(string? parameters, string expectedMessage) + { + string input = Request(id: 11, method: "tools/call", @params: parameters) + Environment.NewLine; + (McpStdioServer server, StringWriter output, _) = CreateServer(input); + + await server.RunAsync(CancellationToken.None); + + AssertError(ParseResponses(output).Single(), 11L, McpStdioJsonRpcErrorCodes.INVALID_PARAMS, expectedMessage); + } + + [TestMethod] + public async Task RunAsync_CallToolSupportsNameAndArguments() + { + RecordingTool tool = new("test_tool"); + McpToolRegistry registry = new(); + registry.RegisterTool(tool); + string input = Request( + id: 12, + method: "tools/call", + @params: "{\"name\":\"test_tool\",\"arguments\":{\"value\":42}}") + Environment.NewLine; + (McpStdioServer server, StringWriter output, _) = CreateServer(input, registry: registry); + + await server.RunAsync(CancellationToken.None); + + Assert.AreEqual("{\"value\":42}", tool.ArgumentsJson); + JsonElement result = ParseResponses(output).Single().GetProperty("result"); + Assert.AreEqual("executed test_tool", result.GetProperty("content")[0].GetProperty("text").GetString()); + Assert.IsFalse(result.TryGetProperty("isError", out _)); + } + + [TestMethod] + public async Task RunAsync_CallToolSupportsLegacyToolNameAndMissingArguments() + { + RecordingTool tool = new("legacy_tool"); + McpToolRegistry registry = new(); + registry.RegisterTool(tool); + string input = Request( + id: 13, + method: "tools/call", + @params: "{\"tool\":\"legacy_tool\"}") + Environment.NewLine; + (McpStdioServer server, StringWriter output, _) = CreateServer(input, registry: registry); + + await server.RunAsync(CancellationToken.None); + + Assert.IsTrue(tool.Executed); + Assert.IsNull(tool.ArgumentsJson); + Assert.AreEqual(13, ParseResponses(output).Single().GetProperty("id").GetInt32()); + } + + [TestMethod] + public async Task RunAsync_CallToolPrefersStandardNameOverLegacyToolName() + { + RecordingTool standardTool = new("standard"); + RecordingTool legacyTool = new("legacy"); + McpToolRegistry registry = new(); + registry.RegisterTool(standardTool); + registry.RegisterTool(legacyTool); + string input = Request( + id: 14, + method: "tools/call", + @params: "{\"name\":\"standard\",\"tool\":\"legacy\",\"arguments\":{}}") + Environment.NewLine; + (McpStdioServer server, StringWriter output, _) = CreateServer(input, registry: registry); + + await server.RunAsync(CancellationToken.None); + + Assert.IsTrue(standardTool.Executed); + Assert.IsFalse(legacyTool.Executed); + Assert.AreEqual(1, ParseResponses(output).Length); + } + + [TestMethod] + public async Task RunAsync_CallToolWithConfiguredRoleProvidesScopedIdentityAndClearsAccessor() + { + RecordingTool tool = new("role_tool"); + McpToolRegistry registry = new(); + registry.RegisterTool(tool); + HttpContextAccessor accessor = new(); + IConfiguration configuration = new ConfigurationBuilder() + .AddInMemoryCollection(new Dictionary { ["MCP:Role"] = "writer" }) + .Build(); + string input = Request( + id: 15, + method: "tools/call", + @params: "{\"name\":\"role_tool\",\"arguments\":{}}") + Environment.NewLine; + (McpStdioServer server, StringWriter output, _) = CreateServer( + input, + registry: registry, + configuration: configuration, + httpContextAccessor: accessor); + + await server.RunAsync(CancellationToken.None); + + Assert.AreEqual("writer", tool.ObservedRoleHeader); + Assert.AreEqual("writer", tool.ObservedRoleClaim); + Assert.IsNotNull(tool.ObservedServiceProvider); + Assert.IsNull(accessor.HttpContext, "The request context must be cleared after tool execution."); + Assert.AreEqual(1, ParseResponses(output).Length); + } + + [TestMethod] + public async Task RunAsync_ThrowingToolReturnsInternalErrorAndClearsRoleContext() + { + RecordingTool tool = new("throwing_tool", exception: new InvalidOperationException("failure")); + McpToolRegistry registry = new(); + registry.RegisterTool(tool); + HttpContextAccessor accessor = new(); + IConfiguration configuration = new ConfigurationBuilder() + .AddInMemoryCollection(new Dictionary { ["MCP:Role"] = "reader" }) + .Build(); + string input = Request( + id: 16, + method: "tools/call", + @params: "{\"name\":\"throwing_tool\"}") + Environment.NewLine; + (McpStdioServer server, StringWriter output, _) = CreateServer( + input, + registry: registry, + configuration: configuration, + httpContextAccessor: accessor); + + await server.RunAsync(CancellationToken.None); + + Assert.IsTrue(tool.Executed); + Assert.IsNull(accessor.HttpContext); + AssertError(ParseResponses(output).Single(), 16L, McpStdioJsonRpcErrorCodes.INTERNAL_ERROR, "Internal error"); + } + + [DataTestMethod] + [DataRow(null, DisplayName = "Missing params")] + [DataRow("{}", DisplayName = "Missing level")] + [DataRow("{\"level\":null}", DisplayName = "Null level")] + [DataRow("{\"level\":\" \"}", DisplayName = "Whitespace level")] + public async Task RunAsync_SetLogLevelRejectsMissingOrInvalidLevel(string? parameters) + { + string input = Request(id: 21, method: "logging/setLevel", @params: parameters) + Environment.NewLine; + (McpStdioServer server, StringWriter output, _) = CreateServer(input); + + await server.RunAsync(CancellationToken.None); + + AssertError( + ParseResponses(output).Single(), + 21L, + McpStdioJsonRpcErrorCodes.INVALID_PARAMS, + "Missing or invalid 'level' parameter"); + } + + [TestMethod] + public async Task RunAsync_SetLogLevelWithoutControllerReturnsSuccess() + { + string input = Request( + id: 22, + method: "logging/setLevel", + @params: "{\"level\":\"debug\"}") + Environment.NewLine; + (McpStdioServer server, StringWriter output, _) = CreateServer(input); + + await server.RunAsync(CancellationToken.None); + + Assert.AreEqual(JsonValueKind.Object, ParseResponses(output).Single().GetProperty("result").ValueKind); + } + + [TestMethod] + public async Task RunAsync_SetLogLevelInvalidValueHasNoSideEffects() + { + RecordingLogLevelController controller = new(updateResult: true); + RecordingNotificationWriter writer = new() { IsEnabled = false }; + string input = Request( + id: 23, + method: "logging/setLevel", + @params: "{\"level\":\"verbose\"}") + Environment.NewLine; + (McpStdioServer server, StringWriter output, _) = CreateServer( + input, + logLevelController: controller, + notificationWriter: writer); + + await server.RunAsync(CancellationToken.None); + + Assert.IsNull(controller.LastLevel); + Assert.IsFalse(writer.IsEnabled); + Assert.AreEqual(JsonValueKind.Object, ParseResponses(output).Single().GetProperty("result").ValueKind); + } + + [TestMethod] + public async Task RunAsync_SetLogLevelNoneDisablesNotifications() + { + RecordingLogLevelController controller = new(updateResult: false); + RecordingNotificationWriter writer = new() { IsEnabled = true }; + string input = Request( + id: 24, + method: "logging/setLevel", + @params: "{\"level\":\"NONE\"}") + Environment.NewLine; + (McpStdioServer server, StringWriter output, _) = CreateServer( + input, + logLevelController: controller, + notificationWriter: writer); + + await server.RunAsync(CancellationToken.None); + + Assert.AreEqual("NONE", controller.LastLevel); + Assert.IsFalse(writer.IsEnabled); + Assert.AreEqual(1, ParseResponses(output).Length); + } + + [TestMethod] + public async Task RunAsync_SetLogLevelValidValueEnablesNotificationsAndRestoresStderr() + { + RecordingLogLevelController controller = new(updateResult: true); + RecordingNotificationWriter writer = new() { IsEnabled = false }; + string input = Request( + id: 25, + method: "logging/setLevel", + @params: "{\"level\":\"warning\"}") + Environment.NewLine; + (McpStdioServer server, StringWriter output, _) = CreateServer( + input, + logLevelController: controller, + notificationWriter: writer); + TextWriter originalError = Console.Error; + + try + { + Console.SetError(TextWriter.Null); + await server.RunAsync(CancellationToken.None); + + Assert.AreNotSame(TextWriter.Null, Console.Error); + } + finally + { + Console.SetError(originalError); + } + + Assert.AreEqual("warning", controller.LastLevel); + Assert.IsTrue(writer.IsEnabled); + Assert.AreEqual(1, ParseResponses(output).Length); + } + + private static (McpStdioServer Server, StringWriter Output, IServiceProvider Services) CreateServer( + string input, + McpToolRegistry? registry = null, + RuntimeConfigProvider? runtimeConfigProvider = null, + IConfiguration? configuration = null, + ILogLevelController? logLevelController = null, + IMcpLogNotificationWriter? notificationWriter = null, + IHttpContextAccessor? httpContextAccessor = null) + { + registry ??= new McpToolRegistry(); + StringWriter output = new(); + ServiceCollection services = new(); + services.AddSingleton(new McpStdoutWriter(output)); + services.AddSingleton(registry); + services.AddSingleton(runtimeConfigProvider ?? new StubRuntimeConfigProvider(CreateRuntimeConfig())); + services.AddSingleton(configuration ?? new ConfigurationBuilder().Build()); + + if (logLevelController is not null) + { + services.AddSingleton(logLevelController); + } + + if (notificationWriter is not null) + { + services.AddSingleton(notificationWriter); + } + + if (httpContextAccessor is not null) + { + services.AddSingleton(httpContextAccessor); + } + + IServiceProvider serviceProvider = services.BuildServiceProvider(); + McpStdioServer server = new(registry, serviceProvider, new StringReader(input)); + return (server, output, serviceProvider); + } + + private static string Request(long? id, string method, string? @params = null) + { + string idProperty = id.HasValue ? $",\"id\":{id.Value}" : string.Empty; + string paramsProperty = @params is not null ? $",\"params\":{@params}" : string.Empty; + return $"{{\"jsonrpc\":\"2.0\"{idProperty},\"method\":\"{method}\"{paramsProperty}}}"; + } + + private static JsonElement[] ParseResponses(StringWriter output) + { + return output.ToString() + .Split(Environment.NewLine, StringSplitOptions.RemoveEmptyEntries) + .Select(line => + { + using JsonDocument response = JsonDocument.Parse(line); + return response.RootElement.Clone(); + }) + .ToArray(); + } + + private static void AssertError(JsonElement response, object? expectedId, int expectedCode, string expectedMessage) + { + Assert.AreEqual(McpStdioJsonRpcErrorCodes.JSON_RPC_VERSION, response.GetProperty("jsonrpc").GetString()); + if (expectedId is null) + { + Assert.AreEqual(JsonValueKind.Null, response.GetProperty("id").ValueKind); + } + else if (expectedId is string expectedString) + { + Assert.AreEqual(expectedString, response.GetProperty("id").GetString()); + } + else + { + Assert.AreEqual(Convert.ToInt64(expectedId), response.GetProperty("id").GetInt64()); + } + + JsonElement error = response.GetProperty("error"); + Assert.AreEqual(expectedCode, error.GetProperty("code").GetInt32()); + Assert.AreEqual(expectedMessage, error.GetProperty("message").GetString()); + } + + private static RuntimeConfig CreateRuntimeConfig() + { + return new RuntimeConfig( + Schema: RuntimeConfig.DEFAULT_CONFIG_SCHEMA_LINK, + DataSource: null, + Entities: new RuntimeEntities(new Dictionary()), + Runtime: new RuntimeOptions( + Rest: null, + GraphQL: null, + Mcp: new McpRuntimeOptions(Enabled: true), + Host: null)); + } + + private sealed class RecordingTool : IMcpTool + { + private readonly string _name; + private readonly bool _isEnabled; + private readonly Exception? _exception; + + public RecordingTool(string name, bool isEnabled = true, Exception? exception = null) + { + _name = name; + _isEnabled = isEnabled; + _exception = exception; + } + + public ToolType ToolType => ToolType.BuiltIn; + + public bool Executed { get; private set; } + + public string? ArgumentsJson { get; private set; } + + public IServiceProvider? ObservedServiceProvider { get; private set; } + + public string? ObservedRoleHeader { get; private set; } + + public string? ObservedRoleClaim { get; private set; } + + public bool IsEnabled(RuntimeConfig config) => _isEnabled; + + public Tool GetToolMetadata() + { + using JsonDocument schema = JsonDocument.Parse("{\"type\":\"object\"}"); + return new Tool + { + Name = _name, + Description = $"Test tool {_name}", + InputSchema = schema.RootElement.Clone() + }; + } + + public Task ExecuteAsync( + JsonDocument? arguments, + IServiceProvider serviceProvider, + CancellationToken cancellationToken = default) + { + Executed = true; + ArgumentsJson = arguments?.RootElement.GetRawText(); + ObservedServiceProvider = serviceProvider; + IHttpContextAccessor? accessor = serviceProvider.GetService(); + ObservedRoleHeader = accessor?.HttpContext?.Request.Headers["X-MS-API-ROLE"].ToString(); + ObservedRoleClaim = accessor?.HttpContext?.User.FindFirst(ClaimTypes.Role)?.Value; + + if (_exception is not null) + { + throw _exception; + } + + return Task.FromResult(new CallToolResult + { + Content = new List + { + new TextContentBlock { Text = $"executed {_name}" } + } + }); + } + } + + private sealed class RecordingLogLevelController : ILogLevelController + { + private readonly bool _updateResult; + + public RecordingLogLevelController(bool updateResult) + { + _updateResult = updateResult; + } + + public bool IsCliOverriding => false; + + public bool IsConfigOverriding => false; + + public bool IsAgentOverriding => LastLevel is not null; + + public string? LastLevel { get; private set; } + + public bool UpdateFromMcp(string mcpLevel) + { + LastLevel = mcpLevel; + return _updateResult; + } + } + + private sealed class RecordingNotificationWriter : IMcpLogNotificationWriter + { + public bool IsEnabled { get; set; } + + public void WriteNotification(LogLevel logLevel, string categoryName, string message) + { + } + } + + private sealed class StubRuntimeConfigProvider : RuntimeConfigProvider + { + private readonly RuntimeConfig _runtimeConfig; + + public StubRuntimeConfigProvider(RuntimeConfig runtimeConfig) + : base(new StubRuntimeConfigLoader()) + { + _runtimeConfig = runtimeConfig; + } + + public override RuntimeConfig GetConfig() => _runtimeConfig; + } + + private sealed class ThrowingRuntimeConfigProvider : RuntimeConfigProvider + { + public ThrowingRuntimeConfigProvider() + : base(new StubRuntimeConfigLoader()) + { + } + + public override RuntimeConfig GetConfig() => throw new InvalidOperationException("Configuration unavailable."); + } + + private sealed class StubRuntimeConfigLoader : RuntimeConfigLoader + { + public override bool TryLoadKnownConfig([NotNullWhen(true)] out RuntimeConfig? config, bool replaceEnvVar = false) + { + config = null; + return false; + } + + public override string GetPublishedDraftSchemaLink() => RuntimeConfig.DEFAULT_CONFIG_SCHEMA_LINK; + } + } +} diff --git a/src/Service.Tests/UnitTests/McpTelemetryTests.cs b/src/Service.Tests/UnitTests/McpTelemetryTests.cs index c2927cd72d..4094f5e94c 100644 --- a/src/Service.Tests/UnitTests/McpTelemetryTests.cs +++ b/src/Service.Tests/UnitTests/McpTelemetryTests.cs @@ -7,6 +7,7 @@ using System.Collections.Generic; using System.Diagnostics; using System.Linq; +using System.Net; using System.Text.Json; using System.Threading; using System.Threading.Tasks; @@ -126,6 +127,48 @@ private static Exception CreateExceptionByTypeName(string typeName) #endregion + #region General Telemetry Trace Helpers + + [TestMethod] + public void TrackQueryActivityStarted_SetsDatabaseTags() + { + using Activity activity = CreateActivity(); + + activity.TrackQueryActivityStarted(DatabaseType.MSSQL, "default"); + + Activity recorded = StopAndGetRecordedActivity(activity); + Assert.AreEqual(DatabaseType.MSSQL, recorded.GetTagItem("data-source.type")); + Assert.AreEqual("default", recorded.GetTagItem("data-source.name")); + } + + [TestMethod] + public void TrackMainControllerActivityFinished_SetsStatusCodeTag() + { + using Activity activity = CreateActivity(); + + activity.TrackMainControllerActivityFinished(HttpStatusCode.Accepted); + + Activity recorded = StopAndGetRecordedActivity(activity); + Assert.AreEqual(HttpStatusCode.Accepted, recorded.GetTagItem("status.code")); + } + + [TestMethod] + public void TrackMainControllerActivityFinishedWithException_SetsErrorDetails() + { + using Activity activity = CreateActivity(); + InvalidOperationException exception = new("failed"); + + activity.TrackMainControllerActivityFinishedWithException(exception, HttpStatusCode.InternalServerError); + + Activity recorded = StopAndGetRecordedActivity(activity); + Assert.AreEqual(ActivityStatusCode.Error, recorded.Status); + Assert.AreEqual(nameof(InvalidOperationException), recorded.GetTagItem("error.type")); + Assert.AreEqual("failed", recorded.GetTagItem("error.message")); + Assert.AreEqual(HttpStatusCode.InternalServerError, recorded.GetTagItem("status.code")); + } + + #endregion + #region TrackMcpToolExecutionStarted /// diff --git a/src/Service.Tests/UnitTests/MultipleCreateOrderHelperEdgeTests.cs b/src/Service.Tests/UnitTests/MultipleCreateOrderHelperEdgeTests.cs new file mode 100644 index 0000000000..cce39ab839 --- /dev/null +++ b/src/Service.Tests/UnitTests/MultipleCreateOrderHelperEdgeTests.cs @@ -0,0 +1,154 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +using System.Collections.Generic; +using Azure.DataApiBuilder.Config.DatabasePrimitives; +using Azure.DataApiBuilder.Core.Resolvers; +using Azure.DataApiBuilder.Core.Services; +using Azure.DataApiBuilder.Service.Exceptions; +using HotChocolate.Language; +using HotChocolate.Resolvers; +using Microsoft.VisualStudio.TestTools.UnitTesting; +using Moq; + +namespace Azure.DataApiBuilder.Service.Tests.UnitTests +{ + [TestClass] + public class MultipleCreateOrderHelperEdgeTests + { + [TestMethod] + public void GetReferencingEntityName_RejectsMissingMetadata() + { + Mock metadata = CreateMetadata(new Dictionary()); + + Assert.ThrowsException(() => MultipleCreateOrderHelper.GetReferencingEntityName( + new Mock().Object, + "relationship", + "Source", + "Target", + metadata.Object, + new Dictionary(), + null, + 2)); + } + + [TestMethod] + public void GetReferencingEntityName_RejectsNonTableEntities() + { + Dictionary objects = new() + { + ["Source"] = new DatabaseView("dbo", "source"), + ["Target"] = new DatabaseTable("dbo", "target") + }; + Mock metadata = CreateMetadata(objects); + + Assert.ThrowsException(() => MultipleCreateOrderHelper.GetReferencingEntityName( + new Mock().Object, + "relationship", + "Source", + "Target", + metadata.Object, + new Dictionary(), + null, + 1)); + } + + [TestMethod] + public void GetReferencingEntityName_RejectsEntitiesBackedBySameTable() + { + Dictionary objects = new() + { + ["Source"] = new DatabaseTable("dbo", "shared"), + ["Target"] = new DatabaseTable("dbo", "shared") + }; + Mock metadata = CreateMetadata(objects); + + Assert.ThrowsException(() => MultipleCreateOrderHelper.GetReferencingEntityName( + new Mock().Object, + "relationship", + "Source", + "Target", + metadata.Object, + new Dictionary(), + null, + 1)); + } + + [TestMethod] + public void GetReferencingEntityName_ManyToManyUsesLinkingTable() + { + Dictionary objects = new() + { + ["Source"] = new DatabaseTable("dbo", "source"), + ["Target"] = new DatabaseTable("dbo", "target") + }; + Mock metadata = CreateMetadata(objects); + + string result = MultipleCreateOrderHelper.GetReferencingEntityName( + new Mock().Object, + "relationship", + "Source", + "Target", + metadata.Object, + new Dictionary(), + null, + 1, + isMNRelationship: true); + + Assert.AreEqual(string.Empty, result); + } + + [TestMethod] + public void GetBackingColumnDataFromFields_ReturnsOnlyMappedScalarFields() + { + Mock metadata = new(); + metadata.Setup(x => x.TryGetArrayElementSyntaxKind("Book", It.IsAny(), out It.Ref.IsAny)) + .Returns(false); + metadata.Setup(x => x.TryGetBackingColumn("Book", "title", out It.Ref.IsAny)) + .Returns((string _, string _, out string? backing) => + { + backing = "book_title"; + return true; + }); + List fields = new() + { + new("title", "DAB"), + new("unmapped", 7), + new("relationship", new ObjectValueNode(new ObjectFieldNode("id", 1))) + }; + Dictionary result = MultipleCreateOrderHelper.GetBackingColumnDataFromFields( + new Mock().Object, + "Book", + fields, + metadata.Object); + + Assert.AreEqual(1, result.Count); + Assert.IsTrue(result.ContainsKey("book_title")); + } + + [TestMethod] + public void GetBackingColumnDataFromFields_ArrayElementKindCanBeScalar() + { + Mock metadata = new(); + SyntaxKind scalarKind = SyntaxKind.FloatValue; + metadata.Setup(x => x.TryGetArrayElementSyntaxKind("Book", "vector", out scalarKind)).Returns(true); + string? backing = "embedding"; + metadata.Setup(x => x.TryGetBackingColumn("Book", "vector", out backing)).Returns(true); + + Dictionary result = MultipleCreateOrderHelper.GetBackingColumnDataFromFields( + new Mock().Object, + "Book", + new[] { new ObjectFieldNode("vector", new ListValueNode(new FloatValueNode(1.5))) }, + metadata.Object); + + Assert.IsTrue(result.ContainsKey("embedding")); + } + + private static Mock CreateMetadata(IReadOnlyDictionary objects) + { + Mock metadata = new(); + metadata.Setup(x => x.GetEntityNamesAndDbObjects()).Returns(objects); + return metadata; + } + } +} diff --git a/src/Service.Tests/UnitTests/PureUtilityCoverageTests.cs b/src/Service.Tests/UnitTests/PureUtilityCoverageTests.cs new file mode 100644 index 0000000000..92c83a5573 --- /dev/null +++ b/src/Service.Tests/UnitTests/PureUtilityCoverageTests.cs @@ -0,0 +1,101 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +using System; +using System.Text.Json; +using Azure.DataApiBuilder.Config; +using Azure.DataApiBuilder.Core.Configurations; +using Azure.DataApiBuilder.Core.Models; +using MetadataTypeConverter = Azure.DataApiBuilder.Core.Services.MetadataProviders.Converters.TypeConverter; +using Azure.DataApiBuilder.Core.Telemetry; +using Microsoft.Extensions.Logging; +using Microsoft.VisualStudio.TestTools.UnitTesting; + +namespace Azure.DataApiBuilder.Service.Tests.UnitTests +{ + [TestClass] + public class PureUtilityCoverageTests + { + [DataTestMethod] + [DataRow("debug", LogLevel.Debug)] + [DataRow("INFO", LogLevel.Information)] + [DataRow("notice", LogLevel.Information)] + [DataRow("warning", LogLevel.Warning)] + [DataRow("error", LogLevel.Error)] + [DataRow("critical", LogLevel.Critical)] + [DataRow("alert", LogLevel.Critical)] + [DataRow("emergency", LogLevel.Critical)] + public void McpLogLevelConverter_RecognizedValuesMapToLogLevels(string value, LogLevel expected) + { + Assert.IsTrue(McpLogLevelConverter.TryConvertFromMcp(value, out LogLevel actual)); + Assert.AreEqual(expected, actual); + } + + [DataTestMethod] + [DataRow(null)] + [DataRow("")] + [DataRow(" ")] + [DataRow("unknown")] + public void McpLogLevelConverter_InvalidValuesReturnFalse(string? value) + { + Assert.IsFalse(McpLogLevelConverter.TryConvertFromMcp(value!, out _)); + } + + [DataTestMethod] + [DataRow(LogLevel.Trace, "debug")] + [DataRow(LogLevel.Debug, "debug")] + [DataRow(LogLevel.Information, "info")] + [DataRow(LogLevel.Warning, "warning")] + [DataRow(LogLevel.Error, "error")] + [DataRow(LogLevel.Critical, "critical")] + [DataRow(LogLevel.None, "debug")] + [DataRow((LogLevel)99, "info")] + public void McpLogLevelConverter_AllLogLevelsMapToProtocolValues(LogLevel value, string expected) + { + Assert.AreEqual(expected, McpLogLevelConverter.ConvertToMcp(value)); + } + + [TestMethod] + public void TryValidateEntityRestPath_RejectsOverlongAndColonPaths() + { + Assert.IsFalse(RuntimeConfigValidatorUtil.TryValidateEntityRestPath(new string('a', 2049), out string? longError)); + StringAssert.Contains(longError, "maximum allowed length"); + + Assert.IsFalse(RuntimeConfigValidatorUtil.TryValidateEntityRestPath("books:archive", out string? colonError)); + StringAssert.Contains(colonError, "reserved character"); + } + + [TestMethod] + public void DabChangeToken_SignalChangeUpdatesStateAndInvokesCallback() + { + DabChangeToken token = new(); + bool callbackInvoked = false; + using IDisposable registration = token.RegisterChangeCallback(_ => callbackInvoked = true, null); + + Assert.IsTrue(token.ActiveChangeCallbacks); + Assert.IsFalse(token.HasChanged); + token.SignalChange(); + Assert.IsTrue(token.HasChanged); + Assert.IsTrue(callbackInvoked); + } + + [TestMethod] + public void MutationResolver_PrimaryConstructorPopulatesProperties() + { + MutationResolver resolver = new("id", null!, "database", "container", "fields", "table"); + + Assert.AreEqual("id", resolver.Id); + Assert.IsNull(resolver.OperationType); + Assert.AreEqual("table", resolver.Table); + } + + [TestMethod] + public void MetadataTypeConverter_NonStringInputThrows() + { + JsonSerializerOptions options = new(); + options.Converters.Add(new MetadataTypeConverter()); + + Assert.ThrowsException(() => JsonSerializer.Deserialize("42", options)); + } + } +} \ No newline at end of file diff --git a/src/Service.Tests/UnitTests/QueryExecutorHelperTests.cs b/src/Service.Tests/UnitTests/QueryExecutorHelperTests.cs new file mode 100644 index 0000000000..93a3ac2169 --- /dev/null +++ b/src/Service.Tests/UnitTests/QueryExecutorHelperTests.cs @@ -0,0 +1,432 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +using System; +using System.Collections.Generic; +using System.Data; +using System.Data.Common; +using System.Linq; +using System.Net; +using System.Reflection; +using System.Runtime.CompilerServices; +using System.Text; +using System.Text.Json; +using System.Text.Json.Nodes; +using System.Threading.Tasks; +using Azure.DataApiBuilder.Core.Resolvers; +using Azure.DataApiBuilder.Core.Models; +using Azure.DataApiBuilder.Core.Services; +using Azure.DataApiBuilder.Config.ObjectModel; +using Azure.DataApiBuilder.Core.Configurations; +using Azure.DataApiBuilder.Service.Exceptions; +using Microsoft.AspNetCore.Http; +using Microsoft.Data.SqlClient; +using Microsoft.Extensions.Logging.Abstractions; +using Microsoft.VisualStudio.TestTools.UnitTesting; +using Moq; + +namespace Azure.DataApiBuilder.Service.Tests.UnitTests +{ + [TestClass] + public class QueryExecutorHelperTests + { + [DataTestMethod] + [DataRow("[1,2]", typeof(JsonArray))] + [DataRow("{\"value\":1}", typeof(JsonObject))] + [DataRow("\"value\"", typeof(JsonValue))] + [DataRow("42", typeof(JsonValue))] + [DataRow("null", null)] + public void FromJsonElement_CreatesExpectedNodeType(string json, Type? expectedType) + { + using JsonDocument document = JsonDocument.Parse(json); + MethodInfo method = typeof(QueryExecutor).GetMethod( + "FromJsonElement", BindingFlags.Static | BindingFlags.NonPublic)!; + + JsonNode? result = (JsonNode?)method.Invoke(null, new object[] { document.RootElement }); + + if (expectedType is null) + { + Assert.IsNull(result); + } + else + { + Assert.IsInstanceOfType(result, expectedType); + } + } + + [DataTestMethod] + [DataRow(100L, 101L, true)] + [DataRow(100L, 100L, false)] + [DataRow(100L, 0L, false)] + public void ValidateSize_EnforcesOnlyValuesOverLimit(long available, long requested, bool throws) + { + MsSqlQueryExecutor executor = (MsSqlQueryExecutor)RuntimeHelpers.GetUninitializedObject(typeof(MsSqlQueryExecutor)); + Type baseType = typeof(MsSqlQueryExecutor).BaseType!; + baseType.GetField("_maxResponseSizeMB", BindingFlags.Instance | BindingFlags.NonPublic)! + .SetValue(executor, 1); + MethodInfo method = baseType.GetMethod("ValidateSize", BindingFlags.Instance | BindingFlags.NonPublic)!; + + if (throws) + { + TargetInvocationException exception = Assert.ThrowsException( + () => method.Invoke(executor, new object[] { available, requested })); + Assert.IsInstanceOfType(exception.InnerException); + } + else + { + method.Invoke(executor, new object[] { available, requested }); + } + } + + [TestMethod] + public void ResultPropertyHandlers_ReturnReaderState() + { + MsSqlQueryExecutor executor = CreateExecutor(); + Mock reader = new(); + reader.SetupGet(x => x.RecordsAffected).Returns(3); + reader.SetupGet(x => x.HasRows).Returns(true); + + Dictionary sync = executor.GetResultProperties(reader.Object); + Dictionary asyncResult = executor.GetResultPropertiesAsync(reader.Object).Result; + + Assert.AreEqual(3, sync[nameof(DbDataReader.RecordsAffected)]); + Assert.AreEqual(true, sync[nameof(DbDataReader.HasRows)]); + CollectionAssert.AreEquivalent(sync, asyncResult); + } + + [TestMethod] + public void StreamCharData_ReadsContentAndHandlesEmptyCells() + { + MsSqlQueryExecutor executor = CreateExecutor(); + Mock reader = new(); + reader.Setup(x => x.GetChars(0, 0, null, 0, 0)).Returns(3); + reader.Setup(x => x.GetChars(0, 0, It.IsAny(), 0, 3)) + .Callback((int _, long _, char[]? buffer, int _, int _) => "DAB".CopyTo(0, buffer!, 0, 3)) + .Returns(3); + StringBuilder result = new(); + + Assert.AreEqual(3, executor.StreamCharData(reader.Object, 3, result, 0)); + Assert.AreEqual("DAB", result.ToString()); + + reader.Setup(x => x.GetChars(0, 0, null, 0, 0)).Returns(0); + Assert.AreEqual(0, executor.StreamCharData(reader.Object, 0, result, 0)); + } + + [TestMethod] + public void StreamByteData_ReadsContentAndHandlesEmptyCells() + { + MsSqlQueryExecutor executor = CreateExecutor(); + Mock reader = new(); + reader.Setup(x => x.GetBytes(0, 0, null, 0, 0)).Returns(2); + reader.Setup(x => x.GetBytes(0, 0, It.IsAny(), 0, 2)) + .Callback((int _, long _, byte[]? buffer, int _, int _) => + { + buffer![0] = 1; + buffer[1] = 2; + }) + .Returns(2); + + Assert.AreEqual(2, executor.StreamByteData(reader.Object, 2, 0, out byte[]? bytes)); + CollectionAssert.AreEqual(new byte[] { 1, 2 }, bytes); + + reader.Setup(x => x.GetBytes(0, 0, null, 0, 0)).Returns(0); + Assert.AreEqual(0, executor.StreamByteData(reader.Object, 0, 0, out bytes)); + Assert.AreEqual(0, bytes!.Length); + } + + [DataTestMethod] + [DataRow(typeof(string), 3)] + [DataRow(typeof(byte[]), 2)] + [DataRow(typeof(int), 4)] + public void StreamDataIntoResultSetRow_HandlesSupportedColumnKinds(Type fieldType, int expectedSize) + { + MsSqlQueryExecutor executor = CreateExecutor(); + Mock reader = new(); + reader.Setup(x => x.GetFieldType(0)).Returns(fieldType); + reader.Setup(x => x.GetChars(0, 0, null, 0, 0)).Returns(3); + reader.Setup(x => x.GetChars(0, 0, It.IsAny(), 0, 3)).Returns(3); + reader.Setup(x => x.GetBytes(0, 0, null, 0, 0)).Returns(2); + reader.Setup(x => x.GetBytes(0, 0, It.IsAny(), 0, 2)).Returns(2); + reader.Setup(x => x["value"]).Returns(42); + DbResultSetRow row = new(); + + int size = executor.StreamDataIntoDbResultSetRow(reader.Object, row, "value", 4, 0, 10); + + Assert.AreEqual(expectedSize, size); + Assert.IsTrue(row.Columns.ContainsKey("value")); + } + + [TestMethod] + public void AddDbExecutionTime_AccumulatesAndIgnoresMissingContext() + { + MsSqlQueryExecutor executor = CreateExecutor(); + DefaultHttpContext context = new(); + typeof(QueryExecutor).GetField("k__BackingField", BindingFlags.Instance | BindingFlags.NonPublic)! + .SetValue(executor, new HttpContextAccessor { HttpContext = context }); + + executor.AddDbExecutionTimeToMiddlewareContext(2); + executor.AddDbExecutionTimeToMiddlewareContext(3); + + Assert.AreEqual(5L, context.Items["TotalDbExecutionTime"]); + typeof(QueryExecutor).GetField("k__BackingField", BindingFlags.Instance | BindingFlags.NonPublic)! + .SetValue(executor, new HttpContextAccessor()); + executor.AddDbExecutionTimeToMiddlewareContext(1); + } + + [DataTestMethod] + [DataRow(false)] + [DataRow(true)] + public async Task GetJsonResultAsync_HandlesRowsAndNoRows(bool hasRows) + { + MsSqlQueryExecutor executor = CreateExecutor(); + Mock reader = new(); + reader.SetupGet(x => x.HasRows).Returns(hasRows); + reader.SetupSequence(x => x.ReadAsync(It.IsAny())) + .ReturnsAsync(hasRows) + .ReturnsAsync(false); + reader.Setup(x => x.GetString(0)).Returns("{\"value\":7}"); + + JsonDocument? result = await executor.GetJsonResultAsync(reader.Object); + + if (hasRows) + { + Assert.AreEqual(7, result!.RootElement.GetProperty("value").GetInt32()); + result.Dispose(); + } + else + { + Assert.IsNull(result); + } + } + + [TestMethod] + public async Task ReadHelpers_TranslateDatabaseExceptions() + { + MsSqlQueryExecutor executor = CreateExecutor(); + Mock reader = new(); + reader.Setup(x => x.Read()).Throws(new TestDbException()); + reader.Setup(x => x.ReadAsync(It.IsAny())) + .ThrowsAsync(new TestDbException()); + + Assert.ThrowsException(() => executor.Read(reader.Object)); + await Assert.ThrowsExceptionAsync(() => executor.ReadAsync(reader.Object)); + } + + [DataTestMethod] + [DataRow(false, false)] + [DataRow(false, true)] + [DataRow(true, false)] + public async Task ExtractResultSet_HandlesColumnsNullsAndFiltering(bool useAsync, bool filterOutColumn) + { + MsSqlQueryExecutor executor = CreateExecutor(maxResponseSizeEnabled: false); + Mock reader = CreateSingleRowReader(valueIsNull: false); + List? columns = filterOutColumn ? new List { "other" } : null; + + DbResultSet result = useAsync + ? await executor.ExtractResultSetFromDbDataReaderAsync(reader.Object, columns) + : executor.ExtractResultSetFromDbDataReader(reader.Object, columns); + + Assert.AreEqual(1, result.Rows.Count); + Assert.AreEqual(filterOutColumn ? 0 : 1, result.Rows[0].Columns.Count); + } + + [TestMethod] + public async Task ExtractResultSet_HandlesNullAndStreamedValues() + { + MsSqlQueryExecutor executor = CreateExecutor(maxResponseSizeEnabled: false); + Mock nullReader = CreateSingleRowReader(valueIsNull: true); + + DbResultSet nullResult = await executor.ExtractResultSetFromDbDataReaderAsync(nullReader.Object); + + Assert.IsNull(nullResult.Rows.Single().Columns["value"]); + + executor = CreateExecutor(maxResponseSizeEnabled: true); + Mock streamedReader = CreateSingleRowReader(valueIsNull: false); + streamedReader.Setup(x => x.GetFieldType(0)).Returns(typeof(string)); + streamedReader.Setup(x => x.GetChars(0, 0, null, 0, 0)).Returns(3); + streamedReader.Setup(x => x.GetChars(0, 0, It.IsAny(), 0, 3)) + .Callback((int _, long _, char[]? buffer, int _, int _) => "DAB".CopyTo(0, buffer!, 0, 3)) + .Returns(3); + + DbResultSet streamedResult = executor.ExtractResultSetFromDbDataReader(streamedReader.Object); + + Assert.AreEqual("DAB", streamedResult.Rows.Single().Columns["value"]); + } + + [TestMethod] + public async Task GetMultipleResultSets_ReturnsFirstPopulatedResultAsUpdate() + { + QueryExecutor executor = CreateBaseExecutor(); + Mock reader = CreateSingleRowReader(valueIsNull: false); + + DbResultSet result = await executor.GetMultipleResultSetsIfAnyAsync(reader.Object); + + Assert.AreEqual(true, result.ResultProperties[SqlMutationEngine.IS_UPDATE_RESULT_SET]); + reader.Verify(x => x.NextResultAsync(It.IsAny()), Times.Never); + } + + [TestMethod] + public async Task GetMultipleResultSets_ThrowsWhenNeitherMutationProducesRows() + { + QueryExecutor executor = CreateBaseExecutor(); + Mock reader = new(); + reader.SetupGet(x => x.RecordsAffected).Returns(0); + reader.SetupGet(x => x.HasRows).Returns(false); + reader.Setup(x => x.ReadAsync(It.IsAny())).ReturnsAsync(false); + reader.Setup(x => x.NextResultAsync(It.IsAny())).ReturnsAsync(false); + + await Assert.ThrowsExceptionAsync( + () => executor.GetMultipleResultSetsIfAnyAsync(reader.Object, new List { "id=1", "Book" })); + } + + [TestMethod] + public async Task MsSqlGetMultipleResultSets_MissingCountResultThrows() + { + MsSqlQueryExecutor executor = CreateExecutor(maxResponseSizeEnabled: false); + Mock reader = new(); + reader.SetupGet(x => x.RecordsAffected).Returns(0); + reader.SetupGet(x => x.HasRows).Returns(false); + reader.Setup(x => x.ReadAsync(It.IsAny())).ReturnsAsync(false); + + await Assert.ThrowsExceptionAsync(() => + executor.GetMultipleResultSetsIfAnyAsync(reader.Object)); + } + + [TestMethod] + public async Task MsSqlGetMultipleResultSets_NoMutationResultWithArgumentsThrowsNotFound() + { + MsSqlQueryExecutor executor = CreateExecutor(maxResponseSizeEnabled: false); + DataTable schema = new(); + schema.Columns.Add("ColumnName", typeof(string)); + schema.Columns.Add("ColumnSize", typeof(int)); + schema.Rows.Add(MsSqlQueryBuilder.COUNT_ROWS_WITH_GIVEN_PK, 4); + Mock reader = new(); + reader.SetupGet(x => x.RecordsAffected).Returns(1); + reader.SetupGet(x => x.HasRows).Returns(true); + reader.SetupSequence(x => x.ReadAsync(It.IsAny())) + .ReturnsAsync(true) + .ReturnsAsync(false); + reader.Setup(x => x.GetSchemaTable()).Returns(schema); + reader.Setup(x => x.GetOrdinal(MsSqlQueryBuilder.COUNT_ROWS_WITH_GIVEN_PK)).Returns(0); + reader.Setup(x => x.IsDBNull(0)).Returns(false); + reader.Setup(x => x[MsSqlQueryBuilder.COUNT_ROWS_WITH_GIVEN_PK]).Returns(0); + reader.Setup(x => x.GetFieldType(0)).Returns(typeof(int)); + reader.Setup(x => x.NextResultAsync(It.IsAny())).ReturnsAsync(false); + + DataApiBuilderException exception = await Assert.ThrowsExceptionAsync(() => + executor.GetMultipleResultSetsIfAnyAsync(reader.Object, new List { "", "Book" })); + + Assert.AreEqual(System.Net.HttpStatusCode.NotFound, exception.StatusCode); + } + + [TestMethod] + public async Task GetJsonResultAsync_StreamsWhenResponseLimitIsEnabled() + { + MsSqlQueryExecutor executor = CreateExecutor(maxResponseSizeEnabled: true); + Mock reader = new(); + reader.SetupGet(x => x.HasRows).Returns(true); + reader.SetupSequence(x => x.ReadAsync(It.IsAny())) + .ReturnsAsync(true) + .ReturnsAsync(false); + reader.Setup(x => x.GetChars(0, 0, null, 0, 0)).Returns(11); + reader.Setup(x => x.GetChars(0, 0, It.IsAny(), 0, 11)) + .Callback((int _, long _, char[]? buffer, int _, int _) => "{\"value\":7}".CopyTo(0, buffer!, 0, 11)) + .Returns(11); + + using JsonDocument? result = await executor.GetJsonResultAsync(reader.Object); + + Assert.AreEqual(7, result!.RootElement.GetProperty("value").GetInt32()); + } + + private static Mock CreateSingleRowReader(bool valueIsNull) + { + DataTable schema = new(); + schema.Columns.Add("ColumnName", typeof(string)); + schema.Columns.Add("ColumnSize", typeof(int)); + schema.Rows.Add("value", 10); + + Mock reader = new(); + reader.SetupGet(x => x.RecordsAffected).Returns(1); + reader.SetupGet(x => x.HasRows).Returns(true); + reader.SetupSequence(x => x.Read()).Returns(true).Returns(false); + reader.SetupSequence(x => x.ReadAsync(It.IsAny())) + .ReturnsAsync(true) + .ReturnsAsync(false); + reader.Setup(x => x.GetSchemaTable()).Returns(schema); + reader.Setup(x => x.GetOrdinal("value")).Returns(0); + reader.Setup(x => x.IsDBNull(0)).Returns(valueIsNull); + reader.Setup(x => x["value"]).Returns(42); + reader.Setup(x => x.GetFieldType(0)).Returns(typeof(int)); + return reader; + } + + private static MsSqlQueryExecutor CreateExecutor(bool maxResponseSizeEnabled = false) + { + MsSqlQueryExecutor executor = (MsSqlQueryExecutor)RuntimeHelpers.GetUninitializedObject(typeof(MsSqlQueryExecutor)); + Type baseType = typeof(MsSqlQueryExecutor).BaseType!; + baseType.GetField("_maxResponseSizeMB", BindingFlags.Instance | BindingFlags.NonPublic)!.SetValue(executor, 1); + baseType.GetField("_maxResponseSizeBytes", BindingFlags.Instance | BindingFlags.NonPublic)!.SetValue(executor, 1024L); + RuntimeConfig runtimeConfig = new( + Schema: string.Empty, + DataSource: new DataSource(DatabaseType.MSSQL, string.Empty), + Entities: new RuntimeEntities(new Dictionary()), + Runtime: new RuntimeOptions( + Rest: new(), + GraphQL: new(), + Mcp: null, + Host: new(Cors: null, Authentication: null, MaxResponseSizeMB: maxResponseSizeEnabled ? 1 : null))); + RuntimeConfigProvider configProvider = TestHelper.GenerateInMemoryRuntimeConfigProvider(runtimeConfig); + baseType.GetField("k__BackingField", BindingFlags.Instance | BindingFlags.NonPublic)! + .SetValue(executor, configProvider); + baseType.GetField("k__BackingField", BindingFlags.Instance | BindingFlags.NonPublic)! + .SetValue(executor, NullLogger.Instance); + baseType.GetField("k__BackingField", BindingFlags.Instance | BindingFlags.NonPublic)! + .SetValue(executor, new TestDbExceptionParser(configProvider)); + return executor; + } + + private static QueryExecutor CreateBaseExecutor() + { + RuntimeConfig runtimeConfig = new( + Schema: string.Empty, + DataSource: new DataSource(DatabaseType.MSSQL, string.Empty), + Entities: new RuntimeEntities(new Dictionary())); + RuntimeConfigProvider configProvider = TestHelper.GenerateInMemoryRuntimeConfigProvider(runtimeConfig); + return new QueryExecutor( + new TestDbExceptionParser(configProvider), + NullLogger.Instance, + configProvider, + new HttpContextAccessor(), + handler: null); + } + + private sealed class TestDbException : DbException + { + public TestDbException() + { + } + + public TestDbException(string message) + : base(message) + { + } + + public TestDbException(string message, Exception innerException) + : base(message, innerException) + { + } + } + + private sealed class TestDbExceptionParser : DbExceptionParser + { + public TestDbExceptionParser(RuntimeConfigProvider configProvider) + : base(configProvider) + { + } + + public override bool IsTransientException(DbException e) => false; + + public override HttpStatusCode GetHttpStatusCodeForException(DbException e) => HttpStatusCode.InternalServerError; + } + } +} diff --git a/src/Service.Tests/UnitTests/RuntimeConfigHelperTests.cs b/src/Service.Tests/UnitTests/RuntimeConfigHelperTests.cs new file mode 100644 index 0000000000..b720aab812 --- /dev/null +++ b/src/Service.Tests/UnitTests/RuntimeConfigHelperTests.cs @@ -0,0 +1,214 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +using System; +using System.Collections.Generic; +using System.Linq; +using Azure.DataApiBuilder.Config.ObjectModel; +using Azure.DataApiBuilder.Service.Exceptions; +using Microsoft.VisualStudio.TestTools.UnitTesting; + +namespace Azure.DataApiBuilder.Service.Tests.UnitTests +{ + [TestClass] + public class RuntimeConfigHelperTests + { + [TestMethod] + public void OptionalRuntimeProperties_UseDocumentedDefaults() + { + RuntimeConfig config = CreateConfig(runtime: null); + + Assert.IsTrue(config.IsGraphQLEnabled); + Assert.IsTrue(config.IsRestEnabled); + Assert.IsTrue(config.IsMcpEnabled); + Assert.IsTrue(config.IsHealthEnabled); + Assert.IsTrue(config.IsUnauthenticatedIdentityProvider); + Assert.IsFalse(config.IsStaticWebAppsIdentityProvider); + Assert.IsFalse(config.IsAppServiceIdentityProvider); + Assert.AreEqual(RestRuntimeOptions.DEFAULT_PATH, config.RestPath); + Assert.AreEqual(GraphQLRuntimeOptions.DEFAULT_PATH, config.GraphQLPath); + Assert.AreEqual(McpRuntimeOptions.DEFAULT_PATH, config.McpPath); + Assert.IsTrue(config.AllowIntrospection); + Assert.IsTrue(config.EnableAggregation); + Assert.IsFalse(config.EnableDwNto1JoinOpt); + Assert.AreEqual(0, config.AllowedRolesForHealth.Count); + Assert.AreEqual(EntityCacheOptions.DEFAULT_TTL_SECONDS, config.CacheTtlSecondsForHealthReport); + } + + [DataTestMethod] + [DataRow("StaticWebApps", true, false, false)] + [DataRow("AppService", false, true, false)] + [DataRow("Unauthenticated", false, false, true)] + [DataRow("Custom", false, false, false)] + public void IdentityProviderProperties_AreCaseInsensitive( + string provider, + bool staticWebApps, + bool appService, + bool unauthenticated) + { + RuntimeOptions runtime = new( + Rest: new RestRuntimeOptions(Enabled: true, Path: "/rest"), + GraphQL: new GraphQLRuntimeOptions(Enabled: true, Path: "/gql", AllowIntrospection: false), + Mcp: new McpRuntimeOptions(Enabled: true, Path: "/tools"), + Host: new HostOptions(null, new AuthenticationOptions(provider))); + RuntimeConfig config = CreateConfig(runtime); + + Assert.AreEqual(staticWebApps, config.IsStaticWebAppsIdentityProvider); + Assert.AreEqual(appService, config.IsAppServiceIdentityProvider); + Assert.AreEqual(unauthenticated, config.IsUnauthenticatedIdentityProvider); + Assert.AreEqual("/rest", config.RestPath); + Assert.AreEqual("/gql", config.GraphQLPath); + Assert.AreEqual("/tools", config.McpPath); + Assert.IsFalse(config.AllowIntrospection); + } + + [TestMethod] + public void CosmosDisablesRestEvenWhenRuntimeEnablesIt() + { + RuntimeConfig config = new( + Schema: string.Empty, + DataSource: new DataSource(DatabaseType.CosmosDB_NoSQL, string.Empty), + Entities: new RuntimeEntities(new Dictionary())); + + Assert.IsFalse(config.IsRestEnabled); + } + + [TestMethod] + public void DataSourceAndEntityMaps_SupportLookupUpdateAndPathOperations() + { + Dictionary entities = new() { ["Book"] = CreateEntity("books") }; + DataSource original = new(DatabaseType.MSSQL, "old"); + RuntimeConfig config = new( + Schema: string.Empty, + DataSource: original, + Entities: new RuntimeEntities(entities)); + string defaultName = config.DefaultDataSourceName; + + Assert.AreSame(original, config.GetDataSourceFromDataSourceName(defaultName)); + Assert.AreSame(original, config.GetDataSourceFromEntityName("Book")); + Assert.AreEqual(defaultName, config.GetDataSourceNameFromEntityName("Book")); + Assert.IsTrue(config.CheckDataSourceExists(defaultName)); + Assert.AreEqual(1, config.ListAllDataSources().Count()); + Assert.AreEqual(defaultName, config.GetDataSourceNamesToDataSourcesIterator().Single().Key); + + DataSource replacement = new(DatabaseType.PostgreSQL, "new"); + config.UpdateDataSourceNameToDataSource(defaultName, replacement); + Assert.AreSame(replacement, config.GetDataSourceFromDataSourceName(defaultName)); + + Assert.IsTrue(config.TryAddEntityPathNameToEntityName("books", "Book")); + Assert.IsFalse(config.TryAddEntityPathNameToEntityName("books", "Other")); + Assert.IsTrue(config.TryGetEntityNameFromPath("books", out string? entityName)); + Assert.AreEqual("Book", entityName); + Assert.IsFalse(config.TryGetEntityNameFromPath("missing", out _)); + + Assert.IsTrue(config.TryAddEntityNameToDataSourceName("Author")); + Assert.IsFalse(config.TryAddEntityNameToDataSourceName("Author")); + Assert.IsTrue(config.RemoveGeneratedAutoentityNameFromDataSourceName("Author")); + Assert.IsFalse(config.RemoveGeneratedAutoentityNameFromDataSourceName("Author")); + } + + [TestMethod] + public void MappingLookups_RejectUnknownNames() + { + RuntimeConfig config = CreateConfig(runtime: null); + + Assert.ThrowsException(() => config.GetDataSourceFromDataSourceName("missing")); + Assert.ThrowsException(() => config.UpdateDataSourceNameToDataSource("missing", new DataSource(DatabaseType.MSSQL, string.Empty))); + Assert.ThrowsException(() => config.GetDataSourceNameFromEntityName("missing")); + Assert.ThrowsException(() => config.GetDataSourceFromEntityName("missing")); + Assert.ThrowsException(() => config.GetDataSourceNameFromAutoentityName("missing")); + Assert.IsFalse(config.TryAddGeneratedAutoentityNameToDataSourceName("Generated", "missing")); + } + + [TestMethod] + public void UpdateDefaultDataSourceName_RekeysDataSourceAndEntities() + { + RuntimeConfig config = new( + Schema: string.Empty, + DataSource: new DataSource(DatabaseType.MSSQL, "connection"), + Entities: new RuntimeEntities(new Dictionary { ["Book"] = CreateEntity("books") })); + + config.UpdateDefaultDataSourceName("stable-name"); + + Assert.AreEqual("stable-name", config.DefaultDataSourceName); + Assert.AreEqual("stable-name", config.GetDataSourceNameFromEntityName("Book")); + Assert.IsTrue(config.CheckDataSourceExists("stable-name")); + } + + [TestMethod] + public void UpdateDefaultDataSourceName_RejectsDuplicateName() + { + DataSource original = new(DatabaseType.MSSQL, string.Empty); + RuntimeConfig config = new( + Schema: string.Empty, + DataSource: original, + Runtime: null!, + Entities: new RuntimeEntities(new Dictionary()), + DefaultDataSourceName: "original", + DataSourceNameToDataSource: new Dictionary + { + ["original"] = original, + ["duplicate"] = new DataSource(DatabaseType.MySQL, string.Empty) + }, + EntityNameToDataSourceName: new Dictionary()); + + Assert.ThrowsException(() => config.UpdateDefaultDataSourceName("duplicate")); + } + + [TestMethod] + public void RuntimeUtilityDefaults_AreStable() + { + RuntimeConfig config = CreateConfig(runtime: null); + + Assert.IsFalse(config.IsDevelopmentMode()); + Assert.IsFalse(RuntimeConfig.IsHotReloadable()); + Assert.IsFalse(config.IsMultipleCreateOperationEnabled()); + Assert.AreEqual(PaginationOptions.DEFAULT_PAGE_SIZE, config.DefaultPageSize()); + Assert.AreEqual(PaginationOptions.MAX_PAGE_SIZE, config.MaxPageSize()); + Assert.IsFalse(config.NextLinkRelative()); + Assert.AreEqual(HostOptions.MAX_RESPONSE_LENGTH_DAB_ENGINE_MB, config.MaxResponseSizeMB()); + Assert.IsFalse(config.MaxResponseSizeLogicEnabled()); + Assert.IsTrue(config.IsLogLevelNull()); + Assert.IsFalse(config.HasExplicitLogLevel()); + Assert.IsFalse(string.IsNullOrWhiteSpace(config.ToJson())); + } + + [TestMethod] + public void ExplicitMappingsConstructor_TracksSqlAndCosmosUsage() + { + DataSource sql = new(DatabaseType.MSSQL, string.Empty); + DataSource cosmos = new(DatabaseType.CosmosDB_NoSQL, string.Empty); + RuntimeConfig config = new( + Schema: string.Empty, + DataSource: sql, + Runtime: null!, + Entities: new RuntimeEntities(new Dictionary()), + DefaultDataSourceName: "sql", + DataSourceNameToDataSource: new Dictionary + { + ["sql"] = sql, + ["cosmos"] = cosmos + }, + EntityNameToDataSourceName: new Dictionary()); + + Assert.IsTrue(config.SqlDataSourceUsed); + Assert.IsTrue(config.CosmosDataSourceUsed); + } + + private static RuntimeConfig CreateConfig(RuntimeOptions? runtime) => new( + Schema: string.Empty, + DataSource: new DataSource(DatabaseType.MSSQL, string.Empty), + Entities: new RuntimeEntities(new Dictionary()), + Runtime: runtime); + + private static Entity CreateEntity(string source) => new( + Source: new EntitySource(source, EntitySourceType.Table, null, null), + GraphQL: null, + Fields: null, + Rest: null, + Permissions: Array.Empty(), + Mappings: null, + Relationships: null, + Mcp: null); + } +} diff --git a/src/Service.Tests/UnitTests/RuntimeConfigValidatorUnitTests.cs b/src/Service.Tests/UnitTests/RuntimeConfigValidatorUnitTests.cs index 1302c02695..5dce43b88a 100644 --- a/src/Service.Tests/UnitTests/RuntimeConfigValidatorUnitTests.cs +++ b/src/Service.Tests/UnitTests/RuntimeConfigValidatorUnitTests.cs @@ -253,6 +253,28 @@ public void ValidateFileSinkPath_Disabled_DoesNotThrow() validator.ValidateFileSinkPath(ConfigWith(telemetry: telemetry)); } + [TestMethod] + public void ValidateFileSinkPath_PathLongerThanRecommendation_DoesNotThrow() + { + RuntimeConfigValidator validator = CreateValidator(); + string path = $"logs/{new string('a', 256)}.txt"; + TelemetryOptions telemetry = new(File: new FileSinkOptions { Enabled = true, Path = path }); + + validator.ValidateFileSinkPath(ConfigWith(telemetry: telemetry)); + } + + [DataTestMethod] + [DataRow("logs/")] + [DataRow("logs/bad\0name.txt")] + public void ValidateFileSinkPath_InvalidFileName_Throws(string path) + { + RuntimeConfigValidator validator = CreateValidator(); + TelemetryOptions telemetry = new(File: new FileSinkOptions { Enabled = true, Path = path }); + + Assert.ThrowsException( + () => validator.ValidateFileSinkPath(ConfigWith(telemetry: telemetry))); + } + #endregion #region ValidateEmbeddingsOptions @@ -322,6 +344,201 @@ public void ValidateEmbeddings_ValidOpenAI_DoesNotThrow() validator.ValidateEmbeddingsOptions(ConfigWith(embeddings: embeddings)); } + [DataTestMethod] + [DataRow(0, null)] + [DataRow(null, 0)] + public void ValidateEmbeddings_NonPositiveNumericOptions_Throw(int? timeoutMs, int? dimensions) + { + RuntimeConfigValidator validator = CreateValidator(); + EmbeddingsOptions embeddings = new( + Provider: EmbeddingProviderType.OpenAI, + BaseUrl: "https://api.openai.com", + ApiKey: "key", + TimeoutMs: timeoutMs, + Dimensions: dimensions); + + Assert.ThrowsException( + () => validator.ValidateEmbeddingsOptions(ConfigWith(embeddings: embeddings))); + } + + [DataTestMethod] + [DataRow(true, 0, "health check", null)] + [DataRow(true, 1000, "", null)] + [DataRow(true, 1000, "health check", 0)] + public void ValidateEmbeddings_InvalidEnabledHealthOptions_Throw( + bool enabled, + int thresholdMs, + string testText, + int? expectedDimensions) + { + RuntimeConfigValidator validator = CreateValidator(); + EmbeddingsOptions embeddings = new( + Provider: EmbeddingProviderType.OpenAI, + BaseUrl: "https://api.openai.com", + ApiKey: "key", + Health: new EmbeddingsHealthCheckConfig(enabled, thresholdMs, testText, expectedDimensions)); + + Assert.ThrowsException( + () => validator.ValidateEmbeddingsOptions(ConfigWith(embeddings: embeddings))); + } + + [TestMethod] + public void ValidateEmbeddings_EnabledEndpointWithEmptyRoles_Throws() + { + RuntimeConfigValidator validator = CreateValidator(); + EmbeddingsOptions embeddings = new( + Provider: EmbeddingProviderType.OpenAI, + BaseUrl: "https://api.openai.com", + ApiKey: "key", + Endpoint: new EmbeddingsEndpointOptions(enabled: true, roles: System.Array.Empty())); + + Assert.ThrowsException( + () => validator.ValidateEmbeddingsOptions(ConfigWith(embeddings: embeddings))); + } + + [TestMethod] + public void ValidateEmbeddings_ProductionEndpointWithoutRoles_Throws() + { + RuntimeConfigValidator validator = CreateValidator(); + EmbeddingsOptions embeddings = new( + Provider: EmbeddingProviderType.OpenAI, + BaseUrl: "https://api.openai.com", + ApiKey: "key", + Endpoint: new EmbeddingsEndpointOptions(enabled: true)); + + Assert.ThrowsException( + () => validator.ValidateEmbeddingsOptions(ConfigWith( + embeddings: embeddings, + hostMode: HostMode.Production))); + } + + [DataTestMethod] + [DataRow(true, 0, false, null)] + [DataRow(true, 24, true, "")] + public void ValidateEmbeddings_InvalidCacheOptions_Throw( + bool enabled, + int ttlHours, + bool level2Enabled, + string? connectionString) + { + RuntimeConfigValidator validator = CreateValidator(); + EmbeddingsCacheOptions cache = new( + Enabled: enabled, + TtlHours: ttlHours, + Level2: new EmbeddingsCacheLevel2Options(level2Enabled, connectionString)); + EmbeddingsOptions embeddings = new( + Provider: EmbeddingProviderType.OpenAI, + BaseUrl: "https://api.openai.com", + ApiKey: "key", + Cache: cache); + + Assert.ThrowsException( + () => validator.ValidateEmbeddingsOptions(ConfigWith(embeddings: embeddings))); + } + + [TestMethod] + public void ValidateEmbeddings_ValidCacheOptions_DoNotThrow() + { + RuntimeConfigValidator validator = CreateValidator(); + EmbeddingsCacheOptions cache = new( + Enabled: true, + TtlHours: 12, + Level2: new EmbeddingsCacheLevel2Options(true, "localhost:6379")); + EmbeddingsOptions embeddings = new( + Provider: EmbeddingProviderType.OpenAI, + BaseUrl: "https://api.openai.com", + ApiKey: "key", + Cache: cache); + + validator.ValidateEmbeddingsOptions(ConfigWith(embeddings: embeddings)); + } + + #endregion + + #region Stored Procedure and MCP Validation + + [TestMethod] + public void ValidateStoredProcedureDuplicateParameters_DuplicateThrows() + { + RuntimeConfigValidator validator = CreateValidator(); + Entity entity = CreateStoredProcedureEntity(new() + { + new ParameterMetadata { Name = "id" }, + new ParameterMetadata { Name = "id" } + }); + + Assert.ThrowsException(() => + validator.ValidateStoredProcedureDuplicateParameters(ConfigWith( + entities: new() { ["Procedure"] = entity }))); + } + + [TestMethod] + public void ValidateStoredProcedureDuplicateParameters_SkipsTablesAndNullParameters() + { + RuntimeConfigValidator validator = CreateValidator(); + Entity procedure = CreateStoredProcedureEntity(parameters: null); + Entity table = CreateStoredProcedureEntity(new() + { + new ParameterMetadata { Name = "id" }, + new ParameterMetadata { Name = "id" } + }) with + { + Source = new("books", EntitySourceType.Table, null, null) + }; + + validator.ValidateStoredProcedureDuplicateParameters(ConfigWith( + entities: new() { ["Procedure"] = procedure, ["Book"] = table })); + } + + [DataTestMethod] + [DataRow(0)] + [DataRow(601)] + public void ValidateMcpUri_OutOfRangeAggregateTimeout_Throws(int timeout) + { + RuntimeConfigValidator validator = CreateValidator(); + McpRuntimeOptions mcp = new( + Enabled: true, + Path: "/mcp", + DmlTools: new DmlToolsConfig(aggregateRecordsQueryTimeout: timeout)); + + Assert.ThrowsException(() => validator.ValidateMcpUri(ConfigWith(mcp: mcp))); + } + + [TestMethod] + public void ValidateRelationshipConfigCorrectness_StoredProcedureRelationshipThrows() + { + RuntimeConfigValidator validator = CreateValidator(); + Entity entity = CreateStoredProcedureEntity(parameters: null) with + { + Relationships = new Dictionary + { + ["author"] = CreateRelationship("Author") + } + }; + + Assert.ThrowsException(() => + validator.ValidateRelationshipConfigCorrectness(ConfigWith( + entities: new() { ["Procedure"] = entity }))); + } + + [TestMethod] + public void ValidateRelationshipConfigCorrectness_UndefinedTargetThrows() + { + RuntimeConfigValidator validator = CreateValidator(); + Entity entity = CreateStoredProcedureEntity(parameters: null) with + { + Source = new("books", EntitySourceType.Table, null, null), + Relationships = new Dictionary + { + ["author"] = CreateRelationship("Missing") + } + }; + + Assert.ThrowsException(() => + validator.ValidateRelationshipConfigCorrectness(ConfigWith( + entities: new() { ["Book"] = entity }))); + } + #endregion #region ValidateGlobalEndpointRouteConfig @@ -400,7 +617,9 @@ private static RuntimeConfig ConfigWith( McpRuntimeOptions? mcp = null, TelemetryOptions? telemetry = null, EmbeddingsOptions? embeddings = null, - string? baseRoute = null) + string? baseRoute = null, + HostMode hostMode = HostMode.Development, + Dictionary? entities = null) { return new RuntimeConfig( Schema: "test-schema", @@ -409,13 +628,41 @@ private static RuntimeConfig ConfigWith( Rest: rest ?? new RestRuntimeOptions(), GraphQL: graphQL ?? new GraphQLRuntimeOptions(), Mcp: mcp, - Host: new(Cors: null, Authentication: null, Mode: HostMode.Development), + Host: new(Cors: null, Authentication: null, Mode: hostMode), BaseRoute: baseRoute, Telemetry: telemetry, Embeddings: embeddings), - Entities: new(new Dictionary())); + Entities: new(entities ?? new Dictionary())); } + private static Entity CreateStoredProcedureEntity(List? parameters) + { + return new Entity( + Source: new("procedure", EntitySourceType.StoredProcedure, parameters, null), + GraphQL: new("Procedure", "Procedures"), + Fields: null, + Rest: new(Enabled: true), + Permissions: new[] + { + new EntityPermission("anonymous", new[] + { + new EntityAction(EntityActionOperation.Execute, null, null) + }) + }, + Mappings: null, + Relationships: null); + } + + private static EntityRelationship CreateRelationship(string targetEntity) => + new( + Cardinality.One, + targetEntity, + System.Array.Empty(), + System.Array.Empty(), + LinkingObject: null, + System.Array.Empty(), + System.Array.Empty()); + #endregion } } diff --git a/src/Service.Tests/UnitTests/RuntimeOptionsConverterCoverageTests.cs b/src/Service.Tests/UnitTests/RuntimeOptionsConverterCoverageTests.cs new file mode 100644 index 0000000000..129851b879 --- /dev/null +++ b/src/Service.Tests/UnitTests/RuntimeOptionsConverterCoverageTests.cs @@ -0,0 +1,137 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +using System.Text.Json; +using Azure.DataApiBuilder.Config; +using Azure.DataApiBuilder.Config.ObjectModel; +using Microsoft.VisualStudio.TestTools.UnitTesting; + +namespace Azure.DataApiBuilder.Service.Tests.UnitTests +{ + [TestClass] + public class RuntimeOptionsConverterCoverageTests + { + private static JsonSerializerOptions Options => RuntimeConfigLoader.GetSerializationOptions(); + + [TestMethod] + public void GraphQLRuntimeOptions_FullObjectRoundTrips() + { + const string Json = """ + { + "enabled": true, + "allow-introspection": false, + "enable-aggregation": true, + "path": "/gql", + "depth-limit": 7, + "multiple-mutations": { "create": { "enabled": true } } + } + """; + + GraphQLRuntimeOptions? value = JsonSerializer.Deserialize(Json, Options); + string serialized = JsonSerializer.Serialize(value, Options); + + Assert.IsNotNull(value); + Assert.IsTrue(value.Enabled); + Assert.IsFalse(value.AllowIntrospection); + Assert.IsTrue(value.EnableAggregation); + Assert.AreEqual("/gql", value.Path); + Assert.AreEqual(7, value.DepthLimit); + Assert.IsTrue(value.MultipleMutationOptions?.MultipleCreateOptions?.Enabled); + StringAssert.Contains(serialized, "\"multiple-mutations\""); + } + + [DataTestMethod] + [DataRow("true", true)] + [DataRow("false", false)] + public void GraphQLRuntimeOptions_BooleanShorthandDeserializes(string json, bool enabled) + { + GraphQLRuntimeOptions? value = JsonSerializer.Deserialize(json, Options); + + Assert.IsNotNull(value); + Assert.AreEqual(enabled, value.Enabled); + } + + [TestMethod] + public void GraphQLRuntimeOptions_NullDepthRoundTrips() + { + GraphQLRuntimeOptions? value = JsonSerializer.Deserialize("{\"depth-limit\":null}", Options); + string serialized = JsonSerializer.Serialize(value, Options); + + Assert.IsNotNull(value); + Assert.IsTrue(value.UserProvidedDepthLimit); + Assert.IsNull(value.DepthLimit); + StringAssert.Contains(serialized, "\"depth-limit\": null"); + } + + [DataTestMethod] + [DataRow("true", true)] + [DataRow("false", false)] + public void RestRuntimeOptions_BooleanShorthandDeserializes(string json, bool expectedEnabled) + { + RestRuntimeOptions? value = JsonSerializer.Deserialize(json, Options); + + Assert.IsNotNull(value); + Assert.AreEqual(expectedEnabled, value.Enabled); + } + + [DataTestMethod] + [DataRow("{\"multiple-mutations\":{\"unknown\":true}}")] + [DataRow("{\"multiple-mutations\":42}")] + [DataRow("{\"multiple-mutations\":{\"create\":{\"unknown\":true}}}")] + [DataRow("{\"multiple-mutations\":{\"create\":42}}")] + public void GraphQLRuntimeOptions_InvalidMultipleMutationValuesThrow(string json) + { + Assert.ThrowsException(() => JsonSerializer.Deserialize(json, Options)); + } + + [DataTestMethod] + [DataRow("{\"enabled\":1}")] + [DataRow("{\"allow-introspection\":1}")] + [DataRow("{\"enable-aggregation\":1}")] + [DataRow("{\"path\":false}")] + [DataRow("{\"depth-limit\":0}")] + [DataRow("{\"depth-limit\":-2}")] + [DataRow("{\"depth-limit\":\"deep\"}")] + [DataRow("{\"unknown\":true}")] + [DataRow("42")] + public void GraphQLRuntimeOptions_InvalidValuesThrow(string json) + { + Assert.ThrowsException(() => JsonSerializer.Deserialize(json, Options)); + } + + [TestMethod] + public void FileSinkOptions_FullObjectRoundTrips() + { + const string Json = """ + { + "enabled": true, + "path": "logs/dab.txt", + "rolling-interval": "day", + "retained-file-count-limit": 5, + "file-size-limit-bytes": 1024 + } + """; + + FileSinkOptions? value = JsonSerializer.Deserialize(Json, Options); + string serialized = JsonSerializer.Serialize(value, Options); + + Assert.IsNotNull(value); + Assert.IsTrue(value.Enabled); + Assert.AreEqual("logs/dab.txt", value.Path); + Assert.AreEqual("Day", value.RollingInterval); + Assert.AreEqual(5, value.RetainedFileCountLimit); + Assert.AreEqual(1024L, value.FileSizeLimitBytes); + StringAssert.Contains(serialized, "\"file-size-limit-bytes\""); + } + + [DataTestMethod] + [DataRow("{\"retained-file-count-limit\":0}")] + [DataRow("{\"retained-file-count-limit\":-1}")] + [DataRow("{\"unknown\":true}")] + [DataRow("42")] + public void FileSinkOptions_InvalidValuesThrow(string json) + { + Assert.ThrowsException(() => JsonSerializer.Deserialize(json, Options)); + } + } +} diff --git a/src/Service.Tests/UnitTests/SqlMetadataProviderHelperTests.cs b/src/Service.Tests/UnitTests/SqlMetadataProviderHelperTests.cs new file mode 100644 index 0000000000..30bf7eb8e2 --- /dev/null +++ b/src/Service.Tests/UnitTests/SqlMetadataProviderHelperTests.cs @@ -0,0 +1,398 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +using System; +using System.Collections.Generic; +using System.Data; +using System.Reflection; +using System.Runtime.CompilerServices; +using Azure.DataApiBuilder.Config.DatabasePrimitives; +using Azure.DataApiBuilder.Config.ObjectModel; +using Azure.DataApiBuilder.Core.Configurations; +using Azure.DataApiBuilder.Core.Models; +using Azure.DataApiBuilder.Core.Resolvers; +using Azure.DataApiBuilder.Core.Resolvers.Factories; +using Azure.DataApiBuilder.Core.Services; +using Azure.DataApiBuilder.Service.Exceptions; +using HotChocolate.Language; +using Microsoft.Data.SqlClient; +using Microsoft.Extensions.Logging; +using Microsoft.VisualStudio.TestTools.UnitTesting; + +namespace Azure.DataApiBuilder.Service.Tests.UnitTests +{ + [TestClass] + public class SqlMetadataProviderHelperTests + { + [TestMethod] + public void InferredDatabaseObjectAccessors_ReturnConfiguredValues() + { + SourceDefinition sourceDefinition = new(); + MsSqlMetadataProvider provider = CreateProvider(new DatabaseTable("dbo", "books") + { + TableDefinition = sourceDefinition + }); + + Assert.AreEqual("dbo", provider.GetSchemaName("Book")); + Assert.AreEqual("books", provider.GetDatabaseObjectName("Book")); + Assert.AreSame(sourceDefinition, provider.GetSourceDefinition("Book")); + Assert.AreEqual(string.Empty, provider.GetDatabaseName()); + Assert.AreEqual(DatabaseType.MSSQL, provider.GetDatabaseType()); + } + + [TestMethod] + public void GetStoredProcedureDefinition_ReturnsConfiguredDefinition() + { + StoredProcedureDefinition definition = new(); + MsSqlMetadataProvider provider = CreateProvider(new DatabaseStoredProcedure("dbo", "get_books") + { + StoredProcedureDefinition = definition + }); + + Assert.AreSame(definition, provider.GetStoredProcedureDefinition("Book")); + } + + [DataTestMethod] + [DataRow("GetSchemaName")] + [DataRow("GetDatabaseObjectName")] + [DataRow("GetSourceDefinition")] + [DataRow("GetStoredProcedureDefinition")] + public void InferredDatabaseObjectAccessors_MissingEntity_Throw(string methodName) + { + MsSqlMetadataProvider provider = CreateProvider(); + MethodInfo method = typeof(MsSqlMetadataProvider).GetMethod(methodName)!; + + TargetInvocationException exception = Assert.ThrowsException( + () => method.Invoke(provider, new object[] { "Missing" })); + + Assert.IsInstanceOfType(exception.InnerException); + } + + [DataTestMethod] + [DataRow("Order Items", "OrderItems")] + [DataRow(" multiple spaces ", "MultipleSpaces")] + [DataRow("NoSpaces", "NoSpaces")] + [DataRow("UPPER CASE", "UPPERCASE")] + [DataRow("", "")] + public void RemoveWhitespaceAddCamelCase_TransformsDeterministically(string input, string expected) + { + MethodInfo method = typeof(MsSqlMetadataProvider).BaseType!.GetMethod( + "RemoveWhitespaceAddCamelCase", BindingFlags.Static | BindingFlags.NonPublic)!; + + Assert.AreEqual(expected, method.Invoke(null, new object[] { input })); + } + + [TestMethod] + public void FieldMappingLookups_UseCachesAndConfiguredFieldAliases() + { + Entity entity = CreateEntity(new List + { + new() { Name = "book_id", Alias = "id" } + }); + MsSqlMetadataProvider provider = CreateProvider(entity: entity); + Dictionary> backingToExposed = GetMap(provider, "EntityBackingColumnsToExposedNames"); + Dictionary> exposedToBacking = GetMap(provider, "EntityExposedNamesToBackingColumnNames"); + backingToExposed["Book"] = new() { ["title"] = "bookTitle" }; + exposedToBacking["Book"] = new() { ["bookTitle"] = "title" }; + + Assert.IsTrue(provider.TryGetExposedColumnName("Book", "title", out string? cachedExposed)); + Assert.AreEqual("bookTitle", cachedExposed); + Assert.IsTrue(provider.TryGetExposedColumnName("Book", "BOOK_ID", out string? configuredExposed)); + Assert.AreEqual("id", configuredExposed); + Assert.IsTrue(provider.TryGetBackingColumn("Book", "bookTitle", out string? cachedBacking)); + Assert.AreEqual("title", cachedBacking); + Assert.IsTrue(provider.TryGetBackingColumn("Book", "ID", out string? configuredBacking)); + Assert.AreEqual("book_id", configuredBacking); + Assert.IsFalse(provider.TryGetExposedColumnName("Book", "missing", out _)); + Assert.IsFalse(provider.TryGetBackingColumn("Book", "missing", out _)); + + Assert.IsTrue(provider.TryGetExposedFieldToBackingFieldMap("Book", out IReadOnlyDictionary? exposedMap)); + Assert.AreSame(exposedToBacking["Book"], exposedMap); + Assert.IsTrue(provider.TryGetBackingFieldToExposedFieldMap("Book", out IReadOnlyDictionary? backingMap)); + Assert.AreSame(backingToExposed["Book"], backingMap); + Assert.IsFalse(provider.TryGetExposedFieldToBackingFieldMap("Missing", out _)); + Assert.IsFalse(provider.TryGetBackingFieldToExposedFieldMap("Missing", out _)); + } + + [TestMethod] + public void FieldMappingLookups_MissingInitializationThrow() + { + MsSqlMetadataProvider provider = CreateProvider(entity: CreateEntity()); + + Assert.ThrowsException(() => provider.TryGetExposedColumnName("Book", "id", out _)); + Assert.ThrowsException(() => provider.TryGetBackingColumn("Book", "id", out _)); + } + + [TestMethod] + public void TryGetArrayElementSyntaxKind_RecognizesSupportedArrayAndRejectsScalar() + { + SourceDefinition definition = new(); + definition.Columns["vector"] = new ColumnDefinition(typeof(float[])) + { + IsArrayType = true, + ElementSystemType = typeof(float) + }; + definition.Columns["title"] = new ColumnDefinition(typeof(string)); + MsSqlMetadataProvider provider = CreateProvider( + new DatabaseTable("dbo", "books") { TableDefinition = definition }, + CreateEntity()); + Dictionary> map = GetMap(provider, "EntityExposedNamesToBackingColumnNames"); + map["Book"] = new() { ["vector"] = "vector", ["title"] = "title" }; + + Assert.IsTrue(provider.TryGetArrayElementSyntaxKind("Book", "vector", out SyntaxKind kind)); + Assert.AreEqual(SyntaxKind.FloatValue, kind); + Assert.IsFalse(provider.TryGetArrayElementSyntaxKind("Book", "title", out _)); + } + + [TestMethod] + public void GetEntityName_ResolvesEntityAndSingularNameAndRejectsUnknownType() + { + MsSqlMetadataProvider provider = CreateProvider(entity: CreateEntity(singular: "Volume")); + + Assert.AreEqual("Book", provider.GetEntityName("Book")); + Assert.AreEqual("Book", provider.GetEntityName("Volume")); + Assert.ThrowsException(() => provider.GetEntityName("Missing")); + } + + [TestMethod] + public void ParseSchemaAndDbTableName_HandlesDefaultExplicitPostgresAndMySqlCases() + { + MsSqlMetadataProvider provider = CreateProvider(); + Assert.AreEqual(("dbo", "books"), provider.ParseSchemaAndDbTableName("books")); + Assert.AreEqual(("custom", "books"), provider.ParseSchemaAndDbTableName("custom.books")); + + SetBaseField(provider, "_databaseType", DatabaseType.PostgreSQL); + SetBaseAutoProperty(provider, "ConnectionString", "Host=localhost;Database=db;SearchPath=tenant"); + Assert.AreEqual(("tenant", "books"), provider.ParseSchemaAndDbTableName("books")); + + SetBaseField(provider, "_databaseType", DatabaseType.MySQL); + Assert.ThrowsException(() => provider.ParseSchemaAndDbTableName("custom.books")); + } + + [TestMethod] + public void RelationalOnlyUnsupportedMetadataMembersThrow() + { + MsSqlMetadataProvider provider = CreateProvider(); + + Assert.ThrowsException(() => provider.GetSchemaGraphQLFieldNamesForEntityName("Book")); + Assert.ThrowsException(() => provider.GetSchemaGraphQLFieldTypeFromFieldName("Book", "id")); + Assert.ThrowsException(() => provider.GetSchemaGraphQLFieldFromFieldName("Book", "id")); + Assert.ThrowsException(() => provider.GetPartitionKeyPath("db", "container")); + Assert.ThrowsException(() => provider.SetPartitionKeyPath("db", "container", "/id")); + } + + [TestMethod] + public void VerifyForeignKeyExistsInDb_ChecksBothDirectionsAndNullMetadata() + { + MsSqlMetadataProvider provider = CreateProvider(); + DatabaseTable first = new("dbo", "books"); + DatabaseTable second = new("dbo", "authors"); + + provider.PairToFkDefinition = null; + Assert.IsFalse(provider.VerifyForeignKeyExistsInDB(first, second)); + + RelationShipPair reverse = new(second, first); + provider.PairToFkDefinition = new() { [reverse] = new ForeignKeyDefinition { Pair = reverse } }; + Assert.IsTrue(provider.VerifyForeignKeyExistsInDB(first, second)); + } + + [TestMethod] + public void TryGetFkDefinition_MissingEntitiesReturnsFalse() + { + MsSqlMetadataProvider provider = CreateProvider(); + + Assert.IsFalse(provider.TryGetFKDefinition("Source", "Target", "Source", "Target", out ForeignKeyDefinition? definition)); + Assert.IsNull(definition); + } + + [TestMethod] + public void InitializeAsync_ReplacesMapsAndGeneratesFieldMappings() + { + SourceDefinition definition = new(); + definition.Columns["title"] = new ColumnDefinition(typeof(string)); + Dictionary databaseObjects = new() + { + ["Book"] = new DatabaseTable("dbo", "books") { TableDefinition = definition } + }; + Dictionary procedures = new() { ["getBooks"] = "Book" }; + MsSqlMetadataProvider provider = CreateProvider(entity: CreateEntity()); + + provider.InitializeAsync(databaseObjects, procedures); + + Assert.AreSame(databaseObjects, provider.EntityToDatabaseObject); + Assert.AreSame(procedures, provider.GraphQLStoredProcedureExposedNameToEntityNameMap); + Assert.IsTrue(provider.TryGetExposedColumnName("Book", "title", out string? exposed)); + Assert.AreEqual("title", exposed); + } + + [TestMethod] + public void BaseVirtualMetadataOperations_UseDefaultBehavior() + { + BaseBehaviorMetadataProvider provider = + (BaseBehaviorMetadataProvider)RuntimeHelpers.GetUninitializedObject(typeof(BaseBehaviorMetadataProvider)); + + Assert.ThrowsException(() => provider.GetDefaultSchemaName()); + Assert.ThrowsException( + () => provider.PopulateTriggerMetadataForTable("Book", "dbo", "books", new SourceDefinition())); + + MethodInfo populateLinkingObject = GetBaseMethod("PopulateMetadataForLinkingObject"); + populateLinkingObject.Invoke(provider, new object[] + { + "Book", "Author", "dbo.book_authors", new Dictionary() + }); + + MethodInfo generateAutoentities = GetBaseMethod("GenerateAutoentitiesIntoEntities"); + TargetInvocationException exception = Assert.ThrowsException( + () => generateAutoentities.Invoke(provider, new object?[] { null })); + Assert.IsInstanceOfType(exception.InnerException); + } + + [TestMethod] + public void GetForeignKeyQueryParams_GeneratesSchemaAndTableParameters() + { + MsSqlMetadataProvider provider = CreateProvider(); + MethodInfo method = GetBaseMethod("GetForeignKeyQueryParams"); + + Dictionary parameters = + (Dictionary)method.Invoke(provider, new object[] + { + new[] { "dbo", "sales" }, new[] { "books", "orders" } + })!; + + Assert.AreEqual(4, parameters.Count); + CollectionAssert.AreEquivalent( + new object[] { "dbo", "sales", "books", "orders" }, + new List(System.Linq.Enumerable.Select(parameters.Values, parameter => parameter.Value))); + } + + [TestMethod] + public async System.Threading.Tasks.Task FillSchemaForStoredProcedureAsync_TranslatesOfflineFailure() + { + BaseBehaviorMetadataProvider provider = + (BaseBehaviorMetadataProvider)RuntimeHelpers.GetUninitializedObject(typeof(BaseBehaviorMetadataProvider)); + MethodInfo method = GetBaseMethod("FillSchemaForStoredProcedureAsync"); + Entity procedure = new( + Source: new EntitySource("dbo.get_books", EntitySourceType.StoredProcedure, null, null), + GraphQL: null, + Fields: null, + Rest: null, + Permissions: Array.Empty(), + Mappings: null, + Relationships: null); + + System.Threading.Tasks.Task task = (System.Threading.Tasks.Task)method.Invoke(provider, new object[] + { + procedure, "Book", "dbo", "get_books", new StoredProcedureDefinition() + })!; + + DataApiBuilderException exception = + await Assert.ThrowsExceptionAsync(() => task); + StringAssert.Contains(exception.Message, "Cannot obtain Schema for entity Book"); + } + + [TestMethod] + public void LogPrimaryKeys_RecordsEntityFailureDuringValidation() + { + MsSqlMetadataProvider provider = CreateProvider(entity: CreateEntity()); + SetBaseField(provider, "_isValidateOnly", true); + SetBaseAutoProperty(provider, "SqlMetadataExceptions", new List()); + + GetBaseMethod("LogPrimaryKeys").Invoke(provider, null); + + Assert.AreEqual(1, provider.SqlMetadataExceptions.Count); + Assert.IsInstanceOfType(provider.SqlMetadataExceptions[0]); + } + + [TestMethod] + public void GenerateRestPathToEntityMap_RecordsConflictingPathDuringValidation() + { + Entity entity = new( + Source: new EntitySource("dbo.books", EntitySourceType.Table, null, null), + GraphQL: null, + Fields: null, + Rest: new EntityRestOptions(Enabled: true, Path: "/graphql"), + Permissions: Array.Empty(), + Mappings: null, + Relationships: null); + MsSqlMetadataProvider provider = CreateProvider(entity: entity); + SetBaseField(provider, "_isValidateOnly", true); + SetBaseAutoProperty(provider, "SqlMetadataExceptions", new List()); + + GetBaseMethod("GenerateRestPathToEntityMap").Invoke(provider, null); + + Assert.AreEqual(1, provider.SqlMetadataExceptions.Count); + Assert.IsInstanceOfType(provider.SqlMetadataExceptions[0]); + } + + private static MsSqlMetadataProvider CreateProvider( + DatabaseObject? databaseObject = null, + Entity? entity = null) + { + MsSqlMetadataProvider provider = (MsSqlMetadataProvider)RuntimeHelpers.GetUninitializedObject(typeof(MsSqlMetadataProvider)); + provider.EntityToDatabaseObject = new Dictionary(StringComparer.InvariantCulture); + if (databaseObject is not null) + { + provider.EntityToDatabaseObject.Add("Book", databaseObject); + } + + SetBaseField(provider, "_databaseType", DatabaseType.MSSQL); + SetBaseField(provider, "_linkingEntities", new Dictionary()); + SetBaseAutoProperty(provider, "EntityBackingColumnsToExposedNames", new Dictionary>()); + SetBaseAutoProperty(provider, "EntityExposedNamesToBackingColumnNames", new Dictionary>()); + + Dictionary entities = entity is null + ? new Dictionary() + : new Dictionary { ["Book"] = entity }; + RuntimeConfig runtimeConfig = new( + Schema: string.Empty, + DataSource: new DataSource(DatabaseType.MSSQL, string.Empty), + Entities: new RuntimeEntities(entities)); + RuntimeConfigProvider configProvider = TestHelper.GenerateInMemoryRuntimeConfigProvider(runtimeConfig); + SetBaseField(provider, "_runtimeConfigProvider", configProvider); + SetBaseField(provider, "_dataSourceName", configProvider.GetConfig().DefaultDataSourceName); + return provider; + } + + private static Entity CreateEntity(List? fields = null, string singular = "Book") => + new( + Source: new EntitySource("dbo.books", EntitySourceType.Table, null, null), + GraphQL: new EntityGraphQLOptions(singular, "Books"), + Fields: fields, + Rest: new EntityRestOptions(Enabled: true), + Permissions: Array.Empty(), + Mappings: null, + Relationships: null); + + private static Dictionary> GetMap(MsSqlMetadataProvider provider, string propertyName) => + (Dictionary>)typeof(MsSqlMetadataProvider).BaseType! + .GetField($"<{propertyName}>k__BackingField", BindingFlags.Instance | BindingFlags.NonPublic)!.GetValue(provider)!; + + private static MethodInfo GetBaseMethod(string methodName) => + typeof(MsSqlMetadataProvider).BaseType!.GetMethod(methodName, BindingFlags.Instance | BindingFlags.NonPublic)!; + + private static void SetBaseField(MsSqlMetadataProvider provider, string fieldName, object value) => + typeof(MsSqlMetadataProvider).BaseType!.GetField( + fieldName, + BindingFlags.Instance | BindingFlags.Public | BindingFlags.NonPublic)! + .SetValue(provider, value); + + private static void SetBaseAutoProperty(MsSqlMetadataProvider provider, string propertyName, object value) => + typeof(MsSqlMetadataProvider).BaseType!.GetField($"<{propertyName}>k__BackingField", BindingFlags.Instance | BindingFlags.NonPublic)! + .SetValue(provider, value); + + private sealed class BaseBehaviorMetadataProvider : SqlMetadataProvider + { + public BaseBehaviorMetadataProvider( + RuntimeConfigProvider runtimeConfigProvider, + RuntimeConfigValidator runtimeConfigValidator, + IAbstractQueryManagerFactory engineFactory, + ILogger logger, + string dataSourceName) + : base(runtimeConfigProvider, runtimeConfigValidator, engineFactory, logger, dataSourceName) + { + } + + public override Type SqlToCLRType(string sqlType) => typeof(string); + } + } +} diff --git a/src/Service.Tests/UnitTests/SqlMutationEngineHelperTests.cs b/src/Service.Tests/UnitTests/SqlMutationEngineHelperTests.cs new file mode 100644 index 0000000000..881a4c1a36 --- /dev/null +++ b/src/Service.Tests/UnitTests/SqlMutationEngineHelperTests.cs @@ -0,0 +1,651 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +using System; +using System.Collections.Generic; +using System.Data.Common; +using System.Reflection; +using System.Runtime.CompilerServices; +using System.Text.Json; +using System.Text.Json.Nodes; +using System.Threading.Tasks; +using System.Transactions; +using Azure.DataApiBuilder.Auth; +using Azure.DataApiBuilder.Config.DatabasePrimitives; +using Azure.DataApiBuilder.Config.ObjectModel; +using Azure.DataApiBuilder.Core.Configurations; +using Azure.DataApiBuilder.Core.Models; +using Azure.DataApiBuilder.Core.Resolvers; +using Azure.DataApiBuilder.Core.Resolvers.Factories; +using Azure.DataApiBuilder.Core.Resolvers.Sql_Query_Structures; +using Azure.DataApiBuilder.Core.Services; +using Azure.DataApiBuilder.Core.Services.MetadataProviders; +using Azure.DataApiBuilder.Service.Exceptions; +using Microsoft.AspNetCore.Http; +using Microsoft.AspNetCore.Mvc; +using Microsoft.VisualStudio.TestTools.UnitTesting; +using Moq; +using HotChocolate.Language; +using HotChocolate.Resolvers; + +namespace Azure.DataApiBuilder.Service.Tests.UnitTests +{ + [TestClass] + public class SqlMutationEngineHelperTests + { + [TestMethod] + public void FetchPrimaryKeyFieldValues_ReturnsMappedNonNullKeys() + { + SourceDefinition definition = new() { PrimaryKey = new() { "book_id", "edition" } }; + Mock metadata = CreateMetadata(definition); + Dictionary values = new() { ["id"] = 7, ["edition"] = 2 }; + + Dictionary result = InvokeStatic>( + "FetchPrimaryKeyFieldValues", metadata.Object, "Book", values); + + CollectionAssert.AreEquivalent(new[] { "book_id", "edition" }, new List(result.Keys)); + Assert.AreEqual(7, result["book_id"]); + Assert.AreEqual(2, result["edition"]); + } + + [DataTestMethod] + [DataRow(false, false)] + [DataRow(true, true)] + public void FetchPrimaryKeyFieldValues_MissingMappingOrNullValue_Throws(bool mappingExists, bool nullValue) + { + SourceDefinition definition = new() { PrimaryKey = new() { "book_id" } }; + Mock metadata = new(); + metadata.Setup(x => x.GetSourceDefinition("Book")).Returns(definition); + string? exposed = mappingExists ? "id" : null; + metadata.Setup(x => x.TryGetExposedColumnName("Book", "book_id", out exposed)).Returns(mappingExists); + Dictionary values = new() { ["id"] = nullValue ? null : 7 }; + + TargetInvocationException exception = Assert.ThrowsException(() => + InvokeStatic>("FetchPrimaryKeyFieldValues", metadata.Object, "Book", values)); + + Assert.IsInstanceOfType(exception.InnerException); + } + + [TestMethod] + public void PopulateReferencingFields_NullComputedFields_DoesNothing() + { + MultipleCreateStructure structure = new("Book", "Publisher"); + ForeignKeyDefinition foreignKey = new() + { + ReferencingColumns = new() { "publisher_id" }, + ReferencedColumns = new() { "id" } + }; + + InvokeStatic("PopulateReferencingFields", new Mock().Object, + structure, foreignKey, null, false, "Publisher"); + + Assert.AreEqual(0, structure.CurrentEntityParams.Count); + Assert.AreEqual(0, structure.LinkingTableParams.Count); + } + + [TestMethod] + public void PopulateReferencingFields_LinkingTable_UsesBackingReferencedNames() + { + MultipleCreateStructure structure = new("BookAuthor", "Book", isLinkingTableInsertionRequired: true); + ForeignKeyDefinition foreignKey = new() + { + ReferencingColumns = new() { "book_id", "author_id" }, + ReferencedColumns = new() { "id", "author_key" } + }; + Dictionary values = new() { ["id"] = 7, ["author_key"] = 9 }; + + InvokeStatic("PopulateReferencingFields", new Mock().Object, + structure, foreignKey, values, true, null); + + Assert.AreEqual(7, structure.LinkingTableParams["book_id"]); + Assert.AreEqual(9, structure.LinkingTableParams["author_id"]); + } + + [DataTestMethod] + [DataRow(true)] + [DataRow(false)] + public void PopulateReferencingFields_CurrentEntity_ResolvesExposedNameWhenAvailable(bool mappingExists) + { + MultipleCreateStructure structure = new("Book", "Publisher"); + ForeignKeyDefinition foreignKey = new() + { + ReferencingColumns = new() { "publisher_id" }, + ReferencedColumns = new() { "publisher_key" } + }; + Mock metadata = new(); + string? exposedName = mappingExists ? "publisherId" : null; + metadata.Setup(x => x.TryGetExposedColumnName("Publisher", "publisher_key", out exposedName)) + .Returns(mappingExists); + string valueName = mappingExists ? "publisherId" : "publisher_key"; + Dictionary values = new() { [valueName] = 11 }; + + InvokeStatic("PopulateReferencingFields", metadata.Object, + structure, foreignKey, values, false, "Publisher"); + + Assert.AreEqual(11, structure.CurrentEntityParams["publisher_id"]); + } + + [TestMethod] + public void GetBackingColumnsFromCollection_MapsNamesAndPreservesValues() + { + Mock metadata = new(); + metadata.Setup(x => x.TryGetBackingColumn("Book", It.IsAny(), out It.Ref.IsAny)) + .Returns((string _, string exposed, out string? backing) => + { + backing = exposed == "id" ? "book_id" : null; + return backing is not null; + }); + Dictionary parameters = new() { ["id"] = 7, ["title"] = null }; + + Dictionary result = SqlMutationEngine.GetBackingColumnsFromCollection( + "Book", parameters, metadata.Object); + + Assert.AreEqual(7, result["book_id"]); + Assert.IsNull(result["title"]); + } + + [TestMethod] + public void GetBackingColumnsFromCollection_EmptyInput_ReturnsEmptyDictionary() + { + Dictionary result = SqlMutationEngine.GetBackingColumnsFromCollection( + "Book", new Dictionary(), new Mock().Object); + + Assert.AreEqual(0, result.Count); + } + + [DataTestMethod] + [DataRow(DatabaseType.MySQL, IsolationLevel.RepeatableRead)] + [DataRow(DatabaseType.MSSQL, IsolationLevel.ReadCommitted)] + [DataRow(DatabaseType.PostgreSQL, IsolationLevel.ReadCommitted)] + public void ConstructTransactionScopeBasedOnDbType_UsesExpectedIsolationLevel( + DatabaseType databaseType, + IsolationLevel expected) + { + Mock metadata = new(); + metadata.Setup(x => x.GetDatabaseType()).Returns(databaseType); + + using TransactionScope scope = InvokeStatic( + "ConstructTransactionScopeBasedOnDbType", metadata.Object); + + Assert.AreEqual(expected, Transaction.Current!.IsolationLevel); + } + + [TestMethod] + public void GetDbOperationResultJsonDocument_ReturnsResultAndEmptyMetadata() + { + Tuple result = InvokeStatic>( + "GetDbOperationResultJsonDocument", "success"); + + Assert.AreEqual("success", result.Item1!.RootElement.GetProperty("result").GetString()); + Assert.IsNotNull(result.Item2); + } + + [DataTestMethod] + [DataRow(EntityActionOperation.UpdateGraphQL, EntityActionOperation.Update, true)] + [DataRow(EntityActionOperation.Create, EntityActionOperation.Create, false)] + public void AreFieldsAuthorizedForEntity_DelegatesColumnOperations( + EntityActionOperation requested, + EntityActionOperation delegated, + bool expected) + { + Mock authorization = new(); + authorization.Setup(x => x.AreColumnsAllowedForOperation( + "Book", "role", delegated, It.IsAny>())).Returns(expected); + SqlMutationEngine engine = CreateUninitializedEngine(authorization.Object); + + bool result = InvokeInstance( + engine, "AreFieldsAuthorizedForEntity", "role", "Book", requested, new[] { "title" }); + + Assert.AreEqual(expected, result); + authorization.Verify(x => x.AreColumnsAllowedForOperation( + "Book", "role", delegated, It.IsAny>()), Times.Once); + } + + [DataTestMethod] + [DataRow(EntityActionOperation.Delete)] + [DataRow(EntityActionOperation.Execute)] + public void AreFieldsAuthorizedForEntity_OperationsWithoutColumnAuthorization_ReturnTrue(EntityActionOperation operation) + { + Mock authorization = new(); + SqlMutationEngine engine = CreateUninitializedEngine(authorization.Object); + + Assert.IsTrue(InvokeInstance( + engine, "AreFieldsAuthorizedForEntity", "role", "Book", operation, Array.Empty())); + authorization.VerifyNoOtherCalls(); + } + + [TestMethod] + public void AreFieldsAuthorizedForEntity_InvalidOperation_Throws() + { + SqlMutationEngine engine = CreateUninitializedEngine(new Mock().Object); + + TargetInvocationException exception = Assert.ThrowsException(() => + InvokeInstance(engine, "AreFieldsAuthorizedForEntity", "role", "Book", EntityActionOperation.Read, Array.Empty())); + + Assert.IsInstanceOfType(exception.InnerException); + } + + [DataTestMethod] + [DataRow(EntityActionOperation.Delete, false, typeof(NoContentResult))] + [DataRow(EntityActionOperation.Insert, true, typeof(CreatedResult))] + [DataRow(EntityActionOperation.Insert, false, typeof(CreatedResult))] + [DataRow(EntityActionOperation.Update, true, typeof(OkObjectResult))] + [DataRow(EntityActionOperation.Update, false, typeof(OkObjectResult))] + [DataRow(EntityActionOperation.UpdateIncremental, false, typeof(OkObjectResult))] + [DataRow(EntityActionOperation.Upsert, false, typeof(OkObjectResult))] + [DataRow(EntityActionOperation.UpsertIncremental, false, typeof(OkObjectResult))] + public async Task ExecuteStoredProcedure_ReturnsResponseForEachSupportedOperation( + EntityActionOperation operation, + bool hasRows, + Type expectedResultType) + { + JsonArray result = hasRows ? new JsonArray(new JsonObject { ["id"] = 7 }) : new JsonArray(); + (SqlMutationEngine engine, StoredProcedureRequestContext context, string dataSourceName) = + CreateStoredProcedureFixture(operation, result); + + IActionResult? response = await engine.ExecuteAsync(context, dataSourceName); + + Assert.IsInstanceOfType(response, expectedResultType); + if (response is CreatedResult created) + { + Assert.AreEqual("https://example.test/api/procedure", created.Location); + } + } + + [TestMethod] + public async Task ExecuteStoredProcedure_RejectsUnsupportedOperationAfterExecution() + { + (SqlMutationEngine engine, StoredProcedureRequestContext context, string dataSourceName) = + CreateStoredProcedureFixture(EntityActionOperation.Create, new JsonArray()); + + await Assert.ThrowsExceptionAsync(() => engine.ExecuteAsync(context, dataSourceName)); + } + + [TestMethod] + public void PopulateCurrentAndLinkingEntityParams_NullInput_DoesNothing() + { + MultipleCreateStructure structure = new("Book", string.Empty); + + InvokeStatic("PopulateCurrentAndLinkingEntityParams", structure, + new Mock().Object, null); + + Assert.AreEqual(0, structure.CurrentEntityParams.Count); + Assert.AreEqual(0, structure.LinkingTableParams.Count); + } + + [TestMethod] + public void PopulateCurrentAndLinkingEntityParams_PartitionsColumnsRelationshipsAndLinkingFields() + { + MultipleCreateStructure structure = new("Book", string.Empty, new Dictionary + { + ["title"] = "DAB", + ["royalty"] = 10, + ["authors"] = new object() + }); + Mock metadata = new(); + metadata.Setup(x => x.TryGetBackingColumn("Book", "title", out It.Ref.IsAny)).Returns(true); + Dictionary relationships = new() + { + ["authors"] = new(Cardinality.Many, "Author", null, null, "book_author", null, null) + }; + + InvokeStatic("PopulateCurrentAndLinkingEntityParams", structure, metadata.Object, relationships); + + Assert.AreEqual("DAB", structure.CurrentEntityParams["title"]); + Assert.AreEqual(10, structure.LinkingTableParams["royalty"]); + Assert.IsFalse(structure.CurrentEntityParams.ContainsKey("authors")); + Assert.IsFalse(structure.LinkingTableParams.ContainsKey("authors")); + } + + [TestMethod] + public void DetermineRelationships_NullMetadata_DoesNothing() + { + MultipleCreateStructure structure = new("Book", string.Empty, new Dictionary()); + + InvokeStatic("DetermineReferencedAndReferencingRelationships", + new Mock().Object, + structure, + new Mock().Object, + null, + new List()); + + Assert.AreEqual(0, structure.ReferencedRelationships.Count); + Assert.AreEqual(0, structure.ReferencingRelationships.Count); + } + + [TestMethod] + public void DetermineRelationships_NullInput_Throws() + { + MultipleCreateStructure structure = new("Book", string.Empty); + + TargetInvocationException exception = Assert.ThrowsException(() => + InvokeStatic("DetermineReferencedAndReferencingRelationships", + new Mock().Object, + structure, + new Mock().Object, + new Dictionary(), + new List())); + + Assert.IsInstanceOfType(exception.InnerException); + } + + [TestMethod] + public void DetermineRelationships_ManyToManyIsReferencingAndUnknownFieldIsIgnored() + { + object relationshipValue = new(); + MultipleCreateStructure structure = new("Book", string.Empty, new Dictionary + { + ["authors"] = relationshipValue, + ["title"] = "DAB" + }); + Dictionary relationships = new() + { + ["authors"] = new(Cardinality.Many, "Author", null, null, "book_author", null, null) + }; + + InvokeStatic("DetermineReferencedAndReferencingRelationships", + new Mock().Object, + structure, + new Mock().Object, + relationships, + new List()); + + Assert.AreEqual(1, structure.ReferencingRelationships.Count); + Assert.AreEqual("authors", structure.ReferencingRelationships[0].Item1); + Assert.AreSame(relationshipValue, structure.ReferencingRelationships[0].Item2); + Assert.AreEqual(0, structure.ReferencedRelationships.Count); + } + + [TestMethod] + public void MultipleCreateArgumentParsing_MissingRootFieldThrowsBadRequest() + { + RuntimeConfig runtimeConfig = new( + Schema: string.Empty, + DataSource: new DataSource(DatabaseType.MSSQL, string.Empty), + Entities: new RuntimeEntities(new Dictionary())); + + Assert.ThrowsException(() => + SqlMutationEngine.GQLMultipleCreateArgumentToDictParams( + Mock.Of(), + "item", + new Dictionary(), + Mock.Of(), + "Book", + runtimeConfig)); + } + + [TestMethod] + public void MultipleCreateArgumentParsing_UnsupportedInputTypeThrowsBadRequest() + { + RuntimeConfig runtimeConfig = new( + Schema: string.Empty, + DataSource: new DataSource(DatabaseType.MSSQL, string.Empty), + Entities: new RuntimeEntities(new Dictionary())); + + Assert.ThrowsException(() => + SqlMutationEngine.GQLMultipleCreateArgumentToDictParamsHelper( + Mock.Of(), + null!, + new object(), + Mock.Of(), + "Book", + runtimeConfig)); + } + + [TestMethod] + public void ProcessMultipleCreateInputField_NullInputThrowsBadRequest() + { + SqlMutationEngine engine = CreateUninitializedEngine(Mock.Of()); + MultipleCreateStructure structure = new("Book", string.Empty); + + TargetInvocationException exception = Assert.ThrowsException(() => + InvokeInstance(engine, "ProcessMultipleCreateInputField", + Mock.Of(), null, Mock.Of(), structure, 0)); + + Assert.IsInstanceOfType(exception.InnerException); + } + + [TestMethod] + public void ProcessMultipleCreateInputField_NonNodeObjectThrowsBadRequest() + { + SqlMutationEngine engine = CreateUninitializedEngine(Mock.Of()); + MultipleCreateStructure structure = new( + "Book", + string.Empty, + new Dictionary { ["title"] = "DAB" }); + + TargetInvocationException exception = Assert.ThrowsException(() => + InvokeInstance(engine, "ProcessMultipleCreateInputField", + Mock.Of(), new object(), Mock.Of(), structure, 0)); + + Assert.IsInstanceOfType(exception.InnerException); + } + + [TestMethod] + public void ProcessMultipleCreateInputField_NullListNodeThrowsBadRequest() + { + SqlMutationEngine engine = CreateUninitializedEngine(Mock.Of()); + MultipleCreateStructure structure = new( + "Book", + string.Empty, + new List> { new Dictionary() }); + + TargetInvocationException exception = Assert.ThrowsException(() => + InvokeInstance(engine, "ProcessMultipleCreateInputField", + Mock.Of(), new List { null! }, Mock.Of(), structure, 0)); + + Assert.IsInstanceOfType(exception.InnerException); + } + + [DataTestMethod] + [DataRow(EntityActionOperation.UpdateGraphQL)] + [DataRow(EntityActionOperation.Delete)] + public async Task PerformMutationOperation_RejectsInvalidInvocation(EntityActionOperation operation) + { + (SqlMutationEngine engine, ISqlMetadataProvider metadata) = CreateMutationOperationFixture(); + MethodInfo method = typeof(SqlMutationEngine).GetMethod("PerformMutationOperation", BindingFlags.Instance | BindingFlags.NonPublic)!; + + Task task = (Task)method.Invoke(engine, new object?[] + { + "Book", + operation, + new Dictionary(), + metadata, + null + })!; + + if (operation is EntityActionOperation.UpdateGraphQL) + { + await Assert.ThrowsExceptionAsync(async () => await task); + } + else + { + await Assert.ThrowsExceptionAsync(async () => await task); + } + } + + [DataTestMethod] + [DataRow(true, false)] + [DataRow(false, true)] + public void BuildAndExecuteInsertDbQueries_ReportsEmptyResults(bool linkingEntity, bool returnNull) + { + (SqlMutationEngine engine, ISqlMetadataProvider metadata) = + CreateMutationOperationFixture(returnNull ? null : new DbResultSet(new Dictionary())); + SourceDefinition definition = new(); + + TargetInvocationException exception = Assert.ThrowsException(() => + InvokeInstance>(engine, "BuildAndExecuteInsertDbQueries", + metadata, + "Book", + "Parent", + new Dictionary(), + definition, + linkingEntity, + 1)); + + DataApiBuilderException error = (DataApiBuilderException)exception.InnerException!; + Assert.AreEqual( + linkingEntity ? System.Net.HttpStatusCode.InternalServerError : System.Net.HttpStatusCode.Forbidden, + error.StatusCode); + } + + private static Mock CreateMetadata(SourceDefinition definition) + { + Mock metadata = new(); + metadata.Setup(x => x.GetSourceDefinition("Book")).Returns(definition); + metadata.Setup(x => x.TryGetExposedColumnName("Book", "book_id", out It.Ref.IsAny)) + .Returns((string _, string backing, out string? exposed) => + { + exposed = backing == "book_id" ? "id" : backing; + return true; + }); + metadata.Setup(x => x.TryGetExposedColumnName("Book", "edition", out It.Ref.IsAny)) + .Returns((string _, string backing, out string? exposed) => + { + exposed = backing; + return true; + }); + return metadata; + } + + private static (SqlMutationEngine Engine, ISqlMetadataProvider Metadata) CreateMutationOperationFixture( + DbResultSet? syncResult = null) + { + Entity entity = new( + Source: new EntitySource("dbo.Books", EntitySourceType.Table, null, null), + GraphQL: null, + Fields: null, + Rest: null, + Permissions: Array.Empty(), + Mappings: null, + Relationships: null, + Mcp: null); + RuntimeConfig runtimeConfig = new( + Schema: string.Empty, + DataSource: new DataSource(DatabaseType.MSSQL, string.Empty), + Entities: new RuntimeEntities(new Dictionary { ["Book"] = entity, ["Parent"] = entity })); + RuntimeConfigProvider runtimeConfigProvider = TestHelper.GenerateInMemoryRuntimeConfigProvider(runtimeConfig); + + SourceDefinition definition = new(); + DatabaseTable table = new("dbo", "Books") { TableDefinition = definition }; + Mock metadata = new(); + metadata.Setup(x => x.GetDatabaseType()).Returns(DatabaseType.MSSQL); + metadata.Setup(x => x.GetSourceDefinition("Book")).Returns(definition); + metadata.SetupGet(x => x.EntityToDatabaseObject) + .Returns(new Dictionary { ["Book"] = table }); + + Mock queryBuilder = new(); + queryBuilder.Setup(x => x.Build(It.IsAny())).Returns("INSERT"); + Mock queryExecutor = new(); + queryExecutor.Setup(x => x.ExecuteQuery( + It.IsAny(), + It.IsAny>(), + It.IsAny?, DbResultSet>>(), + It.IsAny(), + It.IsAny?>(), + It.IsAny())) + .Returns(syncResult); + Mock queryManagerFactory = new(); + queryManagerFactory.Setup(x => x.GetQueryBuilder(DatabaseType.MSSQL)).Returns(queryBuilder.Object); + queryManagerFactory.Setup(x => x.GetQueryExecutor(DatabaseType.MSSQL)).Returns(queryExecutor.Object); + + Mock accessor = new(); + accessor.Setup(x => x.HttpContext).Returns(new DefaultHttpContext()); + Mock authorization = new(); + authorization.Setup(x => x.ResolveDBPolicy( + It.IsAny(), + It.IsAny(), + It.IsAny(), + It.IsAny())) + .Returns(ResolvedDatabasePolicy.Empty); + + SqlMutationEngine engine = new( + queryManagerFactory.Object, + Mock.Of(), + Mock.Of(), + authorization.Object, + null!, + accessor.Object, + runtimeConfigProvider); + return (engine, metadata.Object); + } + + private static SqlMutationEngine CreateUninitializedEngine(IAuthorizationResolver authorizationResolver) + { + SqlMutationEngine engine = (SqlMutationEngine)RuntimeHelpers.GetUninitializedObject(typeof(SqlMutationEngine)); + typeof(SqlMutationEngine).GetField("_authorizationResolver", BindingFlags.Instance | BindingFlags.NonPublic)! + .SetValue(engine, authorizationResolver); + return engine; + } + + private static (SqlMutationEngine Engine, StoredProcedureRequestContext Context, string DataSourceName) + CreateStoredProcedureFixture(EntityActionOperation operation, JsonArray result) + { + RuntimeConfig runtimeConfig = new( + Schema: string.Empty, + DataSource: new DataSource(DatabaseType.MSSQL, string.Empty), + Entities: new RuntimeEntities(new Dictionary())); + RuntimeConfigProvider runtimeConfigProvider = TestHelper.GenerateInMemoryRuntimeConfigProvider(runtimeConfig); + string dataSourceName = runtimeConfigProvider.GetConfig().DefaultDataSourceName; + + StoredProcedureDefinition definition = new(); + DatabaseStoredProcedure procedure = new("dbo", "procedure") + { + StoredProcedureDefinition = definition + }; + Mock metadata = new(); + metadata.SetupGet(x => x.EntityToDatabaseObject) + .Returns(new Dictionary { ["Procedure"] = procedure }); + metadata.Setup(x => x.GetStoredProcedureDefinition("Procedure")).Returns(definition); + metadata.Setup(x => x.GetDatabaseType()).Returns(DatabaseType.MSSQL); + + Mock metadataFactory = new(); + metadataFactory.Setup(x => x.GetMetadataProvider(dataSourceName)).Returns(metadata.Object); + Mock queryBuilder = new(); + queryBuilder.Setup(x => x.Build(It.IsAny())).Returns("EXEC dbo.procedure"); + Mock queryExecutor = new(); + queryExecutor.Setup(x => x.ExecuteQueryAsync( + It.IsAny(), + It.IsAny>(), + It.IsAny?, Task>>(), + dataSourceName, + It.IsAny(), + It.IsAny? >())) + .ReturnsAsync(result); + Mock queryManagerFactory = new(); + queryManagerFactory.Setup(x => x.GetQueryBuilder(DatabaseType.MSSQL)).Returns(queryBuilder.Object); + queryManagerFactory.Setup(x => x.GetQueryExecutor(DatabaseType.MSSQL)).Returns(queryExecutor.Object); + + DefaultHttpContext httpContext = new(); + httpContext.Request.Scheme = "https"; + httpContext.Request.Host = new HostString("example.test"); + httpContext.Request.Path = "/api/procedure"; + Mock accessor = new(); + accessor.Setup(x => x.HttpContext).Returns(httpContext); + + SqlMutationEngine engine = new( + queryManagerFactory.Object, + metadataFactory.Object, + new Mock().Object, + new Mock().Object, + null!, + accessor.Object, + runtimeConfigProvider); + StoredProcedureRequestContext context = new("Procedure", procedure, null, operation); + context.PopulateResolvedParameters(); + return (engine, context, dataSourceName); + } + + private static T InvokeStatic(string methodName, params object?[] arguments) + { + MethodInfo method = typeof(SqlMutationEngine).GetMethod(methodName, BindingFlags.Static | BindingFlags.NonPublic)!; + return (T)method.Invoke(null, arguments)!; + } + + private static T InvokeInstance(SqlMutationEngine instance, string methodName, params object?[] arguments) + { + MethodInfo method = typeof(SqlMutationEngine).GetMethod(methodName, BindingFlags.Instance | BindingFlags.NonPublic)!; + return (T)method.Invoke(instance, arguments)!; + } + } +} diff --git a/src/Service.Tests/UnitTests/SqlPaginationUtilUnitTests.cs b/src/Service.Tests/UnitTests/SqlPaginationUtilUnitTests.cs index 88ecdd2c29..dc14315539 100644 --- a/src/Service.Tests/UnitTests/SqlPaginationUtilUnitTests.cs +++ b/src/Service.Tests/UnitTests/SqlPaginationUtilUnitTests.cs @@ -2,8 +2,14 @@ // Licensed under the MIT License. using System.Collections.Specialized; +using System.Collections.Generic; using System.Text.Json; +using Azure.DataApiBuilder.Config.DatabasePrimitives; +using Azure.DataApiBuilder.Core.Models; using Azure.DataApiBuilder.Core.Resolvers; +using Azure.DataApiBuilder.Service.Exceptions; +using Azure.DataApiBuilder.Service.GraphQLBuilder.GraphQLTypes; +using Microsoft.AspNetCore.Http; using Microsoft.VisualStudio.TestTools.UnitTesting; namespace Azure.DataApiBuilder.Service.Tests.UnitTests @@ -133,5 +139,121 @@ public void GetConsolidatedNextLinkForPagination_Relative_ReturnsPathAndQuery() StringAssert.StartsWith(nextLink, "/api/Book"); Assert.IsFalse(nextLink.Contains("localhost"), "Relative link should not contain the host."); } + + [TestMethod] + public void TryResolveJsonElementToScalarVariable_HandlesEveryJsonKind() + { + using JsonDocument document = JsonDocument.Parse("[\"text\",12.5,null,true,false,{},[]]"); + JsonElement.ArrayEnumerator elements = document.RootElement.EnumerateArray(); + + elements.MoveNext(); + Assert.IsTrue(SqlPaginationUtil.TryResolveJsonElementToScalarVariable(elements.Current, out object? text)); + Assert.AreEqual("text", text); + elements.MoveNext(); + Assert.IsTrue(SqlPaginationUtil.TryResolveJsonElementToScalarVariable(elements.Current, out object? number)); + Assert.AreEqual(12.5, number); + elements.MoveNext(); + Assert.IsTrue(SqlPaginationUtil.TryResolveJsonElementToScalarVariable(elements.Current, out object? nullValue)); + Assert.IsNull(nullValue); + elements.MoveNext(); + Assert.IsTrue(SqlPaginationUtil.TryResolveJsonElementToScalarVariable(elements.Current, out object? trueValue)); + Assert.AreEqual(true, trueValue); + elements.MoveNext(); + Assert.IsTrue(SqlPaginationUtil.TryResolveJsonElementToScalarVariable(elements.Current, out object? falseValue)); + Assert.AreEqual(false, falseValue); + elements.MoveNext(); + Assert.IsFalse(SqlPaginationUtil.TryResolveJsonElementToScalarVariable(elements.Current, out _)); + elements.MoveNext(); + Assert.IsFalse(SqlPaginationUtil.TryResolveJsonElementToScalarVariable(elements.Current, out _)); + } + + [TestMethod] + public void MakeCursorFromJsonElement_IncludesOrderByAndRemainingPrimaryKey() + { + using JsonDocument document = JsonDocument.Parse("{\"name\":\"Ada\",\"id\":7}"); + + string cursor = SqlPaginationUtil.MakeCursorFromJsonElement( + document.RootElement, + new List { "id" }, + new List { new("dbo", "authors", "name", "a", OrderBy.DESC) }, + entityName: "Author"); + string decoded = SqlPaginationUtil.Base64Decode(cursor); + + StringAssert.Contains(decoded, "name"); + StringAssert.Contains(decoded, "Ada"); + StringAssert.Contains(decoded, "id"); + StringAssert.Contains(decoded, "\"Direction\":1"); + } + + [TestMethod] + public void MakeCursorFromJsonElement_GroupByOmitsPrimaryKeys() + { + using JsonDocument document = JsonDocument.Parse("{\"total\":12}"); + + string decoded = SqlPaginationUtil.Base64Decode(SqlPaginationUtil.MakeCursorFromJsonElement( + document.RootElement, + new List { "missing" }, + new List { new("", "", "total", "", OrderBy.ASC) }, + isGroupByQuery: true)); + + StringAssert.Contains(decoded, "total"); + Assert.IsFalse(decoded.Contains("missing")); + } + + [TestMethod] + public void MakeCursorFromJsonElement_RejectsNonScalarCursorValue() + { + using JsonDocument document = JsonDocument.Parse("{\"id\":{}}"); + + Assert.ThrowsException(() => SqlPaginationUtil.MakeCursorFromJsonElement( + document.RootElement, + new List { "id" }, + orderByColumns: null)); + } + + [TestMethod] + public void ConstructBaseUriForPagination_UsesValidForwardedHeadersAndBaseRoute() + { + DefaultHttpContext context = new(); + context.Request.Scheme = "http"; + context.Request.Host = new HostString("internal:5000"); + context.Request.Path = "/api/books"; + context.Request.Headers["X-Forwarded-Proto"] = " HTTPS "; + context.Request.Headers["X-Forwarded-Host"] = " example.com:8443 "; + + string result = SqlPaginationUtil.ConstructBaseUriForPagination(context, "/gateway"); + + Assert.AreEqual("https://example.com:8443/gateway/api/books", result); + } + + [TestMethod] + public void ConstructBaseUriForPagination_InvalidForwardedHeadersFallBackToRequest() + { + DefaultHttpContext context = new(); + context.Request.Scheme = "http"; + context.Request.Host = new HostString("localhost:5000"); + context.Request.Path = "/api/books"; + context.Request.Headers["X-Forwarded-Proto"] = "ftp"; + context.Request.Headers["X-Forwarded-Host"] = "bad host@value"; + + Assert.AreEqual("http://localhost:5000/api/books", SqlPaginationUtil.ConstructBaseUriForPagination(context)); + } + + [TestMethod] + public void FormatQueryString_IgnoresBlankKeysAndValues() + { + NameValueCollection parameters = new() + { + { " ", "ignored" }, + { "$filter", " " }, + { "$first", "10" } + }; + + string result = SqlPaginationUtil.FormatQueryString(parameters); + + StringAssert.Contains(result, "$first=10"); + Assert.IsFalse(result.Contains("ignored")); + Assert.IsFalse(result.Contains("filter")); + } } } diff --git a/src/Service.Tests/UnitTests/SqlQueryEngineHelperTests.cs b/src/Service.Tests/UnitTests/SqlQueryEngineHelperTests.cs new file mode 100644 index 0000000000..3998fbf370 --- /dev/null +++ b/src/Service.Tests/UnitTests/SqlQueryEngineHelperTests.cs @@ -0,0 +1,187 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +using System; +using System.Collections.Generic; +using System.Data.Common; +using System.Reflection; +using System.Runtime.CompilerServices; +using System.Text.Json; +using System.Text.Json.Nodes; +using System.Threading.Tasks; +using Azure.DataApiBuilder.Auth; +using Azure.DataApiBuilder.Config; +using Azure.DataApiBuilder.Config.ObjectModel; +using Azure.DataApiBuilder.Core.Configurations; +using Azure.DataApiBuilder.Core.Models; +using Azure.DataApiBuilder.Core.Resolvers; +using Azure.DataApiBuilder.Core.Resolvers.Factories; +using Azure.DataApiBuilder.Core.Services; +using Azure.DataApiBuilder.Core.Services.Cache; +using Azure.DataApiBuilder.Core.Services.MetadataProviders; +using Microsoft.AspNetCore.Http; +using Microsoft.Extensions.Logging.Abstractions; +using Microsoft.VisualStudio.TestTools.UnitTesting; +using Moq; + +namespace Azure.DataApiBuilder.Service.Tests.UnitTests +{ + [TestClass] + public class SqlQueryEngineHelperTests + { + private const string DataSourceName = "default"; + private const string EntityName = "Book"; + + [DataTestMethod] + [DataRow("{\"value\":1}", true)] + [DataRow(null, false)] + public void ParseResultIntoJsonDocument_HandlesValuesAndNull(string? json, bool hasObject) + { + JsonElement? element = json is null ? null : JsonDocument.Parse(json).RootElement.Clone(); + MethodInfo method = typeof(SqlQueryEngine).GetMethod( + "ParseResultIntoJsonDocument", + BindingFlags.Static | BindingFlags.NonPublic)!; + + using JsonDocument result = (JsonDocument)method.Invoke(null, new object?[] { element })!; + + Assert.AreEqual(hasObject ? JsonValueKind.Object : JsonValueKind.Null, result.RootElement.ValueKind); + } + + [DataTestMethod] + [DataRow("[{\"id\":1}]", true)] + [DataRow("[]", false)] + [DataRow(null, false)] + public async Task ExecuteStoredProcedureCore_HandlesResultShapes(string? json, bool expectsDocument) + { + JsonArray? resultArray = json is null ? null : JsonNode.Parse(json)!.AsArray(); + (SqlQueryEngine engine, Mock executor) = CreateEngine(); + executor.Setup(x => x.ExecuteQueryAsync( + It.IsAny(), + It.IsAny>(), + It.IsAny?, Task>>(), + DataSourceName, + It.IsAny(), + It.IsAny?>())) + .ReturnsAsync(resultArray!); + SqlExecuteStructure structure = CreateUninitializedStructure(); + MethodInfo method = typeof(SqlQueryEngine).GetMethod( + "ExecuteAsync", + BindingFlags.Instance | BindingFlags.NonPublic, + binder: null, + types: new[] { typeof(SqlExecuteStructure), typeof(string) }, + modifiers: null)!; + + using JsonDocument? result = await (Task)method.Invoke( + engine, + new object[] { structure, DataSourceName })!; + + Assert.AreEqual(expectsDocument, result is not null); + } + + [DataTestMethod] + [DataRow(true)] + [DataRow(false)] + public async Task ExecuteListCore_ReturnsExecutorResult(bool returnList) + { + (SqlQueryEngine engine, Mock executor) = CreateEngine(); + List? expected = returnList + ? new List { JsonDocument.Parse("{\"id\":1}") } + : null; + executor.Setup(x => x.ExecuteQueryAsync( + It.IsAny(), + It.IsAny>(), + It.IsAny?, Task>>>(), + DataSourceName, + It.IsAny(), + It.IsAny?>())) + .ReturnsAsync(expected!); + SqlQueryStructure structure = CreateUninitializedStructure(); + MethodInfo method = typeof(SqlQueryEngine).GetMethod( + "ExecuteListAsync", + BindingFlags.Instance | BindingFlags.NonPublic, + binder: null, + types: new[] { typeof(SqlQueryStructure), typeof(string) }, + modifiers: null)!; + + List? result = await (Task?>)method.Invoke( + engine, + new object[] { structure, DataSourceName })!; + + Assert.AreSame(expected, result); + if (expected is not null) + { + foreach (JsonDocument document in expected) + { + document.Dispose(); + } + } + } + + private static (SqlQueryEngine Engine, Mock Executor) CreateEngine() + { + RuntimeConfig runtimeConfig = new( + Schema: string.Empty, + DataSource: new DataSource(DatabaseType.MSSQL, string.Empty), + Entities: new RuntimeEntities(new Dictionary())); + runtimeConfig.UpdateDefaultDataSourceName(DataSourceName); + Mock loader = new(null, null); + Mock configProviderMock = new(loader.Object); + configProviderMock.Setup(x => x.GetConfig()).Returns(runtimeConfig); + RuntimeConfigProvider configProvider = configProviderMock.Object; + Mock queryBuilder = new(); + queryBuilder.Setup(x => x.Build(It.IsAny())).Returns("execute"); + queryBuilder.Setup(x => x.Build(It.IsAny())).Returns("select"); + Mock queryExecutor = new(); + Mock factory = new(); + factory.Setup(x => x.GetQueryBuilder(DatabaseType.MSSQL)).Returns(queryBuilder.Object); + factory.Setup(x => x.GetQueryExecutor(DatabaseType.MSSQL)).Returns(queryExecutor.Object); + Mock metadataProviderFactory = new(); + Mock filterParser = new(configProvider, metadataProviderFactory.Object); + DefaultHttpContext httpContext = new(); + + SqlQueryEngine engine = new( + factory.Object, + metadataProviderFactory.Object, + new HttpContextAccessor { HttpContext = httpContext }, + Mock.Of(), + filterParser.Object, + NullLogger.Instance, + configProvider, + (DabCacheService)RuntimeHelpers.GetUninitializedObject(typeof(DabCacheService))); + return (engine, queryExecutor); + } + + private static T CreateUninitializedStructure() + { + T structure = (T)RuntimeHelpers.GetUninitializedObject(typeof(T)); + SetProperty(structure!, "EntityName", EntityName); + SetProperty(structure!, "Parameters", new Dictionary()); + return structure; + } + + private static void SetProperty(object target, string name, object value) + { + Type? type = target.GetType(); + while (type is not null) + { + PropertyInfo? property = type.GetProperty(name, BindingFlags.Instance | BindingFlags.Public | BindingFlags.NonPublic | BindingFlags.DeclaredOnly); + if (property is not null) + { + property.SetValue(target, value); + return; + } + + FieldInfo? field = type.GetField($"<{name}>k__BackingField", BindingFlags.Instance | BindingFlags.NonPublic | BindingFlags.DeclaredOnly); + if (field is not null) + { + field.SetValue(target, value); + return; + } + + type = type.BaseType; + } + + Assert.Fail($"Member {name} was not found."); + } + } +} diff --git a/src/Service.Tests/UnitTests/SqlQueryExecutorUnitTests.cs b/src/Service.Tests/UnitTests/SqlQueryExecutorUnitTests.cs index b640c79cd9..4f2e7dfe20 100644 --- a/src/Service.Tests/UnitTests/SqlQueryExecutorUnitTests.cs +++ b/src/Service.Tests/UnitTests/SqlQueryExecutorUnitTests.cs @@ -684,7 +684,8 @@ public void ValidateStreamingLogicForEmptyCellsAsync() private static (MsSqlQueryExecutor QueryExecutor, RuntimeConfigProvider Provider) CreateQueryExecutorForPoolingTest( string connectionString, bool enableObo, - Mock httpContextAccessor) + Mock httpContextAccessor, + IOboTokenProvider? oboTokenProvider = null) { DataSource dataSource = new( DatabaseType: DatabaseType.MSSQL, @@ -719,7 +720,7 @@ private static (MsSqlQueryExecutor QueryExecutor, RuntimeConfigProvider Provider Mock>> queryExecutorLogger = new(); DbExceptionParser dbExceptionParser = new MsSqlDbExceptionParser(provider); - MsSqlQueryExecutor queryExecutor = new(provider, dbExceptionParser, queryExecutorLogger.Object, httpContextAccessor.Object); + MsSqlQueryExecutor queryExecutor = new(provider, dbExceptionParser, queryExecutorLogger.Object, httpContextAccessor.Object, oboTokenProvider: oboTokenProvider); return (queryExecutor, provider); } @@ -1013,6 +1014,107 @@ private static Mock CreateHttpContextAccessorWithAuthentic return httpContextAccessor; } + [TestMethod] + public void AddStatementId_NoHttpContextReturns() + { + Mock accessor = new(); + accessor.Setup(x => x.HttpContext).Returns(value: null); + (MsSqlQueryExecutor executor, _) = CreateQueryExecutorForPoolingTest( + "Server=localhost;Database=test;", false, accessor); + + InvokeAddStatementId(executor, "id-1"); + } + + [TestMethod] + public void AddStatementId_AddsThenAppendsValues() + { + DefaultHttpContext context = new(); + Mock accessor = new(); + accessor.Setup(x => x.HttpContext).Returns(context); + (MsSqlQueryExecutor executor, _) = CreateQueryExecutorForPoolingTest( + "Server=localhost;Database=test;", false, accessor); + + InvokeAddStatementId(executor, "id-1"); + InvokeAddStatementId(executor, "id-2"); + + Assert.AreEqual("id-1;id-2", context.Items["QueryIdentifyingIds"]); + } + + [TestMethod] + public void AddStatementId_NonStringExistingValueIsPreserved() + { + DefaultHttpContext context = new(); + context.Items["QueryIdentifyingIds"] = 42; + Mock accessor = new(); + accessor.Setup(x => x.HttpContext).Returns(context); + (MsSqlQueryExecutor executor, _) = CreateQueryExecutorForPoolingTest( + "Server=localhost;Database=test;", false, accessor); + + InvokeAddStatementId(executor, "id-2"); + + Assert.AreEqual(42, context.Items["QueryIdentifyingIds"]); + } + + private static void InvokeAddStatementId(MsSqlQueryExecutor executor, string statementId) => + typeof(MsSqlQueryExecutor).GetMethod( + "AddStatementIDToMiddlewareContext", + BindingFlags.Instance | BindingFlags.NonPublic)!.Invoke(executor, new object[] { statementId }); + + [TestMethod] + public void CreateConnection_MissingDataSourceThrows() + { + Mock accessor = new(); + (MsSqlQueryExecutor executor, _) = CreateQueryExecutorForPoolingTest( + "Server=localhost;Database=test;", false, accessor); + + DataApiBuilderException exception = Assert.ThrowsException(() => + executor.CreateConnection("missing")); + + Assert.AreEqual(DataApiBuilderException.SubStatusCodes.DataSourceNotFound, exception.SubStatusCode); + } + + [TestMethod] + public async Task SetManagedIdentityAccessToken_OboBearerTokenSetsConnectionToken() + { + DefaultHttpContext context = new(); + context.Request.Headers.Authorization = "Bearer incoming-token"; + context.User = new System.Security.Claims.ClaimsPrincipal(new System.Security.Claims.ClaimsIdentity("test")); + Mock accessor = new(); + accessor.Setup(x => x.HttpContext).Returns(context); + Mock tokenProvider = new(); + tokenProvider.Setup(x => x.GetAccessTokenOnBehalfOfAsync( + context.User, "incoming-token", "https://database.windows.net")) + .ReturnsAsync("database-token"); + (MsSqlQueryExecutor executor, RuntimeConfigProvider configProvider) = CreateQueryExecutorForPoolingTest( + "Server=localhost;Database=test;", true, accessor, tokenProvider.Object); + using SqlConnection connection = new(); + + await executor.SetManagedIdentityAccessTokenIfAnyAsync( + connection, configProvider.GetConfig().DefaultDataSourceName); + + Assert.AreEqual("database-token", connection.AccessToken); + } + + [DataTestMethod] + [DataRow("")] + [DataRow("Basic credentials")] + public async Task SetManagedIdentityAccessToken_OboMissingBearerTokenThrows(string authorization) + { + DefaultHttpContext context = new(); + context.Request.Headers.Authorization = authorization; + Mock accessor = new(); + accessor.Setup(x => x.HttpContext).Returns(context); + (MsSqlQueryExecutor executor, RuntimeConfigProvider configProvider) = CreateQueryExecutorForPoolingTest( + "Server=localhost;Database=test;", true, accessor, Mock.Of()); + using SqlConnection connection = new(); + + DataApiBuilderException exception = await Assert.ThrowsExceptionAsync(() => + executor.SetManagedIdentityAccessTokenIfAnyAsync( + connection, configProvider.GetConfig().DefaultDataSourceName)); + + Assert.AreEqual(DataApiBuilderException.SubStatusCodes.OboAuthenticationFailure, exception.SubStatusCode); + } + #endregion /// diff --git a/src/Service.Tests/UnitTests/SqlQueryStructureHelperTests.cs b/src/Service.Tests/UnitTests/SqlQueryStructureHelperTests.cs new file mode 100644 index 0000000000..3a5f10d663 --- /dev/null +++ b/src/Service.Tests/UnitTests/SqlQueryStructureHelperTests.cs @@ -0,0 +1,178 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +using System.Collections.Generic; +using System.Reflection; +using System.Runtime.CompilerServices; +using Azure.DataApiBuilder.Core.Models; +using Azure.DataApiBuilder.Core.Resolvers; +using Azure.DataApiBuilder.Core.Resolvers.Sql_Query_Structures; +using Azure.DataApiBuilder.Service.Exceptions; +using Azure.DataApiBuilder.Service.GraphQLBuilder.Queries; +using HotChocolate.Language; +using Microsoft.AspNetCore.Http; +using Microsoft.VisualStudio.TestTools.UnitTesting; +using Moq; + +namespace Azure.DataApiBuilder.Service.Tests.UnitTests +{ + [TestClass] + public class SqlQueryStructureHelperTests + { + [DataTestMethod] + [DataRow(false, false, 25u, 1u)] + [DataRow(true, false, 25u, 25u)] + [DataRow(false, true, 25u, 25u)] + public void Limit_ReflectsQueryAndPaginationShape(bool isList, bool isPaginated, uint configured, uint expected) + { + SqlQueryStructure structure = CreateStructure(); + structure.IsListQuery = isList; + typeof(SqlQueryStructure).GetField("_limit", BindingFlags.Instance | BindingFlags.NonPublic)! + .SetValue(structure, configured); + structure.PaginationMetadata = new PaginationMetadata(structure) { IsPaginated = isPaginated }; + + Assert.AreEqual(expected, structure.Limit()); + } + + [TestMethod] + public void Limit_ListQueryWithNullLimit_ReturnsNull() + { + SqlQueryStructure structure = CreateStructure(); + structure.IsListQuery = true; + typeof(SqlQueryStructure).GetField("_limit", BindingFlags.Instance | BindingFlags.NonPublic)! + .SetValue(structure, null); + structure.PaginationMetadata = new PaginationMetadata(structure); + + Assert.IsNull(structure.Limit()); + } + + [TestMethod] + public void IsSubqueryColumn_EvaluatesTableAliasAgainstJoinQueries() + { + SqlQueryStructure structure = CreateStructure(); + structure.JoinQueries.Add("joined", CreateStructure()); + + Assert.IsFalse(structure.IsSubqueryColumn(new Column("dbo", "books", "id"))); + Assert.IsFalse(structure.IsSubqueryColumn(new Column("dbo", "books", "id", "missing"))); + Assert.IsTrue(structure.IsSubqueryColumn(new Column("dbo", "books", "id", "joined"))); + } + + [TestMethod] + public void AddCacheControlOptions_CopiesRequestHeader() + { + SqlQueryStructure structure = CreateStructure(); + HeaderDictionary headers = new() { ["Cache-Control"] = "no-store" }; + + GetPrivateMethod("AddCacheControlOptions").Invoke(structure, new object[] { headers }); + + Assert.AreEqual("no-store", structure.CacheControlOption); + } + + [TestMethod] + public void ProcessPaginationFields_SetsEveryRequestedFlag() + { + SqlQueryStructure structure = CreateStructure(); + ISelectionNode[] selections = + { + new FieldNode(QueryBuilder.PAGINATION_FIELD_NAME), + new FieldNode(QueryBuilder.PAGINATION_TOKEN_FIELD_NAME), + new FieldNode(QueryBuilder.HAS_NEXT_PAGE_FIELD_NAME), + new FieldNode(QueryBuilder.GROUP_BY_FIELD_NAME) + }; + + GetPrivateMethod("ProcessPaginationFields").Invoke(structure, new object[] { selections }); + + Assert.IsTrue(structure.PaginationMetadata.RequestedItems); + Assert.IsTrue(structure.PaginationMetadata.RequestedEndCursor); + Assert.IsTrue(structure.PaginationMetadata.RequestedHasNextPage); + Assert.IsTrue(structure.PaginationMetadata.RequestedGroupBy); + } + + [TestMethod] + public void AddGraphQLFields_FragmentSpreadWithoutContextThrows() + { + SqlQueryStructure structure = CreateStructure(); + ISelectionNode[] selections = + { + new FragmentSpreadNode(null, new NameNode("BookFields"), System.Array.Empty()) + }; + + TargetInvocationException exception = Assert.ThrowsException(() => + GetPrivateMethod("AddGraphQLFields").Invoke(structure, new object?[] { selections, null })); + + Assert.IsInstanceOfType(exception.InnerException); + } + + [TestMethod] + public void AddGraphQLFields_InlineFragmentSkipsIntrospectionField() + { + SqlQueryStructure structure = CreateStructure(); + InlineFragmentNode fragment = new( + location: null, + typeCondition: null, + directives: System.Array.Empty(), + selectionSet: new SelectionSetNode(new ISelectionNode[] { new FieldNode("__typename") })); + + GetPrivateMethod("AddGraphQLFields").Invoke( + structure, + new object?[] { new ISelectionNode[] { fragment }, null }); + } + + [TestMethod] + public void AddGraphQLFields_UnsupportedSelectionThrows() + { + SqlQueryStructure structure = CreateStructure(); + Mock selection = new(); + selection.SetupGet(node => node.Kind).Returns(SyntaxKind.Directive); + + TargetInvocationException exception = Assert.ThrowsException(() => + GetPrivateMethod("AddGraphQLFields").Invoke( + structure, + new object?[] { new[] { selection.Object }, null })); + + Assert.IsInstanceOfType(exception.InnerException); + } + + [TestMethod] + public void ProcessGroupByFieldSelections_NullSelectionReturns() + { + SqlQueryStructure structure = CreateStructure(); + + GetPrivateMethod("ProcessGroupByFieldSelections").Invoke( + structure, + new object[] { new FieldNode("fields"), new HashSet() }); + } + + [TestMethod] + public void ProcessGroupByFieldSelections_MismatchedFieldThrows() + { + SqlQueryStructure structure = CreateStructure(); + FieldNode fields = new FieldNode("fields").WithSelectionSet( + new SelectionSetNode(new ISelectionNode[] { new FieldNode("missing") })); + + TargetInvocationException exception = Assert.ThrowsException(() => + GetPrivateMethod("ProcessGroupByFieldSelections").Invoke( + structure, + new object[] { fields, new HashSet { "id" } })); + + Assert.IsInstanceOfType(exception.InnerException); + } + + private static SqlQueryStructure CreateStructure() + { + SqlQueryStructure structure = (SqlQueryStructure)RuntimeHelpers.GetUninitializedObject(typeof(SqlQueryStructure)); + SetAutoProperty(structure, "JoinQueries", new Dictionary()); + structure.PaginationMetadata = new PaginationMetadata(structure); + return structure; + } + + private static void SetAutoProperty(SqlQueryStructure structure, string propertyName, T value) + { + typeof(SqlQueryStructure).GetField($"<{propertyName}>k__BackingField", BindingFlags.Instance | BindingFlags.NonPublic)! + .SetValue(structure, value); + } + + private static MethodInfo GetPrivateMethod(string methodName) => + typeof(SqlQueryStructure).GetMethod(methodName, BindingFlags.Instance | BindingFlags.NonPublic)!; + } +} diff --git a/src/Service.Tests/UnitTests/SqlQueryStructuresModelTests.cs b/src/Service.Tests/UnitTests/SqlQueryStructuresModelTests.cs new file mode 100644 index 0000000000..7ff209eeac --- /dev/null +++ b/src/Service.Tests/UnitTests/SqlQueryStructuresModelTests.cs @@ -0,0 +1,53 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +using System; +using Azure.DataApiBuilder.Core.Models; +using Azure.DataApiBuilder.Core.Resolvers; +using Microsoft.VisualStudio.TestTools.UnitTesting; + +namespace Azure.DataApiBuilder.Service.Tests.UnitTests +{ + [TestClass] + public class SqlQueryStructuresModelTests + { + [TestMethod] + public void LabelledColumn_EqualityHandlesNullIdentityAndValues() + { + LabelledColumn column = new("dbo", "books", "id", "book_id"); + + Assert.IsFalse(column.Equals((LabelledColumn?)null)); + Assert.IsFalse(column.Equals((object?)null)); + Assert.IsTrue(column.Equals(column)); + Assert.IsTrue(column.Equals((object)column)); + Assert.IsTrue(column.Equals(new LabelledColumn("dbo", "books", "id", "book_id"))); + Assert.IsFalse(column.Equals(new LabelledColumn("dbo", "books", "id", "other"))); + Assert.IsFalse(column.Equals(new object())); + Assert.AreNotEqual(0, column.GetHashCode()); + } + + [TestMethod] + public void PredicateOperand_NullConstructorsThrow() + { + Assert.ThrowsException(() => new PredicateOperand((Column?)null)); + Assert.ThrowsException(() => new PredicateOperand((BaseQueryStructure?)null)); + Assert.ThrowsException(() => new PredicateOperand((string?)null)); + Assert.ThrowsException(() => new PredicateOperand((Predicate?)null)); + } + + [TestMethod] + public void PredicateOperand_StringAndPredicateAccessorsReflectStoredType() + { + PredicateOperand text = new("value"); + Predicate predicate = new(null, PredicateOperation.EXISTS, text); + PredicateOperand nested = new(predicate); + + Assert.AreEqual("value", text.AsString()); + Assert.IsNull(text.AsColumn()); + Assert.IsNull(text.AsPredicate()); + Assert.IsFalse(text.IsPredicate()); + Assert.AreSame(predicate, nested.AsPredicate()); + Assert.IsTrue(nested.IsPredicate()); + } + } +} \ No newline at end of file diff --git a/src/Service.Tests/UnitTests/TypeHelperTests.cs b/src/Service.Tests/UnitTests/TypeHelperTests.cs index 5dc7e7ec66..e2bdebdb00 100644 --- a/src/Service.Tests/UnitTests/TypeHelperTests.cs +++ b/src/Service.Tests/UnitTests/TypeHelperTests.cs @@ -32,6 +32,7 @@ public class TypeHelperTests [DataRow(typeof(bool), EdmPrimitiveTypeKind.Boolean)] [DataRow(typeof(DateTime), EdmPrimitiveTypeKind.DateTimeOffset)] [DataRow(typeof(DateTimeOffset), EdmPrimitiveTypeKind.DateTimeOffset)] + [DataRow(typeof(Date), EdmPrimitiveTypeKind.Date)] [DataRow(typeof(TimeOnly), EdmPrimitiveTypeKind.TimeOfDay)] [DataRow(typeof(TimeSpan), EdmPrimitiveTypeKind.TimeOfDay)] public void GetEdmPrimitiveTypeFromSystemType_MapsExpectedKind(Type systemType, EdmPrimitiveTypeKind expected) @@ -58,6 +59,14 @@ public void GetEdmPrimitiveTypeFromSystemType_UnsupportedType_Throws() () => TypeHelper.GetEdmPrimitiveTypeFromSystemType(typeof(System.Text.StringBuilder))); } + [DataTestMethod] + [DataRow("Boolean", EdmPrimitiveTypeKind.Boolean)] + [DataRow("Date", EdmPrimitiveTypeKind.Date)] + public void GetEdmPrimitiveTypeFromITypeNode_MapsExpectedKind(string graphQlType, EdmPrimitiveTypeKind expected) + { + Assert.AreEqual(expected, TypeHelper.GetEdmPrimitiveTypeFromITypeNode(new NamedTypeNode(graphQlType))); + } + [DataTestMethod] [DataRow(typeof(int), JsonDataType.Integer)] [DataRow(typeof(long), JsonDataType.Integer)] @@ -140,6 +149,7 @@ public void GetValue_ConvertsValueNodesToClrValues() Assert.AreEqual(true, TypeHelper.GetValue(new BooleanValueNode(true))); Assert.AreEqual("hi", TypeHelper.GetValue(new StringValueNode("hi"))); Assert.IsNull(TypeHelper.GetValue(NullValueNode.Default)); + Assert.AreEqual("VALUE", TypeHelper.GetValue(new EnumValueNode("VALUE"))); } [DataTestMethod] From b2f5825ab5d4dd6d1467782612797f3171cb83e0 Mon Sep 17 00:00:00 2001 From: Aaron Burtle Date: Sat, 29 Aug 2026 05:57:45 -0700 Subject: [PATCH 02/19] expand MsSql pipeline coverage --- .../UnitTests/MsSqlQueryBuilderHelperTests.cs | 171 ++++++++++++++++++ .../UnitTests/QueryExecutorHelperTests.cs | 30 ++- .../SqlMetadataProviderHelperTests.cs | 150 ++++++++++++++- .../UnitTests/SqlQueryExecutorUnitTests.cs | 107 +++++++++++ 4 files changed, 450 insertions(+), 8 deletions(-) create mode 100644 src/Service.Tests/UnitTests/MsSqlQueryBuilderHelperTests.cs diff --git a/src/Service.Tests/UnitTests/MsSqlQueryBuilderHelperTests.cs b/src/Service.Tests/UnitTests/MsSqlQueryBuilderHelperTests.cs new file mode 100644 index 0000000000..52d77f3d5a --- /dev/null +++ b/src/Service.Tests/UnitTests/MsSqlQueryBuilderHelperTests.cs @@ -0,0 +1,171 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +using System; +using System.Collections.Generic; +using System.Data; +using Azure.DataApiBuilder.Auth; +using Azure.DataApiBuilder.Config.DatabasePrimitives; +using Azure.DataApiBuilder.Config.ObjectModel; +using Azure.DataApiBuilder.Core.Authorization; +using Azure.DataApiBuilder.Core.Configurations; +using Azure.DataApiBuilder.Core.Models; +using Azure.DataApiBuilder.Core.Resolvers; +using Azure.DataApiBuilder.Core.Services; +using Azure.DataApiBuilder.Core.Services.MetadataProviders; +using Microsoft.AspNetCore.Http; +using Microsoft.VisualStudio.TestTools.UnitTesting; +using Moq; + +namespace Azure.DataApiBuilder.Service.Tests.UnitTests +{ + [TestClass, TestCategory(TestCategory.MSSQL)] + public class MsSqlQueryBuilderHelperTests + { + private const string ENTITY_NAME = "Book"; + private const string SCHEMA_NAME = "dbo"; + private const string TABLE_NAME = "books"; + + private delegate void TryGetColumnCallback(string entity, string field, out string? column); + + private static readonly Dictionary _columnMapping = new() + { + { "id", "id" }, + { "title", "title" } + }; + + [TestMethod] + public void BuildInsert_TriggerWithOnlyAutogeneratedPrimaryKeyUsesScopeIdentity() + { + SourceDefinition sourceDefinition = CreateSourceDefinition( + isInsertTriggerEnabled: true, + isUpdateTriggerEnabled: false, + isPrimaryKeyAutogenerated: true); + ISqlMetadataProvider metadataProvider = CreateMetadataProvider(sourceDefinition); + SqlInsertStructure structure = new( + ENTITY_NAME, + metadataProvider, + CreateAuthorizationResolver(), + CreateFilterParser(), + new Dictionary { ["title"] = "The Hobbit" }, + CreateHttpContext()); + + string query = new MsSqlQueryBuilder().Build(structure); + + StringAssert.Contains(query, "INSERT INTO [dbo].[books] ([title])"); + StringAssert.Contains(query, "WHERE [dbo].[books].[id] = SCOPE_IDENTITY()"); + Assert.IsFalse(query.Contains("#books_T", StringComparison.Ordinal)); + } + + [DataTestMethod] + [DataRow(true, false, "OUTPUT Inserted.[id] AS [id]")] + [DataRow(false, true, "SELECT [id] AS [id], [title] AS [title] from [dbo].[books]")] + public void BuildUpsert_AsymmetricTriggerConfigurationUsesCorrectOutputQualifier( + bool isUpdateTriggerEnabled, + bool isInsertTriggerEnabled, + string expectedInsertBranch) + { + SourceDefinition sourceDefinition = CreateSourceDefinition( + isInsertTriggerEnabled, + isUpdateTriggerEnabled, + isPrimaryKeyAutogenerated: false); + ISqlMetadataProvider metadataProvider = CreateMetadataProvider(sourceDefinition); + SqlUpsertQueryStructure structure = new( + ENTITY_NAME, + metadataProvider, + CreateAuthorizationResolver(), + CreateFilterParser(), + new Dictionary + { + ["id"] = 1, + ["title"] = "The Hobbit" + }, + incrementalUpdate: false, + CreateHttpContext()); + + string query = new MsSqlQueryBuilder().Build(structure); + + StringAssert.Contains(query, expectedInsertBranch); + } + + private static SourceDefinition CreateSourceDefinition( + bool isInsertTriggerEnabled, + bool isUpdateTriggerEnabled, + bool isPrimaryKeyAutogenerated) + { + SourceDefinition sourceDefinition = new() + { + PrimaryKey = new() { "id" }, + IsInsertDMLTriggerEnabled = isInsertTriggerEnabled, + IsUpdateDMLTriggerEnabled = isUpdateTriggerEnabled + }; + sourceDefinition.Columns.Add("id", new ColumnDefinition + { + SystemType = typeof(int), + DbType = DbType.Int32, + IsAutoGenerated = isPrimaryKeyAutogenerated + }); + sourceDefinition.Columns.Add("title", new ColumnDefinition + { + SystemType = typeof(string), + DbType = DbType.String + }); + return sourceDefinition; + } + + private static ISqlMetadataProvider CreateMetadataProvider(SourceDefinition sourceDefinition) + { + DatabaseTable table = new(SCHEMA_NAME, TABLE_NAME) + { + TableDefinition = sourceDefinition, + SourceType = EntitySourceType.Table + }; + Mock metadataProvider = new(); + metadataProvider.Setup(x => x.EntityToDatabaseObject) + .Returns(new Dictionary { [ENTITY_NAME] = table }); + metadataProvider.Setup(x => x.GetSourceDefinition(ENTITY_NAME)).Returns(sourceDefinition); + metadataProvider.Setup(x => x.GetDatabaseType()).Returns(DatabaseType.MSSQL); + + string? backingColumn; + metadataProvider.Setup(x => x.TryGetBackingColumn(It.IsAny(), It.IsAny(), out backingColumn)) + .Callback(new TryGetColumnCallback((string entity, string field, out string? column) + => _columnMapping.TryGetValue(field, out column))) + .Returns((string entity, string field, string? column) => _columnMapping.ContainsKey(field)); + + string? exposedColumn; + metadataProvider.Setup(x => x.TryGetExposedColumnName(It.IsAny(), It.IsAny(), out exposedColumn)) + .Callback(new TryGetColumnCallback((string entity, string field, out string? column) + => _columnMapping.TryGetValue(field, out column))) + .Returns((string entity, string field, string? column) => _columnMapping.ContainsKey(field)); + + return metadataProvider.Object; + } + + private static IAuthorizationResolver CreateAuthorizationResolver() + { + Mock authorizationResolver = new(); + authorizationResolver + .Setup(x => x.ResolveDBPolicy( + It.IsAny(), + It.IsAny(), + It.IsAny(), + It.IsAny())) + .Returns(ResolvedDatabasePolicy.Empty); + return authorizationResolver.Object; + } + + private static GQLFilterParser CreateFilterParser() + { + RuntimeConfigProvider runtimeConfigProvider = TestHelper.GetRuntimeConfigProvider(TestHelper.GetRuntimeConfigLoader()); + Mock metadataProviderFactory = new(); + return new GQLFilterParser(runtimeConfigProvider, metadataProviderFactory.Object); + } + + private static DefaultHttpContext CreateHttpContext() + { + DefaultHttpContext httpContext = new(); + httpContext.Request.Headers[AuthorizationResolver.CLIENT_ROLE_HEADER] = "authenticated"; + return httpContext; + } + } +} diff --git a/src/Service.Tests/UnitTests/QueryExecutorHelperTests.cs b/src/Service.Tests/UnitTests/QueryExecutorHelperTests.cs index 93a3ac2169..2076be6162 100644 --- a/src/Service.Tests/UnitTests/QueryExecutorHelperTests.cs +++ b/src/Service.Tests/UnitTests/QueryExecutorHelperTests.cs @@ -27,7 +27,7 @@ namespace Azure.DataApiBuilder.Service.Tests.UnitTests { - [TestClass] + [TestClass, TestCategory(TestCategory.MSSQL)] public class QueryExecutorHelperTests { [DataTestMethod] @@ -319,6 +319,34 @@ public async Task MsSqlGetMultipleResultSets_NoMutationResultWithArgumentsThrows Assert.AreEqual(System.Net.HttpStatusCode.NotFound, exception.StatusCode); } + [TestMethod] + public async Task MsSqlGetMultipleResultSets_NoMutationResultWithoutArgumentsThrowsInternalServerError() + { + MsSqlQueryExecutor executor = CreateExecutor(maxResponseSizeEnabled: false); + DataTable schema = new(); + schema.Columns.Add("ColumnName", typeof(string)); + schema.Columns.Add("ColumnSize", typeof(int)); + schema.Rows.Add(MsSqlQueryBuilder.COUNT_ROWS_WITH_GIVEN_PK, 4); + Mock reader = new(); + reader.SetupGet(x => x.RecordsAffected).Returns(1); + reader.SetupGet(x => x.HasRows).Returns(true); + reader.SetupSequence(x => x.ReadAsync(It.IsAny())) + .ReturnsAsync(true) + .ReturnsAsync(false); + reader.Setup(x => x.GetSchemaTable()).Returns(schema); + reader.Setup(x => x.GetOrdinal(MsSqlQueryBuilder.COUNT_ROWS_WITH_GIVEN_PK)).Returns(0); + reader.Setup(x => x.IsDBNull(0)).Returns(false); + reader.Setup(x => x[MsSqlQueryBuilder.COUNT_ROWS_WITH_GIVEN_PK]).Returns(0); + reader.Setup(x => x.GetFieldType(0)).Returns(typeof(int)); + reader.Setup(x => x.NextResultAsync(It.IsAny())).ReturnsAsync(false); + + DataApiBuilderException exception = await Assert.ThrowsExceptionAsync(() => + executor.GetMultipleResultSetsIfAnyAsync(reader.Object)); + + Assert.AreEqual(HttpStatusCode.InternalServerError, exception.StatusCode); + Assert.AreEqual(DataApiBuilderException.SubStatusCodes.UnexpectedError, exception.SubStatusCode); + } + [TestMethod] public async Task GetJsonResultAsync_StreamsWhenResponseLimitIsEnabled() { diff --git a/src/Service.Tests/UnitTests/SqlMetadataProviderHelperTests.cs b/src/Service.Tests/UnitTests/SqlMetadataProviderHelperTests.cs index 30bf7eb8e2..470c5bf37d 100644 --- a/src/Service.Tests/UnitTests/SqlMetadataProviderHelperTests.cs +++ b/src/Service.Tests/UnitTests/SqlMetadataProviderHelperTests.cs @@ -4,8 +4,10 @@ using System; using System.Collections.Generic; using System.Data; +using System.Net; using System.Reflection; using System.Runtime.CompilerServices; +using System.Text.Json.Nodes; using Azure.DataApiBuilder.Config.DatabasePrimitives; using Azure.DataApiBuilder.Config.ObjectModel; using Azure.DataApiBuilder.Core.Configurations; @@ -16,12 +18,14 @@ using Azure.DataApiBuilder.Service.Exceptions; using HotChocolate.Language; using Microsoft.Data.SqlClient; +using Microsoft.Data.SqlTypes; using Microsoft.Extensions.Logging; using Microsoft.VisualStudio.TestTools.UnitTesting; +using Moq; namespace Azure.DataApiBuilder.Service.Tests.UnitTests { - [TestClass] + [TestClass, TestCategory(TestCategory.MSSQL)] public class SqlMetadataProviderHelperTests { [TestMethod] @@ -155,18 +159,127 @@ public void GetEntityName_ResolvesEntityAndSingularNameAndRejectsUnknownType() } [TestMethod] - public void ParseSchemaAndDbTableName_HandlesDefaultExplicitPostgresAndMySqlCases() + public void ParseSchemaAndDbTableName_HandlesDefaultAndExplicitMsSqlCases() { MsSqlMetadataProvider provider = CreateProvider(); Assert.AreEqual(("dbo", "books"), provider.ParseSchemaAndDbTableName("books")); Assert.AreEqual(("custom", "books"), provider.ParseSchemaAndDbTableName("custom.books")); + } + + [TestMethod] + public void PopulateColumnDefinitionWithHasDefaultAndDbType_MapsMsSqlMetadata() + { + SourceDefinition definition = new(); + definition.Columns["title"] = new ColumnDefinition(typeof(string)); + definition.Columns["published"] = new ColumnDefinition(typeof(DateTime)); + definition.Columns["embedding"] = new ColumnDefinition(typeof(SqlVector)); + + DataTable columns = new(); + columns.Columns.Add("COLUMN_NAME", typeof(string)); + columns.Columns.Add("COLUMN_DEFAULT", typeof(object)); + columns.Columns.Add("DATA_TYPE", typeof(string)); + columns.Rows.Add("title", DBNull.Value, "nvarchar"); + columns.Rows.Add("published", "getdate()", "date"); + columns.Rows.Add("embedding", DBNull.Value, "varbinary"); + columns.Rows.Add("not_configured", DBNull.Value, "int"); + + MsSqlMetadataProvider provider = CreateProvider(); + GetMsSqlMethod("PopulateColumnDefinitionWithHasDefaultAndDbType") + .Invoke(provider, new object[] { definition, columns }); + + ColumnDefinition title = definition.Columns["title"]; + Assert.IsFalse(title.HasDefault); + Assert.IsNull(title.DefaultValue); + Assert.AreEqual(DbType.String, title.DbType); + Assert.AreEqual(SqlDbType.NVarChar, title.SqlDbType); + + ColumnDefinition published = definition.Columns["published"]; + Assert.IsTrue(published.HasDefault); + Assert.AreEqual("getdate()", published.DefaultValue); + Assert.AreEqual(DbType.Date, published.DbType); + Assert.AreEqual(SqlDbType.Date, published.SqlDbType); + + ColumnDefinition embedding = definition.Columns["embedding"]; + Assert.AreEqual(typeof(float[]), embedding.SystemType); + Assert.AreEqual(typeof(float), embedding.ElementSystemType); + Assert.IsTrue(embedding.IsArrayType); + Assert.AreEqual(DbType.Single, embedding.DbType); + Assert.AreEqual(SqlDbType.Vector, embedding.SqlDbType); + } + + [TestMethod] + public void PopulateMetadataForLinkingObject_MultipleCreateDisabledReturnsWithoutChanges() + { + MsSqlMetadataProvider provider = CreateProvider(); + Dictionary sourceObjects = new(); + + GetMsSqlMethod("PopulateMetadataForLinkingObject").Invoke(provider, new object[] + { + "Book", "Author", "dbo.book_authors", sourceObjects + }); + + Assert.AreEqual(0, sourceObjects.Count); + } + + [TestMethod] + public void TryResolveDbType_UnknownSqlTypeReturnsFalse() + { + MsSqlMetadataProvider provider = CreateProvider(); + object?[] arguments = new object?[] { "future_datetime", null }; + + bool result = (bool)GetMsSqlMethod("TryResolveDbType").Invoke(provider, arguments)!; + + Assert.IsFalse(result); + Assert.AreEqual((DbType)0, arguments[1]); + } + + [TestMethod] + public async System.Threading.Tasks.Task GenerateAutoentitiesIntoEntities_NullConfigurationReturns() + { + MsSqlMetadataProvider provider = CreateProvider(); + + System.Threading.Tasks.Task task = (System.Threading.Tasks.Task)GetMsSqlMethod("GenerateAutoentitiesIntoEntities") + .Invoke(provider, new object?[] { null })!; + + await task; + } + + [TestMethod] + public async System.Threading.Tasks.Task GenerateAutoentitiesIntoEntities_NullResultObjectThrows() + { + MsSqlMetadataProvider provider = CreateProvider(); + ConfigureAutoentityQuery(provider, new JsonArray((JsonNode?)null)); + IReadOnlyDictionary autoentities = new Dictionary + { + ["all"] = new Autoentity(null, null, null) + }; + + System.Threading.Tasks.Task task = (System.Threading.Tasks.Task)GetMsSqlMethod("GenerateAutoentitiesIntoEntities") + .Invoke(provider, new object?[] { autoentities })!; - SetBaseField(provider, "_databaseType", DatabaseType.PostgreSQL); - SetBaseAutoProperty(provider, "ConnectionString", "Host=localhost;Database=db;SearchPath=tenant"); - Assert.AreEqual(("tenant", "books"), provider.ParseSchemaAndDbTableName("books")); + DataApiBuilderException exception = await Assert.ThrowsExceptionAsync(() => task); + Assert.AreEqual(HttpStatusCode.InternalServerError, exception.StatusCode); + } - SetBaseField(provider, "_databaseType", DatabaseType.MySQL); - Assert.ThrowsException(() => provider.ParseSchemaAndDbTableName("custom.books")); + [TestMethod] + public async System.Threading.Tasks.Task GenerateAutoentitiesIntoEntities_IncompleteResultObjectIsSkipped() + { + MsSqlMetadataProvider provider = CreateProvider(); + ConfigureAutoentityQuery(provider, new JsonArray(new JsonObject + { + ["entity_name"] = "Book", + ["object"] = "books" + })); + IReadOnlyDictionary autoentities = new Dictionary + { + ["all"] = new Autoentity(null, null, null) + }; + + System.Threading.Tasks.Task task = (System.Threading.Tasks.Task)GetMsSqlMethod("GenerateAutoentitiesIntoEntities") + .Invoke(provider, new object?[] { autoentities })!; + + await task; + Assert.AreEqual(0, provider.EntityToDatabaseObject.Count); } [TestMethod] @@ -339,6 +452,7 @@ private static MsSqlMetadataProvider CreateProvider( SetBaseField(provider, "_linkingEntities", new Dictionary()); SetBaseAutoProperty(provider, "EntityBackingColumnsToExposedNames", new Dictionary>()); SetBaseAutoProperty(provider, "EntityExposedNamesToBackingColumnNames", new Dictionary>()); + SetBaseField(provider, "_logger", Microsoft.Extensions.Logging.Abstractions.NullLogger.Instance); Dictionary entities = entity is null ? new Dictionary() @@ -349,6 +463,8 @@ private static MsSqlMetadataProvider CreateProvider( Entities: new RuntimeEntities(entities)); RuntimeConfigProvider configProvider = TestHelper.GenerateInMemoryRuntimeConfigProvider(runtimeConfig); SetBaseField(provider, "_runtimeConfigProvider", configProvider); + typeof(MsSqlMetadataProvider).GetField("_runtimeConfigProvider", BindingFlags.Instance | BindingFlags.NonPublic)! + .SetValue(provider, configProvider); SetBaseField(provider, "_dataSourceName", configProvider.GetConfig().DefaultDataSourceName); return provider; } @@ -370,6 +486,26 @@ private static Dictionary> GetMap(MsSqlMetada private static MethodInfo GetBaseMethod(string methodName) => typeof(MsSqlMetadataProvider).BaseType!.GetMethod(methodName, BindingFlags.Instance | BindingFlags.NonPublic)!; + private static MethodInfo GetMsSqlMethod(string methodName) => + typeof(MsSqlMetadataProvider).GetMethod(methodName, BindingFlags.Instance | BindingFlags.NonPublic)!; + + private static void ConfigureAutoentityQuery(MsSqlMetadataProvider provider, JsonArray result) + { + Mock queryExecutor = new(); + queryExecutor.Setup(x => x.ExecuteQueryAsync( + It.IsAny(), + It.IsAny>(), + It.IsAny?, System.Threading.Tasks.Task>>(), + It.IsAny(), + It.IsAny(), + It.IsAny>())) + .ReturnsAsync(result); + Mock queryBuilder = new(); + queryBuilder.Setup(x => x.BuildGetAutoentitiesQuery()).Returns("SELECT autoentities"); + SetBaseAutoProperty(provider, "QueryExecutor", queryExecutor.Object); + SetBaseAutoProperty(provider, "SqlQueryBuilder", queryBuilder.Object); + } + private static void SetBaseField(MsSqlMetadataProvider provider, string fieldName, object value) => typeof(MsSqlMetadataProvider).BaseType!.GetField( fieldName, diff --git a/src/Service.Tests/UnitTests/SqlQueryExecutorUnitTests.cs b/src/Service.Tests/UnitTests/SqlQueryExecutorUnitTests.cs index 4f2e7dfe20..af7eadafca 100644 --- a/src/Service.Tests/UnitTests/SqlQueryExecutorUnitTests.cs +++ b/src/Service.Tests/UnitTests/SqlQueryExecutorUnitTests.cs @@ -1060,6 +1060,39 @@ private static void InvokeAddStatementId(MsSqlQueryExecutor executor, string sta "AddStatementIDToMiddlewareContext", BindingFlags.Instance | BindingFlags.NonPublic)!.Invoke(executor, new object[] { statementId }); + private static void InvokeInfoMessageHandler(SqlConnection connection, int errorNumber, string message) + { + SqlException exception = SqlTestHelper.CreateSqlException(errorNumber, message); + ConstructorInfo constructor = typeof(SqlInfoMessageEventArgs) + .GetConstructors(BindingFlags.Instance | BindingFlags.NonPublic).Single(); + ParameterInfo parameter = constructor.GetParameters().Single(); + object constructorArgument = parameter.ParameterType == typeof(SqlException) + ? exception + : exception.Errors; + SqlInfoMessageEventArgs eventArgs = (SqlInfoMessageEventArgs)constructor.Invoke(new[] { constructorArgument }); + + SqlInfoMessageEventHandler? handler = GetInstanceFields(typeof(SqlConnection)) + .Where(field => typeof(Delegate).IsAssignableFrom(field.FieldType)) + .Select(field => field.GetValue(connection)) + .OfType() + .FirstOrDefault(); + + Assert.IsNotNull(handler, "The SqlConnection InfoMessage handler was not registered."); + handler(connection, eventArgs); + } + + private static IEnumerable GetInstanceFields(Type type) + { + for (Type? current = type; current is not null; current = current.BaseType) + { + foreach (FieldInfo field in current.GetFields( + BindingFlags.Instance | BindingFlags.Public | BindingFlags.NonPublic | BindingFlags.DeclaredOnly)) + { + yield return field; + } + } + } + [TestMethod] public void CreateConnection_MissingDataSourceThrows() { @@ -1073,6 +1106,27 @@ public void CreateConnection_MissingDataSourceThrows() Assert.AreEqual(DataApiBuilderException.SubStatusCodes.DataSourceNotFound, exception.SubStatusCode); } + [TestMethod] + public void CreateConnection_InfoMessageHandlerCapturesKnownCodeAndHandlesContextFailure() + { + DefaultHttpContext context = new(); + Mock accessor = new(); + accessor.Setup(x => x.HttpContext).Returns(context); + (MsSqlQueryExecutor executor, RuntimeConfigProvider provider) = CreateQueryExecutorForPoolingTest( + "Server=localhost;Database=test;", false, accessor); + using SqlConnection connection = executor.CreateConnection(provider.GetConfig().DefaultDataSourceName); + + InvokeInfoMessageHandler(connection, 15806, "statement-id"); + + Assert.IsTrue(context.Items.ContainsKey("QueryIdentifyingIds")); + + Mock failingContext = new(); + failingContext.SetupGet(x => x.Items).Throws(new InvalidOperationException("Unavailable items collection.")); + accessor.Setup(x => x.HttpContext).Returns(failingContext.Object); + + InvokeInfoMessageHandler(connection, 15806, "second-statement-id"); + } + [TestMethod] public async Task SetManagedIdentityAccessToken_OboBearerTokenSetsConnectionToken() { @@ -1115,6 +1169,59 @@ public async Task SetManagedIdentityAccessToken_OboMissingBearerTokenThrows(stri Assert.AreEqual(DataApiBuilderException.SubStatusCodes.OboAuthenticationFailure, exception.SubStatusCode); } + [TestMethod] + public async Task SetManagedIdentityAccessToken_OboWithoutTokenProviderThrows() + { + DefaultHttpContext context = new(); + context.Request.Headers.Authorization = "Bearer incoming-token"; + Mock accessor = new(); + accessor.Setup(x => x.HttpContext).Returns(context); + (MsSqlQueryExecutor executor, RuntimeConfigProvider configProvider) = CreateQueryExecutorForPoolingTest( + "Server=localhost;Database=test;", true, accessor); + using SqlConnection connection = new(); + + DataApiBuilderException exception = await Assert.ThrowsExceptionAsync(() => + executor.SetManagedIdentityAccessTokenIfAnyAsync( + connection, configProvider.GetConfig().DefaultDataSourceName)); + + Assert.AreEqual(DataApiBuilderException.SubStatusCodes.OboAuthenticationFailure, exception.SubStatusCode); + } + + [TestMethod] + public async Task SetManagedIdentityAccessToken_OboWithoutHttpContextUsesConfiguredAuthentication() + { + Mock accessor = new(); + accessor.Setup(x => x.HttpContext).Returns(value: null); + (MsSqlQueryExecutor executor, RuntimeConfigProvider configProvider) = CreateQueryExecutorForPoolingTest( + "Server=localhost;Database=test;User ID=user;Password=password;", true, accessor, Mock.Of()); + using SqlConnection connection = new(); + + await executor.SetManagedIdentityAccessTokenIfAnyAsync( + connection, configProvider.GetConfig().DefaultDataSourceName); + + Assert.IsNull(connection.AccessToken); + } + + [TestMethod] + public async Task SetManagedIdentityAccessToken_UnavailableDefaultCredentialIsIgnored() + { + Mock accessor = new(); + accessor.Setup(x => x.HttpContext).Returns(value: null); + (MsSqlQueryExecutor executor, RuntimeConfigProvider configProvider) = CreateQueryExecutorForPoolingTest( + "Server=localhost;Database=test;", false, accessor); + Mock credential = new(); + credential + .Setup(x => x.GetTokenAsync(It.IsAny(), It.IsAny())) + .Returns(ValueTask.FromException(new CredentialUnavailableException("Credential unavailable."))); + executor.AzureCredential = credential.Object; + using SqlConnection connection = new(); + + await executor.SetManagedIdentityAccessTokenIfAnyAsync( + connection, configProvider.GetConfig().DefaultDataSourceName); + + Assert.IsNull(connection.AccessToken); + } + #endregion /// From e33548fdaee64da423f4f4e60f9d6259139a6d5b Mon Sep 17 00:00:00 2001 From: Aaron Burtle Date: Sat, 29 Aug 2026 05:58:00 -0700 Subject: [PATCH 03/19] expand PostgreSQL pipeline coverage --- .../PostgreSqlDbExceptionParserHelperTests.cs | 53 +++++++++ .../PostgreSqlMetadataProviderHelperTests.cs | 101 ++++++++++++++++++ .../PostgreSqlQueryExecutorHelperTests.cs | 86 +++++++++++++++ .../PostgresQueryBuilderHelperTests.cs | 51 +++++++++ 4 files changed, 291 insertions(+) create mode 100644 src/Service.Tests/UnitTests/PostgreSqlDbExceptionParserHelperTests.cs create mode 100644 src/Service.Tests/UnitTests/PostgreSqlMetadataProviderHelperTests.cs create mode 100644 src/Service.Tests/UnitTests/PostgreSqlQueryExecutorHelperTests.cs create mode 100644 src/Service.Tests/UnitTests/PostgresQueryBuilderHelperTests.cs diff --git a/src/Service.Tests/UnitTests/PostgreSqlDbExceptionParserHelperTests.cs b/src/Service.Tests/UnitTests/PostgreSqlDbExceptionParserHelperTests.cs new file mode 100644 index 0000000000..4c40beb91a --- /dev/null +++ b/src/Service.Tests/UnitTests/PostgreSqlDbExceptionParserHelperTests.cs @@ -0,0 +1,53 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +using System.Collections.Generic; +using System.Data.Common; +using System.Net; +using Azure.DataApiBuilder.Config.ObjectModel; +using Azure.DataApiBuilder.Core.Configurations; +using Azure.DataApiBuilder.Core.Resolvers; +using Microsoft.VisualStudio.TestTools.UnitTesting; +using Moq; + +namespace Azure.DataApiBuilder.Service.Tests.UnitTests +{ + [TestClass, TestCategory(TestCategory.POSTGRESQL)] + public class PostgreSqlDbExceptionParserHelperTests + { + [DataTestMethod] + [DataRow("08006", true)] + [DataRow("not-transient", false)] + [DataRow(null, false)] + public void IsTransientException_UsesPostgreSqlState(string? sqlState, bool expected) + { + PostgreSqlDbExceptionParser parser = CreateParser(); + Mock exception = new(); + exception.SetupGet(x => x.SqlState).Returns(sqlState); + + Assert.AreEqual(expected, parser.IsTransientException(exception.Object)); + } + + [DataTestMethod] + [DataRow(null)] + [DataRow("unknown")] + public void GetHttpStatusCodeForException_UnrecognizedStateReturnsInternalServerError(string? sqlState) + { + PostgreSqlDbExceptionParser parser = CreateParser(); + Mock exception = new(); + exception.SetupGet(x => x.SqlState).Returns(sqlState); + + Assert.AreEqual(HttpStatusCode.InternalServerError, parser.GetHttpStatusCodeForException(exception.Object)); + } + + private static PostgreSqlDbExceptionParser CreateParser() + { + RuntimeConfig runtimeConfig = new( + Schema: string.Empty, + DataSource: new DataSource(DatabaseType.PostgreSQL, string.Empty), + Entities: new RuntimeEntities(new Dictionary())); + RuntimeConfigProvider configProvider = TestHelper.GenerateInMemoryRuntimeConfigProvider(runtimeConfig); + return new PostgreSqlDbExceptionParser(configProvider); + } + } +} diff --git a/src/Service.Tests/UnitTests/PostgreSqlMetadataProviderHelperTests.cs b/src/Service.Tests/UnitTests/PostgreSqlMetadataProviderHelperTests.cs new file mode 100644 index 0000000000..78237ee161 --- /dev/null +++ b/src/Service.Tests/UnitTests/PostgreSqlMetadataProviderHelperTests.cs @@ -0,0 +1,101 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +using System; +using System.Data; +using System.Net; +using System.Reflection; +using System.Runtime.CompilerServices; +using Azure.DataApiBuilder.Config.DatabasePrimitives; +using Azure.DataApiBuilder.Config.ObjectModel; +using Azure.DataApiBuilder.Core.Models; +using Azure.DataApiBuilder.Core.Services; +using Azure.DataApiBuilder.Service.Exceptions; +using Microsoft.VisualStudio.TestTools.UnitTesting; + +namespace Azure.DataApiBuilder.Service.Tests.UnitTests +{ + [TestClass, TestCategory(TestCategory.POSTGRESQL)] + public class PostgreSqlMetadataProviderHelperTests + { + [TestMethod] + public void ParseSchemaAndDbTableName_UsesSearchPathFromConnectionString() + { + PostgreSqlMetadataProvider provider = CreateProvider(); + SetBaseField(provider, "k__BackingField", "Host=localhost;Database=db;SearchPath=tenant"); + + Assert.AreEqual(("tenant", "books"), provider.ParseSchemaAndDbTableName("books")); + } + + [TestMethod] + public void TryGetSchemaFromConnectionString_InvalidConnectionStringThrowsInitializationError() + { + DataApiBuilderException exception = Assert.ThrowsException(() => + PostgreSqlMetadataProvider.TryGetSchemaFromConnectionString("Host=localhost;Invalid Keyword=value", out _)); + + Assert.AreEqual(HttpStatusCode.ServiceUnavailable, exception.StatusCode); + Assert.AreEqual(DataApiBuilderException.SubStatusCodes.ErrorInInitialization, exception.SubStatusCode); + } + + [TestMethod] + public void SqlToCLRType_IsNotImplemented() + { + PostgreSqlMetadataProvider provider = CreateProvider(); + + Assert.ThrowsException(() => provider.SqlToCLRType("integer")); + } + + [TestMethod] + public void PopulateColumnDefinitionWithHasDefaultAndDbType_MapsArrayUdtMetadata() + { + SourceDefinition definition = new(); + definition.Columns["numbers"] = new ColumnDefinition(typeof(Array)); + + DataTable columns = new(); + columns.Columns.Add("COLUMN_NAME", typeof(string)); + columns.Columns.Add("COLUMN_DEFAULT", typeof(object)); + columns.Columns.Add("DATA_TYPE", typeof(string)); + columns.Columns.Add("UDT_NAME", typeof(string)); + columns.Rows.Add("numbers", DBNull.Value, "ARRAY", "_int4"); + columns.Rows.Add("not_configured", DBNull.Value, "ARRAY", "_text"); + + PostgreSqlMetadataProvider provider = CreateProvider(); + typeof(PostgreSqlMetadataProvider).GetMethod( + "PopulateColumnDefinitionWithHasDefaultAndDbType", + BindingFlags.Instance | BindingFlags.NonPublic)!.Invoke(provider, new object[] { definition, columns }); + + ColumnDefinition numbers = definition.Columns["numbers"]; + Assert.IsFalse(numbers.HasDefault); + Assert.IsTrue(numbers.IsArrayType); + Assert.IsTrue(numbers.IsReadOnly); + Assert.AreEqual(typeof(int), numbers.ElementSystemType); + Assert.AreEqual(typeof(int[]), numbers.SystemType); + } + + private static PostgreSqlMetadataProvider CreateProvider() + { + PostgreSqlMetadataProvider provider = + (PostgreSqlMetadataProvider)RuntimeHelpers.GetUninitializedObject(typeof(PostgreSqlMetadataProvider)); + SetBaseField(provider, "_databaseType", DatabaseType.PostgreSQL); + return provider; + } + + private static void SetBaseField(object instance, string fieldName, object value) + { + Type? type = instance.GetType(); + while (type is not null) + { + FieldInfo? field = type.GetField(fieldName, BindingFlags.Instance | BindingFlags.NonPublic); + if (field is not null) + { + field.SetValue(instance, value); + return; + } + + type = type.BaseType; + } + + Assert.Fail($"Unable to find field '{fieldName}'."); + } + } +} diff --git a/src/Service.Tests/UnitTests/PostgreSqlQueryExecutorHelperTests.cs b/src/Service.Tests/UnitTests/PostgreSqlQueryExecutorHelperTests.cs new file mode 100644 index 0000000000..901237b748 --- /dev/null +++ b/src/Service.Tests/UnitTests/PostgreSqlQueryExecutorHelperTests.cs @@ -0,0 +1,86 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +using System.Collections.Generic; +using System.Data; +using System.Data.Common; +using System.Net; +using System.Threading; +using System.Threading.Tasks; +using Azure.DataApiBuilder.Config.ObjectModel; +using Azure.DataApiBuilder.Core.Configurations; +using Azure.DataApiBuilder.Core.Resolvers; +using Azure.DataApiBuilder.Service.Exceptions; +using Microsoft.AspNetCore.Http; +using Microsoft.Extensions.Logging.Abstractions; +using Microsoft.VisualStudio.TestTools.UnitTesting; +using Moq; + +namespace Azure.DataApiBuilder.Service.Tests.UnitTests +{ + [TestClass, TestCategory(TestCategory.POSTGRESQL)] + public class PostgreSqlQueryExecutorHelperTests + { + [TestMethod] + public async Task GetMultipleResultSets_MissingCountMetadataThrowsInternalServerError() + { + PostgreSqlQueryExecutor executor = CreateExecutor(); + Mock reader = new(); + reader.Setup(x => x.ReadAsync(It.IsAny())).ReturnsAsync(false); + + DataApiBuilderException exception = await Assert.ThrowsExceptionAsync(() => + executor.GetMultipleResultSetsIfAnyAsync(reader.Object)); + + Assert.AreEqual(HttpStatusCode.InternalServerError, exception.StatusCode); + Assert.AreEqual(DataApiBuilderException.SubStatusCodes.UnexpectedError, exception.SubStatusCode); + } + + [TestMethod] + public async Task GetMultipleResultSets_FallbackUpdateWithoutArgumentsThrowsInternalServerError() + { + PostgreSqlQueryExecutor executor = CreateExecutor(); + DataTable schema = new(); + schema.Columns.Add("ColumnName", typeof(string)); + schema.Columns.Add("ColumnSize", typeof(int)); + schema.Rows.Add(PostgresQueryBuilder.COUNT_ROWS_WITH_GIVEN_PK, 8); + schema.Rows.Add(PostgresQueryBuilder.IS_FALLBACK_TO_UPDATE, 1); + + Mock reader = new(); + reader.SetupGet(x => x.RecordsAffected).Returns(0); + reader.SetupGet(x => x.HasRows).Returns(true); + reader.SetupSequence(x => x.ReadAsync(It.IsAny())) + .ReturnsAsync(true) + .ReturnsAsync(false) + .ReturnsAsync(false); + reader.Setup(x => x.GetSchemaTable()).Returns(schema); + reader.Setup(x => x.GetOrdinal(PostgresQueryBuilder.COUNT_ROWS_WITH_GIVEN_PK)).Returns(0); + reader.Setup(x => x.GetOrdinal(PostgresQueryBuilder.IS_FALLBACK_TO_UPDATE)).Returns(1); + reader.Setup(x => x.IsDBNull(It.IsAny())).Returns(false); + reader.Setup(x => x[PostgresQueryBuilder.COUNT_ROWS_WITH_GIVEN_PK]).Returns(0L); + reader.Setup(x => x[PostgresQueryBuilder.IS_FALLBACK_TO_UPDATE]).Returns(true); + reader.Setup(x => x.NextResultAsync(It.IsAny())).ReturnsAsync(true); + + DataApiBuilderException exception = await Assert.ThrowsExceptionAsync(() => + executor.GetMultipleResultSetsIfAnyAsync(reader.Object)); + + Assert.AreEqual(HttpStatusCode.InternalServerError, exception.StatusCode); + Assert.AreEqual(DataApiBuilderException.SubStatusCodes.UnexpectedError, exception.SubStatusCode); + } + + private static PostgreSqlQueryExecutor CreateExecutor() + { + RuntimeConfig runtimeConfig = new( + Schema: string.Empty, + DataSource: new DataSource( + DatabaseType.PostgreSQL, + "Host=localhost;Database=test;Username=user;Password=password"), + Entities: new RuntimeEntities(new Dictionary())); + RuntimeConfigProvider configProvider = TestHelper.GenerateInMemoryRuntimeConfigProvider(runtimeConfig); + return new PostgreSqlQueryExecutor( + configProvider, + new PostgreSqlDbExceptionParser(configProvider), + NullLogger.Instance, + new HttpContextAccessor()); + } + } +} diff --git a/src/Service.Tests/UnitTests/PostgresQueryBuilderHelperTests.cs b/src/Service.Tests/UnitTests/PostgresQueryBuilderHelperTests.cs new file mode 100644 index 0000000000..ce537ff21e --- /dev/null +++ b/src/Service.Tests/UnitTests/PostgresQueryBuilderHelperTests.cs @@ -0,0 +1,51 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +using System; +using System.Collections.Generic; +using Azure.DataApiBuilder.Core.Resolvers; +using Microsoft.VisualStudio.TestTools.UnitTesting; + +namespace Azure.DataApiBuilder.Service.Tests.UnitTests +{ + [TestClass, TestCategory(TestCategory.POSTGRESQL)] + public class PostgresQueryBuilderHelperTests + { + [TestMethod] + public void BuildExecute_IsNotImplemented() + { + PostgresQueryBuilder builder = new(); + + Assert.ThrowsException(() => builder.Build((SqlExecuteStructure)null!)); + } + + [TestMethod] + public void BuildStoredProcedureResultDetailsQuery_IsNotImplemented() + { + PostgresQueryBuilder builder = new(); + + Assert.ThrowsException(() => + builder.BuildStoredProcedureResultDetailsQuery("get_books")); + } + + [TestMethod] + public void IsInsert_MissingOperationMetadataThrows() + { + Dictionary result = new(); + + Assert.ThrowsException(() => PostgresQueryBuilder.IsInsert(result)); + } + + [TestMethod] + public void IsInsert_InvalidOperationMetadataThrowsAndRemovesMetadata() + { + Dictionary result = new() + { + ["___upsert_op___"] = "invalid" + }; + + Assert.ThrowsException(() => PostgresQueryBuilder.IsInsert(result)); + Assert.IsFalse(result.ContainsKey("___upsert_op___")); + } + } +} From 7c4a34e6d581f9082f89cc1b1f25538539942546 Mon Sep 17 00:00:00 2001 From: Aaron Burtle Date: Sat, 29 Aug 2026 05:58:15 -0700 Subject: [PATCH 04/19] expand MySQL pipeline coverage --- .../MySqlDbExceptionParserHelperTests.cs | 63 +++++++++ .../MySqlMetadataProviderHelperTests.cs | 60 +++++++++ .../UnitTests/MySqlQueryBuilderHelperTests.cs | 30 +++++ .../MySqlQueryExecutorHelperTests.cs | 124 ++++++++++++++++++ 4 files changed, 277 insertions(+) create mode 100644 src/Service.Tests/UnitTests/MySqlDbExceptionParserHelperTests.cs create mode 100644 src/Service.Tests/UnitTests/MySqlMetadataProviderHelperTests.cs create mode 100644 src/Service.Tests/UnitTests/MySqlQueryBuilderHelperTests.cs create mode 100644 src/Service.Tests/UnitTests/MySqlQueryExecutorHelperTests.cs diff --git a/src/Service.Tests/UnitTests/MySqlDbExceptionParserHelperTests.cs b/src/Service.Tests/UnitTests/MySqlDbExceptionParserHelperTests.cs new file mode 100644 index 0000000000..3cbe744f0c --- /dev/null +++ b/src/Service.Tests/UnitTests/MySqlDbExceptionParserHelperTests.cs @@ -0,0 +1,63 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +using System; +using System.Collections.Generic; +using System.Net; +using System.Reflection; +using Azure.DataApiBuilder.Config.ObjectModel; +using Azure.DataApiBuilder.Core.Configurations; +using Azure.DataApiBuilder.Core.Resolvers; +using Microsoft.VisualStudio.TestTools.UnitTesting; +using MySqlConnector; + +namespace Azure.DataApiBuilder.Service.Tests.UnitTests +{ + [TestClass, TestCategory(TestCategory.MYSQL)] + public class MySqlDbExceptionParserHelperTests + { + [DataTestMethod] + [DataRow(1020, true)] + [DataRow(9999, false)] + public void IsTransientException_UsesMySqlErrorNumber(int errorNumber, bool expected) + { + MySqlDbExceptionParser parser = CreateParser(); + + Assert.AreEqual(expected, parser.IsTransientException(CreateException(errorNumber))); + } + + [TestMethod] + public void GetHttpStatusCodeForException_UnknownNumberReturnsInternalServerError() + { + MySqlDbExceptionParser parser = CreateParser(); + + Assert.AreEqual( + HttpStatusCode.InternalServerError, + parser.GetHttpStatusCodeForException(CreateException(9999))); + } + + private static MySqlException CreateException(int errorNumber) + { + ConstructorInfo constructor = typeof(MySqlException).GetConstructor( + BindingFlags.Instance | BindingFlags.NonPublic, + binder: null, + new[] { typeof(MySqlErrorCode), typeof(string) }, + modifiers: null)!; + return (MySqlException)constructor.Invoke(new object[] + { + (MySqlErrorCode)errorNumber, + "Test MySQL exception." + }); + } + + private static MySqlDbExceptionParser CreateParser() + { + RuntimeConfig runtimeConfig = new( + Schema: string.Empty, + DataSource: new DataSource(DatabaseType.MySQL, string.Empty), + Entities: new RuntimeEntities(new Dictionary())); + RuntimeConfigProvider configProvider = TestHelper.GenerateInMemoryRuntimeConfigProvider(runtimeConfig); + return new MySqlDbExceptionParser(configProvider); + } + } +} diff --git a/src/Service.Tests/UnitTests/MySqlMetadataProviderHelperTests.cs b/src/Service.Tests/UnitTests/MySqlMetadataProviderHelperTests.cs new file mode 100644 index 0000000000..6b43b91ec8 --- /dev/null +++ b/src/Service.Tests/UnitTests/MySqlMetadataProviderHelperTests.cs @@ -0,0 +1,60 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +using System; +using System.Reflection; +using System.Runtime.CompilerServices; +using Azure.DataApiBuilder.Config.ObjectModel; +using Azure.DataApiBuilder.Core.Services; +using Azure.DataApiBuilder.Service.Exceptions; +using Microsoft.VisualStudio.TestTools.UnitTesting; + +namespace Azure.DataApiBuilder.Service.Tests.UnitTests +{ + [TestClass, TestCategory(TestCategory.MYSQL)] + public class MySqlMetadataProviderHelperTests + { + [TestMethod] + public void ParseSchemaAndDbTableName_RejectsExplicitSchema() + { + MySqlMetadataProvider provider = CreateProvider(); + + Assert.ThrowsException( + () => provider.ParseSchemaAndDbTableName("custom.books")); + } + + [TestMethod] + public void SqlToCLRType_IsNotImplemented() + { + MySqlMetadataProvider provider = CreateProvider(); + + Assert.ThrowsException(() => provider.SqlToCLRType("int")); + } + + private static MySqlMetadataProvider CreateProvider() + { + MySqlMetadataProvider provider = + (MySqlMetadataProvider)RuntimeHelpers.GetUninitializedObject(typeof(MySqlMetadataProvider)); + SetBaseField(provider, "_databaseType", DatabaseType.MySQL); + return provider; + } + + private static void SetBaseField(object instance, string fieldName, object value) + { + Type? type = instance.GetType(); + while (type is not null) + { + FieldInfo? field = type.GetField(fieldName, BindingFlags.Instance | BindingFlags.NonPublic); + if (field is not null) + { + field.SetValue(instance, value); + return; + } + + type = type.BaseType; + } + + Assert.Fail($"Unable to find field '{fieldName}'."); + } + } +} diff --git a/src/Service.Tests/UnitTests/MySqlQueryBuilderHelperTests.cs b/src/Service.Tests/UnitTests/MySqlQueryBuilderHelperTests.cs new file mode 100644 index 0000000000..a14f621a40 --- /dev/null +++ b/src/Service.Tests/UnitTests/MySqlQueryBuilderHelperTests.cs @@ -0,0 +1,30 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +using System; +using Azure.DataApiBuilder.Core.Resolvers; +using Microsoft.VisualStudio.TestTools.UnitTesting; + +namespace Azure.DataApiBuilder.Service.Tests.UnitTests +{ + [TestClass, TestCategory(TestCategory.MYSQL)] + public class MySqlQueryBuilderHelperTests + { + [TestMethod] + public void BuildExecute_IsNotImplemented() + { + MySqlQueryBuilder builder = new(); + + Assert.ThrowsException(() => builder.Build((SqlExecuteStructure)null!)); + } + + [TestMethod] + public void BuildStoredProcedureResultDetailsQuery_IsNotImplemented() + { + MySqlQueryBuilder builder = new(); + + Assert.ThrowsException(() => + builder.BuildStoredProcedureResultDetailsQuery("get_books")); + } + } +} diff --git a/src/Service.Tests/UnitTests/MySqlQueryExecutorHelperTests.cs b/src/Service.Tests/UnitTests/MySqlQueryExecutorHelperTests.cs new file mode 100644 index 0000000000..6e886530da --- /dev/null +++ b/src/Service.Tests/UnitTests/MySqlQueryExecutorHelperTests.cs @@ -0,0 +1,124 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +using System.Collections.Generic; +using System.Data; +using System.Data.Common; +using System.Net; +using System.Threading; +using System.Threading.Tasks; +using Azure.Core; +using Azure.DataApiBuilder.Config.ObjectModel; +using Azure.DataApiBuilder.Core.Configurations; +using Azure.DataApiBuilder.Core.Resolvers; +using Azure.DataApiBuilder.Service.Exceptions; +using Azure.Identity; +using Microsoft.AspNetCore.Http; +using Microsoft.Extensions.Logging.Abstractions; +using Microsoft.VisualStudio.TestTools.UnitTesting; +using Moq; +using MySqlConnector; + +namespace Azure.DataApiBuilder.Service.Tests.UnitTests +{ + [TestClass, TestCategory(TestCategory.MYSQL)] + public class MySqlQueryExecutorHelperTests + { + [TestMethod] + public async Task SetManagedIdentityAccessToken_UnavailableDefaultCredentialIsIgnored() + { + const string connectionString = "Server=localhost;Database=test;User ID=user;"; + MySqlQueryExecutor executor = CreateExecutor(connectionString); + Mock credential = new(); + credential + .Setup(x => x.GetTokenAsync(It.IsAny(), It.IsAny())) + .Returns(ValueTask.FromException(new CredentialUnavailableException("Credential unavailable."))); + executor.AzureCredential = credential.Object; + using MySqlConnection connection = new(connectionString); + + await executor.SetManagedIdentityAccessTokenIfAnyAsync(connection, string.Empty); + + Assert.AreEqual(string.Empty, new MySqlConnectionStringBuilder(connection.ConnectionString).Password); + } + + [TestMethod] + public async Task GetMultipleResultSets_MissingExistenceMetadataThrowsInternalServerError() + { + MySqlQueryExecutor executor = CreateExecutor(); + Mock reader = new(); + reader.Setup(x => x.ReadAsync(It.IsAny())).ReturnsAsync(false); + + DataApiBuilderException exception = await Assert.ThrowsExceptionAsync(() => + executor.GetMultipleResultSetsIfAnyAsync(reader.Object)); + + Assert.AreEqual(HttpStatusCode.InternalServerError, exception.StatusCode); + Assert.AreEqual(DataApiBuilderException.SubStatusCodes.UnexpectedError, exception.SubStatusCode); + } + + [TestMethod] + public async Task GetMultipleResultSets_UpdateOnlyWithoutArgumentsThrowsInternalServerError() + { + MySqlQueryExecutor executor = CreateExecutor(); + Mock reader = CreateInsertPathReader(includeInsertResultSet: false); + + DataApiBuilderException exception = await Assert.ThrowsExceptionAsync(() => + executor.GetMultipleResultSetsIfAnyAsync(reader.Object)); + + Assert.AreEqual(HttpStatusCode.InternalServerError, exception.StatusCode); + Assert.AreEqual(DataApiBuilderException.SubStatusCodes.UnexpectedError, exception.SubStatusCode); + } + + [TestMethod] + public async Task GetMultipleResultSets_EmptyInsertResultThrowsInternalServerError() + { + MySqlQueryExecutor executor = CreateExecutor(); + Mock reader = CreateInsertPathReader(includeInsertResultSet: true); + + DataApiBuilderException exception = await Assert.ThrowsExceptionAsync(() => + executor.GetMultipleResultSetsIfAnyAsync(reader.Object)); + + Assert.AreEqual(HttpStatusCode.InternalServerError, exception.StatusCode); + Assert.AreEqual(DataApiBuilderException.SubStatusCodes.UnexpectedError, exception.SubStatusCode); + } + + private static Mock CreateInsertPathReader(bool includeInsertResultSet) + { + DataTable schema = new(); + schema.Columns.Add("ColumnName", typeof(string)); + schema.Columns.Add("ColumnSize", typeof(int)); + schema.Rows.Add(MySqlQueryBuilder.ROW_EXISTED_BEFORE_UPSERT, 4); + + Mock reader = new(); + reader.SetupGet(x => x.RecordsAffected).Returns(0); + reader.SetupGet(x => x.HasRows).Returns(true); + reader.SetupSequence(x => x.ReadAsync(It.IsAny())) + .ReturnsAsync(true) + .ReturnsAsync(false) + .ReturnsAsync(false) + .ReturnsAsync(false); + reader.Setup(x => x.GetSchemaTable()).Returns(schema); + reader.Setup(x => x.GetOrdinal(MySqlQueryBuilder.ROW_EXISTED_BEFORE_UPSERT)).Returns(0); + reader.Setup(x => x.IsDBNull(0)).Returns(false); + reader.Setup(x => x[MySqlQueryBuilder.ROW_EXISTED_BEFORE_UPSERT]).Returns(0); + reader.SetupSequence(x => x.NextResultAsync(It.IsAny())) + .ReturnsAsync(true) + .ReturnsAsync(includeInsertResultSet); + return reader; + } + + private static MySqlQueryExecutor CreateExecutor( + string connectionString = "Server=localhost;Database=test;User ID=user;Password=password;") + { + RuntimeConfig runtimeConfig = new( + Schema: string.Empty, + DataSource: new DataSource(DatabaseType.MySQL, connectionString), + Entities: new RuntimeEntities(new Dictionary())); + RuntimeConfigProvider configProvider = TestHelper.GenerateInMemoryRuntimeConfigProvider(runtimeConfig); + return new MySqlQueryExecutor( + configProvider, + new MySqlDbExceptionParser(configProvider), + NullLogger.Instance, + new HttpContextAccessor()); + } + } +} From 3dc6aeca19270ce2c00eb49596248b024f5925c6 Mon Sep 17 00:00:00 2001 From: Aaron Burtle Date: Sat, 29 Aug 2026 05:58:29 -0700 Subject: [PATCH 05/19] expand DWSQL pipeline coverage --- .../UnitTests/DwSqlQueryBuilderHelperTests.cs | 37 ++++++++++++++++++- 1 file changed, 36 insertions(+), 1 deletion(-) diff --git a/src/Service.Tests/UnitTests/DwSqlQueryBuilderHelperTests.cs b/src/Service.Tests/UnitTests/DwSqlQueryBuilderHelperTests.cs index 2152cab39a..588b64851b 100644 --- a/src/Service.Tests/UnitTests/DwSqlQueryBuilderHelperTests.cs +++ b/src/Service.Tests/UnitTests/DwSqlQueryBuilderHelperTests.cs @@ -2,6 +2,7 @@ // Licensed under the MIT License. using System.Collections.Generic; +using System.Linq; using System.Reflection; using System.Runtime.CompilerServices; using Azure.DataApiBuilder.Config.DatabasePrimitives; @@ -14,7 +15,7 @@ namespace Azure.DataApiBuilder.Service.Tests.UnitTests { - [TestClass] + [TestClass, TestCategory(TestCategory.DWSQL)] public class DwSqlQueryBuilderHelperTests { [TestMethod] @@ -95,6 +96,32 @@ public void Build_UnoptimizedSimpleQueryUsesStringAggregation() Assert.IsFalse(query.Contains("FOR JSON PATH")); } + [TestMethod] + public void BuildWithJsonFunc_SubqueryWrapsColumnsAsJsonObject() + { + SqlQueryStructure structure = CreateBuildableStructure(isList: false); + + string query = InvokeInstance( + new DwSqlQueryBuilder(enableNto1JoinOpt: true), + "BuildWithJsonFunc", + structure, + true); + + StringAssert.StartsWith(query, "SELECT JSON_OBJECT('id': [id])"); + StringAssert.Contains(query, "FROM (SELECT TOP 1"); + StringAssert.Contains(query, "AS [table0]"); + } + + [TestMethod] + public void BuildFetchEnabledTriggersQuery_ReturnsTriggerMetadataQuery() + { + string query = new DwSqlQueryBuilder(enableNto1JoinOpt: true).BuildFetchEnabledTriggersQuery(); + + StringAssert.Contains(query, "FROM sys.triggers"); + StringAssert.Contains(query, "ST.parent_id = object_id(@param0 + '.' + @param1)"); + StringAssert.Contains(query, "ST.is_disabled = 0"); + } + private static bool InvokeHasToOneOrNoRelation(SqlQueryStructure? structure, bool isSubQuery) => InvokeStatic("HasToOneOrNoRelation", structure, isSubQuery); @@ -104,6 +131,14 @@ private static T InvokeStatic(string methodName, params object?[] arguments) return (T)method.Invoke(null, arguments)!; } + private static T InvokeInstance(DwSqlQueryBuilder builder, string methodName, params object?[] arguments) + { + MethodInfo method = typeof(DwSqlQueryBuilder) + .GetMethods(BindingFlags.Instance | BindingFlags.NonPublic) + .Single(candidate => candidate.Name == methodName && candidate.GetParameters().Length == arguments.Length); + return (T)method.Invoke(builder, arguments)!; + } + private static SqlQueryStructure CreateStructure(bool isList, params (string Alias, SqlQueryStructure Query)[] joins) { SqlQueryStructure structure = (SqlQueryStructure)RuntimeHelpers.GetUninitializedObject(typeof(SqlQueryStructure)); From ec822f1a87404837b5dd1cd0f9cfdc3e1b63209e Mon Sep 17 00:00:00 2001 From: Aaron Burtle Date: Sat, 29 Aug 2026 05:58:52 -0700 Subject: [PATCH 06/19] expand Cosmos DB pipeline coverage --- .../CosmosClientProviderHelperTests.cs | 161 +++++++++++ .../UnitTests/CosmosEngineHelperTests.cs | 259 ++++++++++++++++- .../CosmosODataVisitorHelperTests.cs | 134 +++++++++ .../CosmosQueryBuilderHelperTests.cs | 107 +++++++ .../CosmosQueryEngineExecutionHelperTests.cs | 268 ++++++++++++++++++ .../CosmosQueryStructureHelperTests.cs | 86 ++++++ .../UnitTests/CosmosSamplerHelperTests.cs | 69 +++++ .../CosmosSqlMetadataProviderHelperTests.cs | 71 ++++- 8 files changed, 1151 insertions(+), 4 deletions(-) create mode 100644 src/Service.Tests/UnitTests/CosmosClientProviderHelperTests.cs create mode 100644 src/Service.Tests/UnitTests/CosmosODataVisitorHelperTests.cs create mode 100644 src/Service.Tests/UnitTests/CosmosQueryBuilderHelperTests.cs create mode 100644 src/Service.Tests/UnitTests/CosmosQueryEngineExecutionHelperTests.cs create mode 100644 src/Service.Tests/UnitTests/CosmosQueryStructureHelperTests.cs create mode 100644 src/Service.Tests/UnitTests/CosmosSamplerHelperTests.cs diff --git a/src/Service.Tests/UnitTests/CosmosClientProviderHelperTests.cs b/src/Service.Tests/UnitTests/CosmosClientProviderHelperTests.cs new file mode 100644 index 0000000000..3388607694 --- /dev/null +++ b/src/Service.Tests/UnitTests/CosmosClientProviderHelperTests.cs @@ -0,0 +1,161 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +using System; +using System.Collections.Generic; +using System.Reflection; +using System.Runtime.CompilerServices; +using System.Threading.Tasks; +using Azure.Core; +using Azure.DataApiBuilder.Config; +using Azure.DataApiBuilder.Config.ObjectModel; +using Azure.DataApiBuilder.Core.Configurations; +using Azure.DataApiBuilder.Core.Resolvers; +using Microsoft.VisualStudio.TestTools.UnitTesting; +using System.IO.Abstractions.TestingHelpers; + +namespace Azure.DataApiBuilder.Service.Tests.UnitTests +{ + [TestClass, TestCategory(TestCategory.COSMOSDBNOSQL)] + public class CosmosClientProviderHelperTests + { + [TestMethod] + public async Task Constructor_ConfigNotLoadedRegistersInitializationHandler() + { + RuntimeConfigProvider runtimeConfigProvider = new( + new FileSystemRuntimeConfigLoader(new MockFileSystem())); + + CosmosClientProvider provider = new(runtimeConfigProvider); + + Assert.AreEqual(0, provider.Clients.Count); + Assert.AreEqual(1, runtimeConfigProvider.RuntimeConfigLoadedHandlers.Count); + + RuntimeConfig configuration = new( + Schema: string.Empty, + DataSource: new DataSource(DatabaseType.MSSQL, "Server=localhost"), + Entities: new RuntimeEntities(new Dictionary())); + bool initialized = await runtimeConfigProvider.RuntimeConfigLoadedHandlers[0](runtimeConfigProvider, configuration); + + Assert.IsTrue(initialized); + Assert.AreEqual(0, provider.Clients.Count); + } + + [TestMethod] + public void InitializeClient_NullConfigurationThrows() + { + CosmosClientProvider provider = CreateUninitializedProvider(); + + TargetInvocationException exception = Assert.ThrowsException(() => + InvokeInitializeClient(provider, null)); + + Assert.IsInstanceOfType(exception.InnerException); + } + + [TestMethod] + public void InitializeClient_NonCosmosConfigurationReturnsWithoutClients() + { + CosmosClientProvider provider = CreateUninitializedProvider(); + RuntimeConfig configuration = new( + Schema: string.Empty, + DataSource: new DataSource(DatabaseType.MSSQL, "Server=localhost"), + Entities: new RuntimeEntities(new Dictionary())); + + InvokeInitializeClient(provider, configuration); + + Assert.AreEqual(0, provider.Clients.Count); + } + + [TestMethod] + public void InitializeClient_CosmosWithoutAccountKeyCreatesCredentialClient() + { + RuntimeConfig configuration = new( + Schema: string.Empty, + DataSource: new DataSource( + DatabaseType.CosmosDB_NoSQL, + "AccountEndpoint=https://localhost:8081/;", + new Dictionary()), + Entities: new RuntimeEntities(new Dictionary())); + CosmosClientProvider provider = CreateUninitializedProvider(); + + InvokeInitializeClient(provider, configuration); + + Assert.IsTrue(provider.Clients.ContainsKey(configuration.DefaultDataSourceName)); + provider.Clients[configuration.DefaultDataSourceName]!.Dispose(); + } + + [TestMethod] + public void Constructor_LoadedCosmosConfigurationInitializesImmediately() + { + RuntimeConfig configuration = new( + Schema: string.Empty, + DataSource: new DataSource( + DatabaseType.CosmosDB_NoSQL, + "AccountEndpoint=https://localhost:8081/;", + new Dictionary()), + Entities: new RuntimeEntities(new Dictionary())); + RuntimeConfigProvider runtimeConfigProvider = TestHelper.GenerateInMemoryRuntimeConfigProvider(configuration); + Assert.AreEqual(DatabaseType.CosmosDB_NoSQL, runtimeConfigProvider.GetConfig().DataSource.DatabaseType); + + CosmosClientProvider provider = new(runtimeConfigProvider); + + Assert.AreEqual(0, runtimeConfigProvider.RuntimeConfigLoadedHandlers.Count); + Assert.AreEqual(1, provider.Clients.Count); + foreach (Microsoft.Azure.Cosmos.CosmosClient? client in provider.Clients.Values) + { + client?.Dispose(); + } + } + + [DataTestMethod] + [DataRow("AccountEndpoint=https://localhost:8081/;AccountKey=secret", "https://localhost:8081/", "secret")] + [DataRow("ApplicationName=dab", null, null)] + public void ParseCosmosConnectionString_ReturnsAvailableComponents( + string connectionString, + string? expectedEndpoint, + string? expectedKey) + { + MethodInfo method = typeof(CosmosClientProvider).GetMethod( + "ParseCosmosConnectionString", + BindingFlags.Static | BindingFlags.NonPublic)!; + + (string? endpoint, string? key) = ((string?, string?))method.Invoke(null, new object[] { connectionString })!; + + Assert.AreEqual(expectedEndpoint, endpoint); + Assert.AreEqual(expectedKey, key); + } + + [TestMethod] + public void AadTokenCredential_InvalidTokenThrows() + { + Type credentialType = typeof(CosmosClientProvider).GetNestedType( + "AADTokenCredential", + BindingFlags.NonPublic)!; + TokenCredential credential = (TokenCredential)Activator.CreateInstance( + credentialType, + BindingFlags.Instance | BindingFlags.Public | BindingFlags.NonPublic, + binder: null, + args: new object[] { "not-a-jwt" }, + culture: null)!; + + Assert.ThrowsException(() => + credential.GetToken(new TokenRequestContext(Array.Empty()), default)); + } + + private static CosmosClientProvider CreateUninitializedProvider() + { + CosmosClientProvider provider = + (CosmosClientProvider)RuntimeHelpers.GetUninitializedObject(typeof(CosmosClientProvider)); + typeof(CosmosClientProvider).GetField("k__BackingField", BindingFlags.Instance | BindingFlags.NonPublic)! + .SetValue(provider, new Dictionary()); + typeof(CosmosClientProvider).GetField("_accessToken", BindingFlags.Instance | BindingFlags.NonPublic)! + .SetValue(provider, new Dictionary()); + return provider; + } + + private static void InvokeInitializeClient(CosmosClientProvider provider, RuntimeConfig? configuration) + { + typeof(CosmosClientProvider).GetMethod("InitializeClient", BindingFlags.Instance | BindingFlags.NonPublic)! + .Invoke(provider, new object?[] { configuration }); + } + } +} diff --git a/src/Service.Tests/UnitTests/CosmosEngineHelperTests.cs b/src/Service.Tests/UnitTests/CosmosEngineHelperTests.cs index f87ab0b986..ed0e079d32 100644 --- a/src/Service.Tests/UnitTests/CosmosEngineHelperTests.cs +++ b/src/Service.Tests/UnitTests/CosmosEngineHelperTests.cs @@ -3,8 +3,11 @@ using System; using System.Collections.Generic; +using System.IO; using System.Linq; using System.Reflection; +using System.Runtime.CompilerServices; +using System.Threading; using System.Threading.Tasks; using Azure.DataApiBuilder.Auth; using Azure.DataApiBuilder.Config.ObjectModel; @@ -16,13 +19,16 @@ using HotChocolate.Language; using HotChocolate.Resolvers; using Microsoft.Extensions.Primitives; +using Microsoft.Azure.Cosmos; using Microsoft.VisualStudio.TestTools.UnitTesting; using Moq; +using Azure.DataApiBuilder.Service.GraphQLBuilder.Mutations; +using Azure.DataApiBuilder.Service.GraphQLBuilder.Queries; using Newtonsoft.Json.Linq; namespace Azure.DataApiBuilder.Service.Tests.UnitTests { - [TestClass] + [TestClass, TestCategory(TestCategory.COSMOSDBNOSQL)] public class CosmosEngineHelperTests { [DataTestMethod] @@ -80,7 +86,13 @@ public void AuthorizeMutation_UnsupportedOperationThrows() CosmosMutationEngine engine = new(null!, null!, new Mock().Object); Assert.ThrowsException(() => engine.AuthorizeMutation( - CreateContext(), new Dictionary(), "Book", EntityActionOperation.Read)); + CreateContext(), + new Dictionary + { + [MutationBuilder.ITEM_INPUT_ARGUMENT_NAME] = new List() + }, + "Book", + EntityActionOperation.Read)); } [TestMethod] @@ -126,6 +138,24 @@ public void ParseInlineInputItem_HandlesObjectListArrayAndPrimitiveValues() Assert.AreEqual(9, InvokeMutation("ParseInlineInputItem", 9)); } + [TestMethod] + public void ParseInlineInputItem_HandlesNestedObjectValue() + { + Mock nested = new(); + nested.SetupGet(x => x.Kind).Returns(SyntaxKind.ObjectValue); + nested.SetupGet(x => x.Value).Returns(new List { new("id", 7) }); + JObject result = (JObject)InvokeMutation( + "ParseInlineInputItem", + new List { new("nested", nested.Object) })!; + + Assert.AreEqual(7, result["nested"]!["id"]!.Value()); + + JObject directResult = (JObject)InvokeMutation( + "ParseInlineInputItem", + new ObjectFieldNode("nested", nested.Object))!; + Assert.AreEqual(7, directResult["nested"]!["id"]!.Value()); + } + [TestMethod] public void GeneratePatchOperations_CreatesLeafAndArrayOperations() { @@ -151,13 +181,165 @@ public void Base64Helpers_RoundTrip(string? plain, string? encoded) public async Task UnsupportedCosmosEngineEntryPointsThrow() { CosmosMutationEngine mutation = new(null!, null!, new Mock().Object); - CosmosQueryEngine query = (CosmosQueryEngine)System.Runtime.CompilerServices.RuntimeHelpers.GetUninitializedObject(typeof(CosmosQueryEngine)); + CosmosQueryEngine query = (CosmosQueryEngine)RuntimeHelpers.GetUninitializedObject(typeof(CosmosQueryEngine)); await Assert.ThrowsExceptionAsync(() => mutation.ExecuteAsync((RestRequestContext)null!)); + await Assert.ThrowsExceptionAsync(() => mutation.ExecuteAsync((StoredProcedureRequestContext)null!, string.Empty)); await Assert.ThrowsExceptionAsync(() => query.ExecuteAsync((FindRequestContext)null!)); await Assert.ThrowsExceptionAsync(() => query.ExecuteAsync((StoredProcedureRequestContext)null!, string.Empty)); } + [TestMethod] + public async Task ExecuteMutation_NullArgumentsAndMissingClientThrow() + { + CosmosClientProvider clientProvider = CreateClientProvider(new Dictionary + { + ["cosmos"] = null + }); + CosmosMutationEngine engine = new(clientProvider, null!, Mock.Of()); + CosmosOperationMetadata operation = new("db", "container", EntityActionOperation.Create); + + await Assert.ThrowsExceptionAsync(() => + InvokeMutationAsync(engine, null!, operation, "cosmos")); + DataApiBuilderException exception = await Assert.ThrowsExceptionAsync(() => + InvokeMutationAsync(engine, new Dictionary(), operation, "cosmos")); + Assert.AreEqual(DataApiBuilderException.SubStatusCodes.DatabaseOperationFailed, exception.SubStatusCode); + } + + [TestMethod] + public async Task MutationHandlers_ValidateRequiredArguments() + { + Container container = Mock.Of(); + + await AssertPrivateThrowsAsync( + "HandleDeleteAsync", + new Dictionary(), + container); + await AssertPrivateThrowsAsync( + "HandleDeleteAsync", + new Dictionary { [QueryBuilder.ID_FIELD_NAME] = "1" }, + container); + await AssertPrivateThrowsAsync( + "HandleUpdateAsync", + new Dictionary(), + container); + await AssertPrivateThrowsAsync( + "HandleUpdateAsync", + new Dictionary { [QueryBuilder.ID_FIELD_NAME] = "1" }, + container); + await AssertPrivateThrowsAsync( + "HandlePatchAsync", + new Dictionary(), + container); + await AssertPrivateThrowsAsync( + "HandlePatchAsync", + new Dictionary { [QueryBuilder.ID_FIELD_NAME] = "1" }, + container); + } + + [DataTestMethod] + [DataRow("HandleCreateAsync", false)] + [DataRow("HandleUpdateAsync", true)] + [DataRow("HandlePatchAsync", true)] + public async Task MutationHandlers_InvalidInputThrows(string methodName, bool requiresKeys) + { + Dictionary arguments = new() + { + [MutationBuilder.ITEM_INPUT_ARGUMENT_NAME] = null + }; + if (requiresKeys) + { + arguments[QueryBuilder.ID_FIELD_NAME] = "1"; + arguments[QueryBuilder.PARTITION_KEY_FIELD_NAME] = "tenant"; + } + + await AssertPrivateThrowsAsync(methodName, arguments, Mock.Of()); + } + + [DataTestMethod] + [DataRow("HandleCreateAsync", false)] + [DataRow("HandleUpdateAsync", true)] + public async Task MutationHandlers_AcceptVariableInput(string methodName, bool requiresKeys) + { + Dictionary arguments = new() + { + [MutationBuilder.ITEM_INPUT_ARGUMENT_NAME] = new Dictionary { ["title"] = "DAB" } + }; + if (requiresKeys) + { + arguments[QueryBuilder.ID_FIELD_NAME] = "1"; + arguments[QueryBuilder.PARTITION_KEY_FIELD_NAME] = "tenant"; + } + + Mock> response = new(); + response.SetupGet(x => x.Resource).Returns(JObject.Parse(@"{ ""title"": ""DAB"" }")); + Mock container = new(); + container.Setup(x => x.CreateItemAsync( + It.IsAny(), + It.IsAny(), + It.IsAny(), + It.IsAny())) + .ReturnsAsync(response.Object); + container.Setup(x => x.ReplaceItemAsync( + It.IsAny(), + It.IsAny(), + It.IsAny(), + It.IsAny(), + It.IsAny())) + .ReturnsAsync(response.Object); + + await InvokePrivateMutationAsync(methodName, arguments, container.Object); + } + + [TestMethod] + public async Task HandlePatchAsync_VariableInputWithinLimitReturnsResource() + { + Dictionary arguments = CreatePatchArguments(1); + JObject resource = JObject.Parse(@"{ ""id"": ""1"" } "); + Mock> response = new(); + response.SetupGet(x => x.Resource).Returns(resource); + Mock container = new(); + container.Setup(x => x.PatchItemAsync( + "1", + It.IsAny(), + It.IsAny>(), + It.IsAny(), + It.IsAny())) + .ReturnsAsync(response.Object); + + JObject result = (JObject)(await InvokePrivateMutationAsync("HandlePatchAsync", arguments, container.Object))!; + + Assert.AreSame(resource, result); + } + + [TestMethod] + public async Task HandlePatchAsync_FailedTransactionalBatchThrows() + { + Dictionary arguments = CreatePatchArguments(11); + Mock response = new(); + response.SetupGet(x => x.IsSuccessStatusCode).Returns(false); + Mock batch = new(); + batch.Setup(x => x.PatchItem(It.IsAny(), It.IsAny>(), It.IsAny())) + .Returns(batch.Object); + batch.Setup(x => x.ExecuteAsync(It.IsAny())).ReturnsAsync(response.Object); + Mock container = new(); + container.Setup(x => x.CreateTransactionalBatch(It.IsAny())).Returns(batch.Object); + + await Assert.ThrowsExceptionAsync(async () => + await InvokePrivateMutationAsync("HandlePatchAsync", arguments, container.Object)); + } + + [TestMethod] + public void QueryValueHelpers_HandleMissingInputs() + { + Assert.IsNull(InvokeQuery("GetPartitionKeyValue", CreateContext(), null, null)); + List filter = new() + { + new("id", new ObjectValueNode(new ObjectFieldNode("ne", 1))) + }; + Assert.IsNull(InvokeQuery("GetIdValue", CreateContext(), filter)); + } + private static IMiddlewareContext CreateContext() { Mock context = new(); @@ -179,5 +361,76 @@ private static IMiddlewareContext CreateContext() MethodInfo method = typeof(CosmosQueryEngine).GetMethod(methodName, BindingFlags.Static | BindingFlags.NonPublic)!; return method.Invoke(null, args); } + + private static CosmosClientProvider CreateClientProvider(Dictionary clients) + { + CosmosClientProvider provider = + (CosmosClientProvider)RuntimeHelpers.GetUninitializedObject(typeof(CosmosClientProvider)); + typeof(CosmosClientProvider).GetField("k__BackingField", BindingFlags.Instance | BindingFlags.NonPublic)! + .SetValue(provider, clients); + return provider; + } + + private static async Task InvokeMutationAsync( + CosmosMutationEngine engine, + IDictionary arguments, + CosmosOperationMetadata operation, + string dataSourceName) + { + MethodInfo method = typeof(CosmosMutationEngine).GetMethod( + "ExecuteAsync", + BindingFlags.Instance | BindingFlags.NonPublic, + binder: null, + new[] + { + typeof(IMiddlewareContext), + typeof(IDictionary), + typeof(CosmosOperationMetadata), + typeof(string) + }, + modifiers: null)!; + Task task = (Task)method.Invoke(engine, new object?[] + { + CreateContext(), arguments, operation, dataSourceName + })!; + await task; + } + + private static async Task AssertPrivateThrowsAsync( + string methodName, + IDictionary arguments, + Container container) + where TException : Exception + { + MethodInfo method = typeof(CosmosMutationEngine).GetMethod( + methodName, + BindingFlags.Static | BindingFlags.NonPublic)!; + Task task = (Task)method.Invoke(null, new object[] { arguments, container })!; + await Assert.ThrowsExceptionAsync(() => task); + } + + private static Dictionary CreatePatchArguments(int propertyCount) + { + Dictionary item = Enumerable.Range(1, propertyCount) + .ToDictionary(index => $"property{index}", index => (object?)index); + return new Dictionary + { + [QueryBuilder.ID_FIELD_NAME] = "1", + [QueryBuilder.PARTITION_KEY_FIELD_NAME] = "tenant", + [MutationBuilder.ITEM_INPUT_ARGUMENT_NAME] = item + }; + } + + private static async Task InvokePrivateMutationAsync( + string methodName, + IDictionary arguments, + Container container) + { + MethodInfo method = typeof(CosmosMutationEngine).GetMethod( + methodName, + BindingFlags.Static | BindingFlags.NonPublic)!; + dynamic task = method.Invoke(null, new object[] { arguments, container })!; + return await task; + } } } diff --git a/src/Service.Tests/UnitTests/CosmosODataVisitorHelperTests.cs b/src/Service.Tests/UnitTests/CosmosODataVisitorHelperTests.cs new file mode 100644 index 0000000000..05a704012b --- /dev/null +++ b/src/Service.Tests/UnitTests/CosmosODataVisitorHelperTests.cs @@ -0,0 +1,134 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +using System; +using System.Reflection; +using Azure.DataApiBuilder.Auth; +using Azure.DataApiBuilder.Core.Parsers; +using Azure.DataApiBuilder.Core.Resolvers; +using Azure.DataApiBuilder.Core.Services; +using Microsoft.OData.Edm; +using Microsoft.OData.UriParser; +using Microsoft.VisualStudio.TestTools.UnitTesting; +using Moq; + +namespace Azure.DataApiBuilder.Service.Tests.UnitTests +{ + [TestClass, TestCategory(TestCategory.COSMOSDBNOSQL)] + public class CosmosODataVisitorHelperTests + { + [DataTestMethod] + [DataRow(BinaryOperatorKind.Equal, "=")] + [DataRow(BinaryOperatorKind.GreaterThan, ">")] + [DataRow(BinaryOperatorKind.GreaterThanOrEqual, ">=")] + [DataRow(BinaryOperatorKind.LessThan, "<")] + [DataRow(BinaryOperatorKind.LessThanOrEqual, "<=")] + [DataRow(BinaryOperatorKind.NotEqual, "!=")] + [DataRow(BinaryOperatorKind.And, "AND")] + [DataRow(BinaryOperatorKind.Or, "OR")] + public void BinaryOperatorMapping_ReturnsExpectedOperator(BinaryOperatorKind operation, string expected) + { + Assert.AreEqual(expected, InvokeStatic("GetFilterPredicateOperator", operation)); + } + + [TestMethod] + public void BinaryOperatorMapping_UnknownOperationThrows() + { + TargetInvocationException exception = Assert.ThrowsException(() => + InvokeStatic("GetFilterPredicateOperator", (BinaryOperatorKind)int.MaxValue)); + + Assert.IsInstanceOfType(exception.InnerException); + } + + [TestMethod] + public void UnaryOperatorMapping_HandlesNotAndRejectsUnknownOperation() + { + Assert.AreEqual("NOT", InvokeStatic("GetFilterPredicateOperator", UnaryOperatorKind.Not)); + TargetInvocationException exception = Assert.ThrowsException(() => + InvokeStatic("GetFilterPredicateOperator", (UnaryOperatorKind)int.MaxValue)); + Assert.IsInstanceOfType(exception.InnerException); + } + + [DataTestMethod] + [DataRow(BinaryOperatorKind.Equal, "field", "NULL", "(field IS NULL)")] + [DataRow(BinaryOperatorKind.Equal, "NULL", "field", "(field IS NULL)")] + [DataRow(BinaryOperatorKind.NotEqual, "field", "NULL", "(field IS NOT NULL)")] + [DataRow(BinaryOperatorKind.NotEqual, "NULL", "field", "(field IS NOT NULL)")] + [DataRow(BinaryOperatorKind.GreaterThan, "field", "NULL", "(field > NULL)")] + [DataRow(BinaryOperatorKind.GreaterThanOrEqual, "field", "NULL", "(field >= NULL)")] + [DataRow(BinaryOperatorKind.LessThan, "field", "NULL", "(field < NULL)")] + [DataRow(BinaryOperatorKind.LessThanOrEqual, "field", "NULL", "(field <= NULL)")] + public void CreateNullResult_FormatsSupportedOperations( + BinaryOperatorKind operation, + string left, + string right, + string expected) + { + Assert.AreEqual(expected, InvokeStatic("CreateNullResult", operation, left, right)); + } + + [TestMethod] + public void CreateNullResult_RejectsUnsupportedOperation() + { + TargetInvocationException exception = Assert.ThrowsException(() => + InvokeStatic("CreateNullResult", BinaryOperatorKind.And, "field", "NULL")); + + Assert.IsInstanceOfType(exception.InnerException); + } + + [TestMethod] + public void Visitor_HandlesNullBinaryUnaryConvertAndTypedConstants() + { + TestQueryStructure structure = new(); + ODataASTCosmosVisitor visitor = new("c", structure); + ConstantNode nullNode = CreateConstantNode(null!, "null", EdmPrimitiveTypeKind.String, isNull: true); + ConstantNode valueNode = CreateConstantNode(7, "7", EdmPrimitiveTypeKind.Int32); + BinaryOperatorNode binary = new(BinaryOperatorKind.Equal, valueNode, nullNode); + UnaryOperatorNode unary = new(UnaryOperatorKind.Not, CreateConstantNode(true, "true", EdmPrimitiveTypeKind.Boolean)); + EdmPrimitiveTypeReference intType = new(EdmCoreModel.Instance.GetPrimitiveType(EdmPrimitiveTypeKind.Int32), false); + ConvertNode convert = new(valueNode, intType); + + Assert.AreEqual("(@param0 IS NULL)", binary.Accept(visitor)); + Assert.AreEqual("(NOT @param1 )", unary.Accept(visitor)); + Assert.AreEqual("@param2", convert.Accept(visitor)); + Assert.AreEqual("NULL", nullNode.Accept(visitor)); + CollectionAssert.AreEqual(new object[] { 7, true, 7 }, + new System.Collections.Generic.List(System.Linq.Enumerable.Select(structure.Parameters.Values, p => p.Value))); + } + + private static ConstantNode CreateConstantNode( + object value, + string literal, + EdmPrimitiveTypeKind kind, + bool isNull = false) + { + EdmPrimitiveTypeReference? type = isNull + ? null + : new EdmPrimitiveTypeReference(EdmCoreModel.Instance.GetPrimitiveType(kind), false); + return new ConstantNode(value, literal, type); + } + + private static string InvokeStatic(string methodName, params object[] arguments) + { + Type[] parameterTypes = Array.ConvertAll(arguments, argument => argument.GetType()); + MethodInfo method = typeof(ODataASTCosmosVisitor).GetMethod( + methodName, + BindingFlags.Static | BindingFlags.NonPublic, + binder: null, + parameterTypes, + modifiers: null)!; + return (string)method.Invoke(null, arguments)!; + } + + private sealed class TestQueryStructure : BaseQueryStructure + { + public TestQueryStructure() + : base( + Mock.Of(), + Mock.Of(), + gQLFilterParser: null!) + { + } + } + } +} diff --git a/src/Service.Tests/UnitTests/CosmosQueryBuilderHelperTests.cs b/src/Service.Tests/UnitTests/CosmosQueryBuilderHelperTests.cs new file mode 100644 index 0000000000..6bc2d688b9 --- /dev/null +++ b/src/Service.Tests/UnitTests/CosmosQueryBuilderHelperTests.cs @@ -0,0 +1,107 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +using System; +using System.Collections.Generic; +using System.Reflection; +using System.Runtime.CompilerServices; +using Azure.DataApiBuilder.Config.ObjectModel; +using Azure.DataApiBuilder.Core.Models; +using Azure.DataApiBuilder.Core.Resolvers; +using Microsoft.VisualStudio.TestTools.UnitTesting; + +namespace Azure.DataApiBuilder.Service.Tests.UnitTests +{ + [TestClass, TestCategory(TestCategory.COSMOSDBNOSQL)] + public class CosmosQueryBuilderHelperTests + { + [TestMethod] + public void BuildPaginationPredicate_ReturnsEmptyString() + { + TestCosmosQueryBuilder builder = new(); + + Assert.AreEqual(string.Empty, builder.BuildPaginationPredicate(null)); + } + + [TestMethod] + public void BuildPredicate_NullAndUnknownOperationThrow() + { + TestCosmosQueryBuilder builder = new(); + + Assert.ThrowsException(() => builder.BuildPredicate(null)); + Assert.ThrowsException(() => builder.BuildOperation(PredicateOperation.None)); + Assert.ThrowsException(() => builder.BuildOperation(PredicateOperation.IN)); + } + + [TestMethod] + public void ResolveOperand_NullAndEmptyOperandThrow() + { + TestCosmosQueryBuilder builder = new(); + PredicateOperand emptyOperand = (PredicateOperand)RuntimeHelpers.GetUninitializedObject(typeof(PredicateOperand)); + + Assert.ThrowsException(() => builder.Resolve(null)); + Assert.ThrowsException(() => builder.Resolve(emptyOperand)); + } + + [TestMethod] + public void ResolveOperand_CosmosQueryStructureBuildsNestedQuery() + { + TestCosmosQueryBuilder builder = new(); + CosmosQueryStructure structure = CreateStructure(); + + string query = builder.Resolve(new PredicateOperand(structure)); + + Assert.AreEqual("SELECT c.id FROM c", query); + } + + [TestMethod] + public void BuildExistsQueryForCosmos_WithoutPredicatesOmitsWhereClause() + { + Assert.AreEqual( + "EXISTS (SELECT VALUE 1 FROM item IN c.items )", + CosmosQueryBuilder.BuildExistsQueryForCosmos("item IN c.items", null)); + } + + private static CosmosQueryStructure CreateStructure() + { + CosmosQueryStructure structure = + (CosmosQueryStructure)RuntimeHelpers.GetUninitializedObject(typeof(CosmosQueryStructure)); + SetField(structure, "Columns", new List + { + new(string.Empty, "c", "id", "id") + }); + SetField(structure, "Predicates", new List()); + SetField(structure, "DbPolicyPredicatesForOperations", new Dictionary()); + SetField(structure, "OrderByColumns", new List()); + return structure; + } + + private static void SetField(object instance, string propertyName, object value) + { + for (Type? type = instance.GetType(); type is not null; type = type.BaseType) + { + FieldInfo? field = type.GetField( + $"<{propertyName}>k__BackingField", + BindingFlags.Instance | BindingFlags.NonPublic); + if (field is not null) + { + field.SetValue(instance, value); + return; + } + } + + Assert.Fail($"Unable to find backing field for '{propertyName}'."); + } + + private sealed class TestCosmosQueryBuilder : CosmosQueryBuilder + { + public string BuildPaginationPredicate(KeysetPaginationPredicate? predicate) => Build(predicate); + + public string BuildOperation(PredicateOperation operation) => Build(operation); + + public string BuildPredicate(Predicate? predicate) => Build(predicate); + + public string Resolve(PredicateOperand? operand) => ResolveOperand(operand); + } + } +} diff --git a/src/Service.Tests/UnitTests/CosmosQueryEngineExecutionHelperTests.cs b/src/Service.Tests/UnitTests/CosmosQueryEngineExecutionHelperTests.cs new file mode 100644 index 0000000000..f739fc0cc8 --- /dev/null +++ b/src/Service.Tests/UnitTests/CosmosQueryEngineExecutionHelperTests.cs @@ -0,0 +1,268 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +using System; +using System.Collections.Generic; +using System.Linq; +using System.Net; +using System.Reflection; +using System.Runtime.CompilerServices; +using System.Threading; +using System.Threading.Tasks; +using Azure.DataApiBuilder.Auth; +using Azure.DataApiBuilder.Config.ObjectModel; +using Azure.DataApiBuilder.Core.Authorization; +using Azure.DataApiBuilder.Core.Configurations; +using Azure.DataApiBuilder.Core.Models; +using Azure.DataApiBuilder.Core.Resolvers; +using Azure.DataApiBuilder.Core.Services; +using Azure.DataApiBuilder.Core.Services.MetadataProviders; +using Azure.DataApiBuilder.Service.GraphQLBuilder.Queries; +using HotChocolate; +using HotChocolate.Execution; +using HotChocolate.Language; +using HotChocolate.Resolvers; +using Microsoft.AspNetCore.Http; +using Microsoft.Azure.Cosmos; +using Microsoft.VisualStudio.TestTools.UnitTesting; +using Moq; +using Newtonsoft.Json.Linq; + +namespace Azure.DataApiBuilder.Service.Tests.UnitTests +{ + [TestClass, TestCategory(TestCategory.COSMOSDBNOSQL)] + public class CosmosQueryEngineExecutionHelperTests + { + [TestMethod] + public async Task ExecuteListAsync_ReturnsAllItemsFromCosmosPages() + { + JObject first = JObject.Parse(@"{ ""id"": ""1"" }"); + JObject second = JObject.Parse(@"{ ""id"": ""2"" }"); + Mock> page = new(); + page.Setup(x => x.GetEnumerator()).Returns(() => new[] { first, second }.AsEnumerable().GetEnumerator()); + Mock> iterator = new(); + iterator.SetupSequence(x => x.HasMoreResults).Returns(true).Returns(false); + iterator.Setup(x => x.ReadNextAsync(It.IsAny())).ReturnsAsync(page.Object); + Mock container = CreateQueryContainer(iterator.Object); + Mock database = new(); + database.Setup(x => x.GetContainer("books")).Returns(container.Object); + Mock client = new(); + client.Setup(x => x.GetDatabase("db")).Returns(database.Object); + CosmosClientProvider clientProvider = + (CosmosClientProvider)RuntimeHelpers.GetUninitializedObject(typeof(CosmosClientProvider)); + typeof(CosmosClientProvider).GetField("k__BackingField", BindingFlags.Instance | BindingFlags.NonPublic)! + .SetValue(clientProvider, new Dictionary { ["cosmos"] = client.Object }); + + CosmosSqlMetadataProvider metadata = CreateMetadataProvider(); + Mock metadataFactory = new(); + metadataFactory.Setup(x => x.GetMetadataProvider("cosmos")).Returns(metadata); + RuntimeConfig configuration = new( + Schema: string.Empty, + DataSource: new DataSource(DatabaseType.CosmosDB_NoSQL, string.Empty), + Entities: new RuntimeEntities(new Dictionary())); + RuntimeConfigProvider runtimeConfigProvider = TestHelper.GenerateInMemoryRuntimeConfigProvider(configuration); + CosmosQueryEngine engine = new( + clientProvider, + metadataFactory.Object, + Mock.Of(), + new GQLFilterParser(runtimeConfigProvider, metadataFactory.Object), + runtimeConfigProvider, + null!); + + ISchemaBuilder schemaBuilder = SchemaBuilder.New() + .AddDocumentFromString("type Query { books: [Book!]! } type Book { id: String }") + .AddResolver("Book", "id", _ => "id") + .AddResolver("Query", "books", async context => + { + Tuple, IMetadata> result = + await engine.ExecuteListAsync( + (IMiddlewareContext)context, + new Dictionary { ["id"] = "1" }, + "cosmos"); + return result.Item1.Select(document => new + { + id = document.RootElement.GetProperty("id").GetString() + }); + }); + + IOperationRequest request = OperationRequestBuilder.New() + .SetDocument("{ books { id } }") + .SetGlobalState(nameof(HttpContext), CreateHttpContext()) + .Build(); + IExecutionResult executionResult = await schemaBuilder.Create().MakeExecutable().ExecuteAsync(request); + + OperationResult operationResult = executionResult.ExpectOperationResult(); + Assert.AreEqual(0, operationResult.Errors.Count, string.Join(" | ", operationResult.Errors.Select(error => error.ToString()))); + container.Verify(x => x.GetItemQueryIterator( + It.IsAny(), + It.IsAny(), + It.IsAny()), Times.Once); + } + + [TestMethod] + public async Task ExecuteQueryAsync_NoCrossPartitionResultsReturnsNull() + { + Mock> page = new(); + page.SetupGet(x => x.Count).Returns(0); + Mock> iterator = new(); + iterator.Setup(x => x.ReadNextAsync(It.IsAny())).ReturnsAsync(page.Object); + iterator.SetupGet(x => x.HasMoreResults).Returns(false); + Mock container = CreateQueryContainer(iterator.Object); + + JObject? result = await InvokeExecuteQueryAsync( + CreateStructure(), new QueryRequestOptions(), container.Object, string.Empty, string.Empty); + + Assert.IsNull(result); + } + + [TestMethod] + public async Task ExecuteQueryAsync_PaginatedCrossPartitionResultIncludesContinuation() + { + JObject first = JObject.Parse(@"{ ""id"": ""1"" }"); + JObject second = JObject.Parse(@"{ ""id"": ""2"" }"); + Mock> page = new(); + page.Setup(x => x.GetEnumerator()).Returns(() => new[] { first, second }.AsEnumerable().GetEnumerator()); + page.SetupGet(x => x.ContinuationToken).Returns("next-token"); + Mock> iterator = new(); + iterator.Setup(x => x.ReadNextAsync(It.IsAny())).ReturnsAsync(page.Object); + Mock container = CreateQueryContainer(iterator.Object); + CosmosQueryStructure structure = CreateStructure(isPaginated: true); + structure.MaxItemCount = 2; + structure.Continuation = "cHJldmlvdXMtdG9rZW4="; + QueryRequestOptions options = new(); + + JObject result = (await InvokeExecuteQueryAsync( + structure, options, container.Object, string.Empty, string.Empty))!; + + Assert.AreEqual(2, options.MaxItemCount); + Assert.AreEqual("bmV4dC10b2tlbg==", result[QueryBuilder.PAGINATION_TOKEN_FIELD_NAME]!.Value()); + Assert.IsTrue(result[QueryBuilder.HAS_NEXT_PAGE_FIELD_NAME]!.Value()); + Assert.AreEqual(2, ((JArray)result[QueryBuilder.PAGINATION_FIELD_NAME]!).Count); + container.Verify(x => x.GetItemQueryIterator( + It.IsAny(), + "previous-token", + options), Times.Once); + } + + [DataTestMethod] + [DataRow(false)] + [DataRow(true)] + public async Task QueryByIdAndPartitionKey_SuccessReturnsExpectedShape(bool isPaginated) + { + JObject item = JObject.Parse(@"{ ""id"": ""1"" }"); + Mock> response = new(); + response.SetupGet(x => x.Resource).Returns(item); + Mock container = new(); + container.Setup(x => x.ReadItemAsync( + "1", + It.IsAny(), + It.IsAny(), + It.IsAny())) + .ReturnsAsync(response.Object); + + JObject result = await InvokeQueryByIdAndPartitionKey(container.Object, isPaginated); + + if (isPaginated) + { + Assert.IsFalse(result[QueryBuilder.HAS_NEXT_PAGE_FIELD_NAME]!.Value()); + Assert.AreEqual("1", result[QueryBuilder.PAGINATION_FIELD_NAME]![0]!["id"]!.Value()); + } + else + { + Assert.AreSame(item, result); + } + } + + [TestMethod] + public async Task QueryByIdAndPartitionKey_NotFoundReturnsNull() + { + Mock container = new(); + container.Setup(x => x.ReadItemAsync( + It.IsAny(), + It.IsAny(), + It.IsAny(), + It.IsAny())) + .ThrowsAsync(new CosmosException("missing", HttpStatusCode.NotFound, 0, string.Empty, 0)); + + JObject? result = await InvokeQueryByIdAndPartitionKey(container.Object, false); + + Assert.IsNull(result); + } + + private static Mock CreateQueryContainer(FeedIterator iterator) + { + Mock container = new(); + container.Setup(x => x.GetItemQueryIterator( + It.IsAny(), + It.IsAny(), + It.IsAny())) + .Returns(iterator); + return container; + } + + private static CosmosSqlMetadataProvider CreateMetadataProvider() + { + CosmosSqlMetadataProvider provider = + (CosmosSqlMetadataProvider)RuntimeHelpers.GetUninitializedObject(typeof(CosmosSqlMetadataProvider)); + Entity entity = new( + Source: new EntitySource("db.books", EntitySourceType.Table, null, null), + GraphQL: new EntityGraphQLOptions("Book", "Books"), + Fields: null, + Rest: new EntityRestOptions(Enabled: true), + Permissions: Array.Empty(), + Mappings: null, + Relationships: null); + SetPrivateField(provider, "_runtimeConfigEntities", new RuntimeEntities(new Dictionary { ["Book"] = entity })); + SetPrivateField(provider, "_cosmosDb", new CosmosDbNoSQLDataSourceOptions("db", "books", null, null)); + SetPrivateField(provider, "_databaseType", DatabaseType.CosmosDB_NoSQL); + provider.GraphQLSchemaRoot = Utf8GraphQLParser.Parse("type Book { id: String }"); + provider.EntityWithJoins = new Dictionary>(); + return provider; + } + + private static DefaultHttpContext CreateHttpContext() + { + DefaultHttpContext context = new(); + context.Request.Headers[AuthorizationResolver.CLIENT_ROLE_HEADER] = AuthorizationResolver.ROLE_ANONYMOUS; + return context; + } + + private static void SetPrivateField(object instance, string name, object value) + { + instance.GetType().GetField(name, BindingFlags.Instance | BindingFlags.NonPublic)!.SetValue(instance, value); + } + + private static CosmosQueryStructure CreateStructure(bool isPaginated = false) + { + CosmosQueryStructure structure = + (CosmosQueryStructure)RuntimeHelpers.GetUninitializedObject(typeof(CosmosQueryStructure)); + structure.IsPaginated = isPaginated; + return structure; + } + + private static async Task InvokeExecuteQueryAsync( + CosmosQueryStructure structure, + QueryRequestOptions options, + Container container, + string id, + string partitionKey) + { + MethodInfo method = typeof(CosmosQueryEngine).GetMethod( + "ExecuteQueryAsync", + BindingFlags.Static | BindingFlags.NonPublic)!; + return await (Task)method.Invoke( + null, + new object[] { structure, new QueryDefinition("SELECT * FROM c"), options, container, id, partitionKey })!; + } + + private static async Task InvokeQueryByIdAndPartitionKey(Container container, bool isPaginated) + { + MethodInfo method = typeof(CosmosQueryEngine).GetMethod( + "QueryByIdAndPartitionKey", + BindingFlags.Static | BindingFlags.NonPublic)!; + return await (Task)method.Invoke( + null, + new object[] { container, "1", "tenant", isPaginated })!; + } + } +} diff --git a/src/Service.Tests/UnitTests/CosmosQueryStructureHelperTests.cs b/src/Service.Tests/UnitTests/CosmosQueryStructureHelperTests.cs new file mode 100644 index 0000000000..e5dfc5421c --- /dev/null +++ b/src/Service.Tests/UnitTests/CosmosQueryStructureHelperTests.cs @@ -0,0 +1,86 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +using System; +using System.Collections.Generic; +using System.Linq; +using System.Reflection; +using System.Runtime.CompilerServices; +using Azure.DataApiBuilder.Config.ObjectModel; +using Azure.DataApiBuilder.Core.Models; +using Azure.DataApiBuilder.Core.Resolvers; +using Azure.DataApiBuilder.Service.GraphQLBuilder.GraphQLTypes; +using HotChocolate.Language; +using Microsoft.VisualStudio.TestTools.UnitTesting; + +namespace Azure.DataApiBuilder.Service.Tests.UnitTests +{ + [TestClass, TestCategory(TestCategory.COSMOSDBNOSQL)] + public class CosmosQueryStructureHelperTests + { + [TestMethod] + public void GetTableAlias_IncrementsCounter() + { + CosmosQueryStructure structure = CreateStructure(); + + Assert.AreEqual("table0", structure.GetTableAlias()); + Assert.AreEqual("table1", structure.GetTableAlias()); + } + + [TestMethod] + public void GenerateQueryColumns_ExpandsFragmentSpreadsAndInlineFragments() + { + DocumentNode document = Utf8GraphQLParser.Parse(@" + query { + book { + id + ...BookFields + ... on Book { title } + } + } + fragment BookFields on Book { name }"); + OperationDefinitionNode operation = document.Definitions.OfType().Single(); + FieldNode book = operation.SelectionSet.Selections.OfType().Single(); + MethodInfo method = typeof(CosmosQueryStructure).GetMethod( + "GenerateQueryColumns", + BindingFlags.Static | BindingFlags.NonPublic)!; + + IEnumerable columns = (IEnumerable)method.Invoke( + null, + new object[] { book.SelectionSet!, document, "c" })!; + + CollectionAssert.AreEqual(new[] { "id", "name", "title" }, columns.Select(column => column.Label).ToArray()); + } + + [TestMethod] + public void ProcessGraphQLOrderByArg_SkipsNullAndMapsDescending() + { + CosmosQueryStructure structure = CreateStructure(); + List orderBy = new() + { + new("ignored", NullValueNode.Default), + new("title", new EnumValueNode("DESC")), + new("id", new EnumValueNode("ASC")) + }; + MethodInfo method = typeof(CosmosQueryStructure).GetMethod( + "ProcessGraphQLOrderByArg", + BindingFlags.Instance | BindingFlags.NonPublic)!; + + List columns = (List)method.Invoke(structure, new object[] { orderBy })!; + + Assert.AreEqual(2, columns.Count); + Assert.AreEqual(OrderBy.DESC, columns[0].Direction); + Assert.AreEqual(OrderBy.ASC, columns[1].Direction); + } + + private static CosmosQueryStructure CreateStructure() + { + CosmosQueryStructure structure = + (CosmosQueryStructure)RuntimeHelpers.GetUninitializedObject(typeof(CosmosQueryStructure)); + structure.TableCounter = new IncrementingInteger(); + typeof(CosmosQueryStructure).GetField("_containerAlias", BindingFlags.Instance | BindingFlags.NonPublic)! + .SetValue(structure, CosmosQueryStructure.COSMOSDB_CONTAINER_DEFAULT_ALIAS); + return structure; + } + } +} diff --git a/src/Service.Tests/UnitTests/CosmosSamplerHelperTests.cs b/src/Service.Tests/UnitTests/CosmosSamplerHelperTests.cs new file mode 100644 index 0000000000..d168beb4c7 --- /dev/null +++ b/src/Service.Tests/UnitTests/CosmosSamplerHelperTests.cs @@ -0,0 +1,69 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +using System; +using System.Collections.Generic; +using System.IO; +using System.Net; +using System.Text; +using System.Threading; +using System.Threading.Tasks; +using Azure.DataApiBuilder.Core.Generator.Sampler; +using Microsoft.Azure.Cosmos; +using Microsoft.Extensions.Logging; +using Microsoft.VisualStudio.TestTools.UnitTesting; +using Moq; + +namespace Azure.DataApiBuilder.Service.Tests.UnitTests +{ + [TestClass, TestCategory(TestCategory.COSMOSDBNOSQL)] + public class CosmosSamplerHelperTests + { + [TestMethod] + public async Task ExecuteQueryAsync_ProcessesArrayResponseAndInvokesCallback() + { + using ResponseMessage response = new(HttpStatusCode.OK) + { + Content = new MemoryStream(Encoding.UTF8.GetBytes("{\"Documents\":[{\"Id\":1},{\"Id\":2}]}")) + }; + CosmosExecutor executor = CreateExecutor(response); + List callbacks = new(); + + List results = await executor.ExecuteQueryAsync( + "SELECT * FROM c", + item => callbacks.Add(item)); + + CollectionAssert.AreEqual(new[] { 1, 2 }, results.ConvertAll(item => item.Id)); + Assert.AreEqual(2, callbacks.Count); + } + + [TestMethod] + public async Task ExecuteQueryAsync_UnsuccessfulResponseThrows() + { + using ResponseMessage response = new(HttpStatusCode.BadRequest); + CosmosExecutor executor = CreateExecutor(response); + + await Assert.ThrowsExceptionAsync(() => + executor.ExecuteQueryAsync("SELECT * FROM c")); + } + + private static CosmosExecutor CreateExecutor(ResponseMessage response) + { + Mock iterator = new(); + iterator.SetupSequence(x => x.HasMoreResults).Returns(true).Returns(false); + iterator.Setup(x => x.ReadNextAsync(It.IsAny())).ReturnsAsync(response); + Mock container = new(); + container.Setup(x => x.GetItemQueryStreamIterator( + It.IsAny(), + It.IsAny(), + It.IsAny())) + .Returns(iterator.Object); + return new CosmosExecutor(container.Object, Mock.Of()); + } + + private sealed class SampleItem + { + public int Id { get; set; } + } + } +} diff --git a/src/Service.Tests/UnitTests/CosmosSqlMetadataProviderHelperTests.cs b/src/Service.Tests/UnitTests/CosmosSqlMetadataProviderHelperTests.cs index f6547532aa..fc0922823a 100644 --- a/src/Service.Tests/UnitTests/CosmosSqlMetadataProviderHelperTests.cs +++ b/src/Service.Tests/UnitTests/CosmosSqlMetadataProviderHelperTests.cs @@ -8,18 +8,21 @@ using System.Runtime.CompilerServices; using Azure.DataApiBuilder.Config.DatabasePrimitives; using Azure.DataApiBuilder.Config.ObjectModel; +using Azure.DataApiBuilder.Core.Configurations; using Azure.DataApiBuilder.Core.Models; using Azure.DataApiBuilder.Core.Parsers; using Azure.DataApiBuilder.Core.Services.MetadataProviders; using Azure.DataApiBuilder.Service.Exceptions; using HotChocolate.Language; +using Microsoft.Extensions.Logging.Abstractions; using Microsoft.VisualStudio.TestTools.UnitTesting; using Moq; using System.IO.Abstractions; +using System.IO.Abstractions.TestingHelpers; namespace Azure.DataApiBuilder.Service.Tests.UnitTests { - [TestClass] + [TestClass, TestCategory(TestCategory.COSMOSDBNOSQL)] public class CosmosSqlMetadataProviderHelperTests { [TestMethod] @@ -138,6 +141,56 @@ public void GetSchemaName_MissingDatabaseThrows() Assert.ThrowsException(() => provider.GetSchemaName("Book")); } + [TestMethod] + public void GetDatabaseObjectName_NullSourceAndMissingContainerThrows() + { + CosmosSqlMetadataProvider provider = CreateProvider( + entities: new() { ["Book"] = CreateEntity(null!) }, + options: new("db", null, null, null)); + + Assert.ThrowsException(() => provider.GetDatabaseObjectName("Book")); + } + + [TestMethod] + public void GetDatabaseObjectName_EmptySourceAndMissingContainerReturnsEmptyName() + { + CosmosSqlMetadataProvider provider = CreateProvider( + entities: new() { ["Book"] = CreateEntity(string.Empty) }, + options: new("db", string.Empty, null, null)); + + Assert.AreEqual(string.Empty, provider.GetDatabaseObjectName("Book")); + } + + [TestMethod] + public void GetSchemaName_OnePartSourceAndMissingConfiguredDatabaseThrows() + { + CosmosSqlMetadataProvider provider = CreateProvider( + entities: new() { ["Book"] = CreateEntity("books") }, + options: new(null, "container", null, null)); + + Assert.ThrowsException(() => provider.GetSchemaName("Book")); + } + + [TestMethod] + public void Constructor_MissingCosmosOptionsThrowsInitializationError() + { + RuntimeConfig runtimeConfig = new( + Schema: string.Empty, + DataSource: new DataSource(DatabaseType.CosmosDB_NoSQL, string.Empty, Options: null), + Entities: new RuntimeEntities(new Dictionary())); + RuntimeConfigProvider configProvider = TestHelper.GenerateInMemoryRuntimeConfigProvider(runtimeConfig); + MockFileSystem fileSystem = new(); + RuntimeConfigValidator validator = new( + configProvider, + fileSystem, + NullLogger.Instance); + + DataApiBuilderException exception = Assert.ThrowsException(() => + new CosmosSqlMetadataProvider(configProvider, validator, fileSystem)); + + Assert.AreEqual(DataApiBuilderException.SubStatusCodes.ErrorInInitialization, exception.SubStatusCode); + } + [TestMethod] public void GetEntityName_ResolvesDirectModelDirectiveAndSingularNames() { @@ -184,6 +237,21 @@ public void ParseSchemaGraphQLDocument_LoadsSchemaFromConfiguredFile() Assert.IsNull(provider.GetSchemaGraphQLFieldFromFieldName("Missing", "id")); } + [TestMethod] + public void ParseSchemaGraphQLFieldsForJoins_AddsRepeatedModelPaths() + { + DocumentNode schema = Utf8GraphQLParser.Parse(@" + type FirstBook @model(name: ""Book"") { id: ID } + type SecondBook @model(name: ""Book"") { id: ID }"); + CosmosSqlMetadataProvider provider = CreateProvider( + entities: new() { ["Book"] = CreateEntity("db.books") }, + schema: schema); + + InvokePrivate(provider, "ParseSchemaGraphQLFieldsForJoins"); + + Assert.AreEqual(2, provider.EntityWithJoins["Book"].Count); + } + [TestMethod] public void AssertIfEntityIsAvailableInConfig_MissingEntityThrows() { @@ -210,6 +278,7 @@ private static CosmosSqlMetadataProvider CreateProvider( SetField(provider, "_partitionKeyPaths", new ConcurrentDictionary()); SetField(provider, "_oDataParser", new ODataParser()); SetField(provider, "_graphQLTypeToFieldsMap", new Dictionary>()); + provider.EntityWithJoins = new Dictionary>(); provider.GraphQLSchemaRoot = schema ?? new DocumentNode(Array.Empty()); return provider; } From 450b21b43d98031d3dcfa9c40466a0f5fc400cc1 Mon Sep 17 00:00:00 2001 From: Aaron Burtle Date: Sat, 29 Aug 2026 06:56:40 -0700 Subject: [PATCH 07/19] fix unit pipeline formatting --- .../Mcp/AggregateRecordsToolTests.cs | 4 ++-- .../BaseSqlQueryStructureHelperTests.cs | 1 - .../UnitTests/ConfigureJwtBearerOptionsTests.cs | 1 - .../UnitTests/GraphQLFilterParserUnitTests.cs | 1 - .../UnitTests/PureUtilityCoverageTests.cs | 2 +- .../UnitTests/SqlMutationEngineHelperTests.cs | 4 ++-- .../UnitTests/SqlPaginationUtilUnitTests.cs | 3 +-- .../UnitTests/SqlQueryEngineHelperTests.cs | 17 ++++++++--------- .../UnitTests/SqlQueryStructureHelperTests.cs | 1 - 9 files changed, 14 insertions(+), 20 deletions(-) diff --git a/src/Service.Tests/Mcp/AggregateRecordsToolTests.cs b/src/Service.Tests/Mcp/AggregateRecordsToolTests.cs index e57f48d1d6..213b0d87f4 100644 --- a/src/Service.Tests/Mcp/AggregateRecordsToolTests.cs +++ b/src/Service.Tests/Mcp/AggregateRecordsToolTests.cs @@ -666,8 +666,8 @@ public void BuildSimpleResponse_NullEmptyAndPopulatedResults_ReturnExpectedArray foreach (JsonArray? input in new JsonArray?[] { null, - new JsonArray(), - new JsonArray(new JsonObject { ["count"] = 3 }) + new(), + new(new JsonObject { ["count"] = 3 }) }) { CallToolResult result = InvokePrivateResponseBuilder( diff --git a/src/Service.Tests/UnitTests/BaseSqlQueryStructureHelperTests.cs b/src/Service.Tests/UnitTests/BaseSqlQueryStructureHelperTests.cs index d7a399e06e..0d05eb43ef 100644 --- a/src/Service.Tests/UnitTests/BaseSqlQueryStructureHelperTests.cs +++ b/src/Service.Tests/UnitTests/BaseSqlQueryStructureHelperTests.cs @@ -7,7 +7,6 @@ using Azure.DataApiBuilder.Auth; using Azure.DataApiBuilder.Config.DatabasePrimitives; using Azure.DataApiBuilder.Config.ObjectModel; -using Azure.DataApiBuilder.Core.Models; using Azure.DataApiBuilder.Core.Resolvers; using Azure.DataApiBuilder.Core.Services; using Azure.DataApiBuilder.Service.Exceptions; diff --git a/src/Service.Tests/UnitTests/ConfigureJwtBearerOptionsTests.cs b/src/Service.Tests/UnitTests/ConfigureJwtBearerOptionsTests.cs index a6ee19f470..edad902844 100644 --- a/src/Service.Tests/UnitTests/ConfigureJwtBearerOptionsTests.cs +++ b/src/Service.Tests/UnitTests/ConfigureJwtBearerOptionsTests.cs @@ -6,7 +6,6 @@ using Azure.DataApiBuilder.Config; using Azure.DataApiBuilder.Config.ObjectModel; using Azure.DataApiBuilder.Core.Configurations; -using Azure.DataApiBuilder.Service; using Microsoft.AspNetCore.Authentication.JwtBearer; using Microsoft.VisualStudio.TestTools.UnitTesting; diff --git a/src/Service.Tests/UnitTests/GraphQLFilterParserUnitTests.cs b/src/Service.Tests/UnitTests/GraphQLFilterParserUnitTests.cs index 8a051af4eb..46d31b98de 100644 --- a/src/Service.Tests/UnitTests/GraphQLFilterParserUnitTests.cs +++ b/src/Service.Tests/UnitTests/GraphQLFilterParserUnitTests.cs @@ -6,7 +6,6 @@ using Azure.DataApiBuilder.Config.ObjectModel; using Azure.DataApiBuilder.Core.Configurations; using Azure.DataApiBuilder.Core.Models; -using Azure.DataApiBuilder.Core.Resolvers; using Azure.DataApiBuilder.Core.Services.MetadataProviders; using Azure.DataApiBuilder.Service.Exceptions; using HotChocolate.Language; diff --git a/src/Service.Tests/UnitTests/PureUtilityCoverageTests.cs b/src/Service.Tests/UnitTests/PureUtilityCoverageTests.cs index 92c83a5573..975b9b7dbd 100644 --- a/src/Service.Tests/UnitTests/PureUtilityCoverageTests.cs +++ b/src/Service.Tests/UnitTests/PureUtilityCoverageTests.cs @@ -6,10 +6,10 @@ using Azure.DataApiBuilder.Config; using Azure.DataApiBuilder.Core.Configurations; using Azure.DataApiBuilder.Core.Models; -using MetadataTypeConverter = Azure.DataApiBuilder.Core.Services.MetadataProviders.Converters.TypeConverter; using Azure.DataApiBuilder.Core.Telemetry; using Microsoft.Extensions.Logging; using Microsoft.VisualStudio.TestTools.UnitTesting; +using MetadataTypeConverter = Azure.DataApiBuilder.Core.Services.MetadataProviders.Converters.TypeConverter; namespace Azure.DataApiBuilder.Service.Tests.UnitTests { diff --git a/src/Service.Tests/UnitTests/SqlMutationEngineHelperTests.cs b/src/Service.Tests/UnitTests/SqlMutationEngineHelperTests.cs index 881a4c1a36..f881e53f39 100644 --- a/src/Service.Tests/UnitTests/SqlMutationEngineHelperTests.cs +++ b/src/Service.Tests/UnitTests/SqlMutationEngineHelperTests.cs @@ -21,12 +21,12 @@ using Azure.DataApiBuilder.Core.Services; using Azure.DataApiBuilder.Core.Services.MetadataProviders; using Azure.DataApiBuilder.Service.Exceptions; +using HotChocolate.Language; +using HotChocolate.Resolvers; using Microsoft.AspNetCore.Http; using Microsoft.AspNetCore.Mvc; using Microsoft.VisualStudio.TestTools.UnitTesting; using Moq; -using HotChocolate.Language; -using HotChocolate.Resolvers; namespace Azure.DataApiBuilder.Service.Tests.UnitTests { diff --git a/src/Service.Tests/UnitTests/SqlPaginationUtilUnitTests.cs b/src/Service.Tests/UnitTests/SqlPaginationUtilUnitTests.cs index dc14315539..f58794836f 100644 --- a/src/Service.Tests/UnitTests/SqlPaginationUtilUnitTests.cs +++ b/src/Service.Tests/UnitTests/SqlPaginationUtilUnitTests.cs @@ -1,10 +1,9 @@ // Copyright (c) Microsoft Corporation. // Licensed under the MIT License. -using System.Collections.Specialized; using System.Collections.Generic; +using System.Collections.Specialized; using System.Text.Json; -using Azure.DataApiBuilder.Config.DatabasePrimitives; using Azure.DataApiBuilder.Core.Models; using Azure.DataApiBuilder.Core.Resolvers; using Azure.DataApiBuilder.Service.Exceptions; diff --git a/src/Service.Tests/UnitTests/SqlQueryEngineHelperTests.cs b/src/Service.Tests/UnitTests/SqlQueryEngineHelperTests.cs index 3998fbf370..b7ba612ec8 100644 --- a/src/Service.Tests/UnitTests/SqlQueryEngineHelperTests.cs +++ b/src/Service.Tests/UnitTests/SqlQueryEngineHelperTests.cs @@ -16,7 +16,6 @@ using Azure.DataApiBuilder.Core.Models; using Azure.DataApiBuilder.Core.Resolvers; using Azure.DataApiBuilder.Core.Resolvers.Factories; -using Azure.DataApiBuilder.Core.Services; using Azure.DataApiBuilder.Core.Services.Cache; using Azure.DataApiBuilder.Core.Services.MetadataProviders; using Microsoft.AspNetCore.Http; @@ -29,8 +28,8 @@ namespace Azure.DataApiBuilder.Service.Tests.UnitTests [TestClass] public class SqlQueryEngineHelperTests { - private const string DataSourceName = "default"; - private const string EntityName = "Book"; + private const string DATA_SOURCE_NAME = "default"; + private const string ENTITY_NAME = "Book"; [DataTestMethod] [DataRow("{\"value\":1}", true)] @@ -59,7 +58,7 @@ public async Task ExecuteStoredProcedureCore_HandlesResultShapes(string? json, b It.IsAny(), It.IsAny>(), It.IsAny?, Task>>(), - DataSourceName, + DATA_SOURCE_NAME, It.IsAny(), It.IsAny?>())) .ReturnsAsync(resultArray!); @@ -73,7 +72,7 @@ public async Task ExecuteStoredProcedureCore_HandlesResultShapes(string? json, b using JsonDocument? result = await (Task)method.Invoke( engine, - new object[] { structure, DataSourceName })!; + new object[] { structure, DATA_SOURCE_NAME })!; Assert.AreEqual(expectsDocument, result is not null); } @@ -91,7 +90,7 @@ public async Task ExecuteListCore_ReturnsExecutorResult(bool returnList) It.IsAny(), It.IsAny>(), It.IsAny?, Task>>>(), - DataSourceName, + DATA_SOURCE_NAME, It.IsAny(), It.IsAny?>())) .ReturnsAsync(expected!); @@ -105,7 +104,7 @@ public async Task ExecuteListCore_ReturnsExecutorResult(bool returnList) List? result = await (Task?>)method.Invoke( engine, - new object[] { structure, DataSourceName })!; + new object[] { structure, DATA_SOURCE_NAME })!; Assert.AreSame(expected, result); if (expected is not null) @@ -123,7 +122,7 @@ private static (SqlQueryEngine Engine, Mock Executor) CreateEngi Schema: string.Empty, DataSource: new DataSource(DatabaseType.MSSQL, string.Empty), Entities: new RuntimeEntities(new Dictionary())); - runtimeConfig.UpdateDefaultDataSourceName(DataSourceName); + runtimeConfig.UpdateDefaultDataSourceName(DATA_SOURCE_NAME); Mock loader = new(null, null); Mock configProviderMock = new(loader.Object); configProviderMock.Setup(x => x.GetConfig()).Returns(runtimeConfig); @@ -154,7 +153,7 @@ private static (SqlQueryEngine Engine, Mock Executor) CreateEngi private static T CreateUninitializedStructure() { T structure = (T)RuntimeHelpers.GetUninitializedObject(typeof(T)); - SetProperty(structure!, "EntityName", EntityName); + SetProperty(structure!, "EntityName", ENTITY_NAME); SetProperty(structure!, "Parameters", new Dictionary()); return structure; } diff --git a/src/Service.Tests/UnitTests/SqlQueryStructureHelperTests.cs b/src/Service.Tests/UnitTests/SqlQueryStructureHelperTests.cs index 3a5f10d663..8550c1123f 100644 --- a/src/Service.Tests/UnitTests/SqlQueryStructureHelperTests.cs +++ b/src/Service.Tests/UnitTests/SqlQueryStructureHelperTests.cs @@ -6,7 +6,6 @@ using System.Runtime.CompilerServices; using Azure.DataApiBuilder.Core.Models; using Azure.DataApiBuilder.Core.Resolvers; -using Azure.DataApiBuilder.Core.Resolvers.Sql_Query_Structures; using Azure.DataApiBuilder.Service.Exceptions; using Azure.DataApiBuilder.Service.GraphQLBuilder.Queries; using HotChocolate.Language; From 5efb26de656699f59235718b9f80b7a6727da2db Mon Sep 17 00:00:00 2001 From: Aaron Burtle Date: Sat, 29 Aug 2026 06:56:50 -0700 Subject: [PATCH 08/19] fix MsSql pipeline formatting --- src/Service.Tests/UnitTests/QueryExecutorHelperTests.cs | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/src/Service.Tests/UnitTests/QueryExecutorHelperTests.cs b/src/Service.Tests/UnitTests/QueryExecutorHelperTests.cs index 2076be6162..278ca5a919 100644 --- a/src/Service.Tests/UnitTests/QueryExecutorHelperTests.cs +++ b/src/Service.Tests/UnitTests/QueryExecutorHelperTests.cs @@ -13,11 +13,11 @@ using System.Text.Json; using System.Text.Json.Nodes; using System.Threading.Tasks; -using Azure.DataApiBuilder.Core.Resolvers; -using Azure.DataApiBuilder.Core.Models; -using Azure.DataApiBuilder.Core.Services; using Azure.DataApiBuilder.Config.ObjectModel; using Azure.DataApiBuilder.Core.Configurations; +using Azure.DataApiBuilder.Core.Models; +using Azure.DataApiBuilder.Core.Resolvers; +using Azure.DataApiBuilder.Core.Services; using Azure.DataApiBuilder.Service.Exceptions; using Microsoft.AspNetCore.Http; using Microsoft.Data.SqlClient; From fa9c9579284a5268400f975eb68b7d04c1df429c Mon Sep 17 00:00:00 2001 From: Aaron Burtle Date: Sat, 29 Aug 2026 06:57:00 -0700 Subject: [PATCH 09/19] fix MySQL pipeline formatting --- src/Service.Tests/UnitTests/MySqlDbExceptionParserHelperTests.cs | 1 - 1 file changed, 1 deletion(-) diff --git a/src/Service.Tests/UnitTests/MySqlDbExceptionParserHelperTests.cs b/src/Service.Tests/UnitTests/MySqlDbExceptionParserHelperTests.cs index 3cbe744f0c..6535eb7f64 100644 --- a/src/Service.Tests/UnitTests/MySqlDbExceptionParserHelperTests.cs +++ b/src/Service.Tests/UnitTests/MySqlDbExceptionParserHelperTests.cs @@ -1,7 +1,6 @@ // Copyright (c) Microsoft Corporation. // Licensed under the MIT License. -using System; using System.Collections.Generic; using System.Net; using System.Reflection; From 080b9b89f6ec78e14d8493ef6267cd9a7127a15c Mon Sep 17 00:00:00 2001 From: Aaron Burtle Date: Sat, 29 Aug 2026 06:57:08 -0700 Subject: [PATCH 10/19] fix PostgreSQL pipeline formatting --- .../UnitTests/PostgreSqlMetadataProviderHelperTests.cs | 1 - 1 file changed, 1 deletion(-) diff --git a/src/Service.Tests/UnitTests/PostgreSqlMetadataProviderHelperTests.cs b/src/Service.Tests/UnitTests/PostgreSqlMetadataProviderHelperTests.cs index 78237ee161..789d0883b9 100644 --- a/src/Service.Tests/UnitTests/PostgreSqlMetadataProviderHelperTests.cs +++ b/src/Service.Tests/UnitTests/PostgreSqlMetadataProviderHelperTests.cs @@ -8,7 +8,6 @@ using System.Runtime.CompilerServices; using Azure.DataApiBuilder.Config.DatabasePrimitives; using Azure.DataApiBuilder.Config.ObjectModel; -using Azure.DataApiBuilder.Core.Models; using Azure.DataApiBuilder.Core.Services; using Azure.DataApiBuilder.Service.Exceptions; using Microsoft.VisualStudio.TestTools.UnitTesting; From b194e0843072f4ddaa6b39c9b5b3bebd1a7e5bc0 Mon Sep 17 00:00:00 2001 From: Aaron Burtle Date: Sat, 29 Aug 2026 06:57:17 -0700 Subject: [PATCH 11/19] fix Cosmos DB pipeline formatting --- .../UnitTests/CosmosClientProviderHelperTests.cs | 2 +- src/Service.Tests/UnitTests/CosmosEngineHelperTests.cs | 7 +++---- .../UnitTests/CosmosQueryEngineExecutionHelperTests.cs | 1 - .../UnitTests/CosmosQueryStructureHelperTests.cs | 2 -- .../UnitTests/CosmosSqlMetadataProviderHelperTests.cs | 5 ++--- 5 files changed, 6 insertions(+), 11 deletions(-) diff --git a/src/Service.Tests/UnitTests/CosmosClientProviderHelperTests.cs b/src/Service.Tests/UnitTests/CosmosClientProviderHelperTests.cs index 3388607694..49e0d25284 100644 --- a/src/Service.Tests/UnitTests/CosmosClientProviderHelperTests.cs +++ b/src/Service.Tests/UnitTests/CosmosClientProviderHelperTests.cs @@ -3,6 +3,7 @@ using System; using System.Collections.Generic; +using System.IO.Abstractions.TestingHelpers; using System.Reflection; using System.Runtime.CompilerServices; using System.Threading.Tasks; @@ -12,7 +13,6 @@ using Azure.DataApiBuilder.Core.Configurations; using Azure.DataApiBuilder.Core.Resolvers; using Microsoft.VisualStudio.TestTools.UnitTesting; -using System.IO.Abstractions.TestingHelpers; namespace Azure.DataApiBuilder.Service.Tests.UnitTests { diff --git a/src/Service.Tests/UnitTests/CosmosEngineHelperTests.cs b/src/Service.Tests/UnitTests/CosmosEngineHelperTests.cs index ed0e079d32..eca39879d1 100644 --- a/src/Service.Tests/UnitTests/CosmosEngineHelperTests.cs +++ b/src/Service.Tests/UnitTests/CosmosEngineHelperTests.cs @@ -14,16 +14,15 @@ using Azure.DataApiBuilder.Core.Authorization; using Azure.DataApiBuilder.Core.Models; using Azure.DataApiBuilder.Core.Resolvers; -using Azure.DataApiBuilder.Core.Services; using Azure.DataApiBuilder.Service.Exceptions; +using Azure.DataApiBuilder.Service.GraphQLBuilder.Mutations; +using Azure.DataApiBuilder.Service.GraphQLBuilder.Queries; using HotChocolate.Language; using HotChocolate.Resolvers; -using Microsoft.Extensions.Primitives; using Microsoft.Azure.Cosmos; +using Microsoft.Extensions.Primitives; using Microsoft.VisualStudio.TestTools.UnitTesting; using Moq; -using Azure.DataApiBuilder.Service.GraphQLBuilder.Mutations; -using Azure.DataApiBuilder.Service.GraphQLBuilder.Queries; using Newtonsoft.Json.Linq; namespace Azure.DataApiBuilder.Service.Tests.UnitTests diff --git a/src/Service.Tests/UnitTests/CosmosQueryEngineExecutionHelperTests.cs b/src/Service.Tests/UnitTests/CosmosQueryEngineExecutionHelperTests.cs index f739fc0cc8..6fd234c0d3 100644 --- a/src/Service.Tests/UnitTests/CosmosQueryEngineExecutionHelperTests.cs +++ b/src/Service.Tests/UnitTests/CosmosQueryEngineExecutionHelperTests.cs @@ -15,7 +15,6 @@ using Azure.DataApiBuilder.Core.Configurations; using Azure.DataApiBuilder.Core.Models; using Azure.DataApiBuilder.Core.Resolvers; -using Azure.DataApiBuilder.Core.Services; using Azure.DataApiBuilder.Core.Services.MetadataProviders; using Azure.DataApiBuilder.Service.GraphQLBuilder.Queries; using HotChocolate; diff --git a/src/Service.Tests/UnitTests/CosmosQueryStructureHelperTests.cs b/src/Service.Tests/UnitTests/CosmosQueryStructureHelperTests.cs index e5dfc5421c..b58b6d2dd5 100644 --- a/src/Service.Tests/UnitTests/CosmosQueryStructureHelperTests.cs +++ b/src/Service.Tests/UnitTests/CosmosQueryStructureHelperTests.cs @@ -1,12 +1,10 @@ // Copyright (c) Microsoft Corporation. // Licensed under the MIT License. -using System; using System.Collections.Generic; using System.Linq; using System.Reflection; using System.Runtime.CompilerServices; -using Azure.DataApiBuilder.Config.ObjectModel; using Azure.DataApiBuilder.Core.Models; using Azure.DataApiBuilder.Core.Resolvers; using Azure.DataApiBuilder.Service.GraphQLBuilder.GraphQLTypes; diff --git a/src/Service.Tests/UnitTests/CosmosSqlMetadataProviderHelperTests.cs b/src/Service.Tests/UnitTests/CosmosSqlMetadataProviderHelperTests.cs index fc0922823a..2f7fbcf840 100644 --- a/src/Service.Tests/UnitTests/CosmosSqlMetadataProviderHelperTests.cs +++ b/src/Service.Tests/UnitTests/CosmosSqlMetadataProviderHelperTests.cs @@ -4,12 +4,13 @@ using System; using System.Collections.Concurrent; using System.Collections.Generic; +using System.IO.Abstractions; +using System.IO.Abstractions.TestingHelpers; using System.Reflection; using System.Runtime.CompilerServices; using Azure.DataApiBuilder.Config.DatabasePrimitives; using Azure.DataApiBuilder.Config.ObjectModel; using Azure.DataApiBuilder.Core.Configurations; -using Azure.DataApiBuilder.Core.Models; using Azure.DataApiBuilder.Core.Parsers; using Azure.DataApiBuilder.Core.Services.MetadataProviders; using Azure.DataApiBuilder.Service.Exceptions; @@ -17,8 +18,6 @@ using Microsoft.Extensions.Logging.Abstractions; using Microsoft.VisualStudio.TestTools.UnitTesting; using Moq; -using System.IO.Abstractions; -using System.IO.Abstractions.TestingHelpers; namespace Azure.DataApiBuilder.Service.Tests.UnitTests { From 9e752f5bc424b0d8bdf1109134f346bb01b13abb Mon Sep 17 00:00:00 2001 From: Aaron Burtle Date: Sat, 29 Aug 2026 07:09:42 -0700 Subject: [PATCH 12/19] remove unused Cosmos DB test import --- .../UnitTests/CosmosSqlMetadataProviderHelperTests.cs | 1 - 1 file changed, 1 deletion(-) diff --git a/src/Service.Tests/UnitTests/CosmosSqlMetadataProviderHelperTests.cs b/src/Service.Tests/UnitTests/CosmosSqlMetadataProviderHelperTests.cs index 2f7fbcf840..2918ddcbfc 100644 --- a/src/Service.Tests/UnitTests/CosmosSqlMetadataProviderHelperTests.cs +++ b/src/Service.Tests/UnitTests/CosmosSqlMetadataProviderHelperTests.cs @@ -8,7 +8,6 @@ using System.IO.Abstractions.TestingHelpers; using System.Reflection; using System.Runtime.CompilerServices; -using Azure.DataApiBuilder.Config.DatabasePrimitives; using Azure.DataApiBuilder.Config.ObjectModel; using Azure.DataApiBuilder.Core.Configurations; using Azure.DataApiBuilder.Core.Parsers; From 67560e8bb3305433a78eeda6a07193a0c6105c5a Mon Sep 17 00:00:00 2001 From: Aaron Burtle Date: Sat, 29 Aug 2026 07:09:52 -0700 Subject: [PATCH 13/19] remove unused MsSql test import --- src/Service.Tests/UnitTests/QueryExecutorHelperTests.cs | 1 - 1 file changed, 1 deletion(-) diff --git a/src/Service.Tests/UnitTests/QueryExecutorHelperTests.cs b/src/Service.Tests/UnitTests/QueryExecutorHelperTests.cs index 278ca5a919..88eb72c2b6 100644 --- a/src/Service.Tests/UnitTests/QueryExecutorHelperTests.cs +++ b/src/Service.Tests/UnitTests/QueryExecutorHelperTests.cs @@ -17,7 +17,6 @@ using Azure.DataApiBuilder.Core.Configurations; using Azure.DataApiBuilder.Core.Models; using Azure.DataApiBuilder.Core.Resolvers; -using Azure.DataApiBuilder.Core.Services; using Azure.DataApiBuilder.Service.Exceptions; using Microsoft.AspNetCore.Http; using Microsoft.Data.SqlClient; From ba6deb82823b07cbf15b78649caac9f8d470acf1 Mon Sep 17 00:00:00 2001 From: Aaron Burtle Date: Sat, 29 Aug 2026 07:18:15 -0700 Subject: [PATCH 14/19] fix entity health threshold error message --- src/Config/Converters/EntityHealthOptionsConvertorFactory.cs | 2 +- .../UnitTests/EntityHealthOptionsConverterTests.cs | 4 ++-- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/src/Config/Converters/EntityHealthOptionsConvertorFactory.cs b/src/Config/Converters/EntityHealthOptionsConvertorFactory.cs index ea95e212e4..98bd549d59 100644 --- a/src/Config/Converters/EntityHealthOptionsConvertorFactory.cs +++ b/src/Config/Converters/EntityHealthOptionsConvertorFactory.cs @@ -79,7 +79,7 @@ private class HealthCheckOptionsConverter : JsonConverter Date: Sat, 29 Aug 2026 07:33:41 -0700 Subject: [PATCH 15/19] fix aggregate records test indentation --- .../Mcp/AggregateRecordsToolTests.cs | 78 +++++++++---------- 1 file changed, 39 insertions(+), 39 deletions(-) diff --git a/src/Service.Tests/Mcp/AggregateRecordsToolTests.cs b/src/Service.Tests/Mcp/AggregateRecordsToolTests.cs index 213b0d87f4..4d2ac2f9c2 100644 --- a/src/Service.Tests/Mcp/AggregateRecordsToolTests.cs +++ b/src/Service.Tests/Mcp/AggregateRecordsToolTests.cs @@ -392,49 +392,49 @@ public async Task AggregateRecords_InvalidFieldFunctionCombination_ReturnsInvali $"Error message must contain '{expectedInMessage}'. Actual: '{message}'"); } - [DataTestMethod] - [DataRow("{\"entity\":\"Book\",\"function\":\"count\",\"field\":\"\"}", "EntityNotFound")] - [DataRow("{\"entity\":\"Book\",\"function\":\"sum\",\"field\":\"\"}", "InvalidArguments")] - [DataRow("{\"entity\":\"Book\",\"function\":\"count\",\"field\":\"id\",\"distinct\":\"yes\"}", "InvalidArguments")] - [DataRow("{\"entity\":\"Book\",\"function\":\"count\",\"first\":0}", "InvalidArguments")] - [DataRow("{\"entity\":\"Book\",\"function\":\"count\",\"first\":1}", "InvalidArguments")] - [DataRow("{\"entity\":\"Book\",\"function\":\"count\",\"after\":\"MA==\"}", "InvalidArguments")] - [DataRow("{\"entity\":\"Book\",\"function\":\"count\",\"groupby\":[\"title\"],\"after\":\"MA==\"}", "InvalidArguments")] - [DataRow("{\"entity\":\"Book\",\"function\":\"count\",\"groupby\":[\"title\"],\"having\":{\"in\":[]}}", "InvalidArguments")] - public async Task AggregateRecords_AdditionalArgumentEdges_ReturnExpectedError(string json, string expectedError) - { - CallToolResult result = await ExecuteToolAsync(CreateDefaultServiceProvider(), json); + [DataTestMethod] + [DataRow("{\"entity\":\"Book\",\"function\":\"count\",\"field\":\"\"}", "EntityNotFound")] + [DataRow("{\"entity\":\"Book\",\"function\":\"sum\",\"field\":\"\"}", "InvalidArguments")] + [DataRow("{\"entity\":\"Book\",\"function\":\"count\",\"field\":\"id\",\"distinct\":\"yes\"}", "InvalidArguments")] + [DataRow("{\"entity\":\"Book\",\"function\":\"count\",\"first\":0}", "InvalidArguments")] + [DataRow("{\"entity\":\"Book\",\"function\":\"count\",\"first\":1}", "InvalidArguments")] + [DataRow("{\"entity\":\"Book\",\"function\":\"count\",\"after\":\"MA==\"}", "InvalidArguments")] + [DataRow("{\"entity\":\"Book\",\"function\":\"count\",\"groupby\":[\"title\"],\"after\":\"MA==\"}", "InvalidArguments")] + [DataRow("{\"entity\":\"Book\",\"function\":\"count\",\"groupby\":[\"title\"],\"having\":{\"in\":[]}}", "InvalidArguments")] + public async Task AggregateRecords_AdditionalArgumentEdges_ReturnExpectedError(string json, string expectedError) + { + CallToolResult result = await ExecuteToolAsync(CreateDefaultServiceProvider(), json); - AssertErrorResult(result, expectedError); - } + AssertErrorResult(result, expectedError); + } - [TestMethod] - public async Task AggregateRecords_ValidHavingOperators_PassArgumentValidation() + [TestMethod] + public async Task AggregateRecords_ValidHavingOperators_PassArgumentValidation() + { + const string json = """ { - const string json = """ - { - "entity": "Book", - "function": "count", - "groupby": ["title", "TITLE", ""], - "having": { - "eq": 1, - "neq": 2, - "gt": 3, - "gte": 4, - "lt": 5, - "lte": 6, - "in": [7, 8] - } - } - """; - - CallToolResult result = await ExecuteToolAsync(CreateDefaultServiceProvider(), json); - - JsonElement content = ParseContent(result); - Assert.AreNotEqual( - "InvalidArguments", - content.GetProperty("error").GetProperty("type").GetString()); + "entity": "Book", + "function": "count", + "groupby": ["title", "TITLE", ""], + "having": { + "eq": 1, + "neq": 2, + "gt": 3, + "gte": 4, + "lt": 5, + "lte": 6, + "in": [7, 8] + } } + """; + + CallToolResult result = await ExecuteToolAsync(CreateDefaultServiceProvider(), json); + + JsonElement content = ParseContent(result); + Assert.AreNotEqual( + "InvalidArguments", + content.GetProperty("error").GetProperty("type").GetString()); + } #endregion From 4dbc4440661e7ac6f28faadcc82698d9a6b5d616 Mon Sep 17 00:00:00 2001 From: Aaron Burtle Date: Sat, 29 Aug 2026 07:49:07 -0700 Subject: [PATCH 16/19] fix remaining unit test formatting --- src/Service.Tests/UnitTests/PureUtilityCoverageTests.cs | 2 +- src/Service.Tests/UnitTests/SqlMutationEngineHelperTests.cs | 2 +- src/Service.Tests/UnitTests/SqlQueryStructuresModelTests.cs | 2 +- 3 files changed, 3 insertions(+), 3 deletions(-) diff --git a/src/Service.Tests/UnitTests/PureUtilityCoverageTests.cs b/src/Service.Tests/UnitTests/PureUtilityCoverageTests.cs index 975b9b7dbd..ac7585fee7 100644 --- a/src/Service.Tests/UnitTests/PureUtilityCoverageTests.cs +++ b/src/Service.Tests/UnitTests/PureUtilityCoverageTests.cs @@ -98,4 +98,4 @@ public void MetadataTypeConverter_NonStringInputThrows() Assert.ThrowsException(() => JsonSerializer.Deserialize("42", options)); } } -} \ No newline at end of file +} diff --git a/src/Service.Tests/UnitTests/SqlMutationEngineHelperTests.cs b/src/Service.Tests/UnitTests/SqlMutationEngineHelperTests.cs index f881e53f39..5dad2b0b04 100644 --- a/src/Service.Tests/UnitTests/SqlMutationEngineHelperTests.cs +++ b/src/Service.Tests/UnitTests/SqlMutationEngineHelperTests.cs @@ -610,7 +610,7 @@ private static (SqlMutationEngine Engine, StoredProcedureRequestContext Context, It.IsAny?, Task>>(), dataSourceName, It.IsAny(), - It.IsAny? >())) + It.IsAny?>())) .ReturnsAsync(result); Mock queryManagerFactory = new(); queryManagerFactory.Setup(x => x.GetQueryBuilder(DatabaseType.MSSQL)).Returns(queryBuilder.Object); diff --git a/src/Service.Tests/UnitTests/SqlQueryStructuresModelTests.cs b/src/Service.Tests/UnitTests/SqlQueryStructuresModelTests.cs index 7ff209eeac..1ba429fdf1 100644 --- a/src/Service.Tests/UnitTests/SqlQueryStructuresModelTests.cs +++ b/src/Service.Tests/UnitTests/SqlQueryStructuresModelTests.cs @@ -50,4 +50,4 @@ public void PredicateOperand_StringAndPredicateAccessorsReflectStoredType() Assert.IsTrue(nested.IsPredicate()); } } -} \ No newline at end of file +} From 2138ee61e923def8a3990c94586a17a5ddd05c8e Mon Sep 17 00:00:00 2001 From: Aaron Burtle Date: Sat, 29 Aug 2026 19:42:58 -0700 Subject: [PATCH 17/19] stabilize Cosmos time-partitioned sampler tests --- src/Service.Tests/CosmosTests/SamplerTests.cs | 84 ++++++++++++++----- .../UnitTests/CosmosSamplerHelperTests.cs | 58 +++++++++++++ 2 files changed, 119 insertions(+), 23 deletions(-) diff --git a/src/Service.Tests/CosmosTests/SamplerTests.cs b/src/Service.Tests/CosmosTests/SamplerTests.cs index 8ddb6c3d7e..009e42854c 100644 --- a/src/Service.Tests/CosmosTests/SamplerTests.cs +++ b/src/Service.Tests/CosmosTests/SamplerTests.cs @@ -39,6 +39,8 @@ public class SamplerTests : TestBase private const string CONTAINER_NAME_ID_PK = "containerWithIdPk"; private const string CONTAINER_NAME_NAME_PK = "containerWithNamePk"; + private const int DEFAULT_TIME_GROUP_COUNT = 10; + private const int DEFAULT_RECORDS_PER_TIME_GROUP = 10; /// /// Initializes the test environment by creating Cosmos DB containers and populating them with sample data. @@ -61,9 +63,9 @@ public async Task Initialize() // Retrieve timestamps from the container to use in validation. CosmosExecutor executor = new(_containerWithIdPk, new Mock().Object); - await executor - .ExecuteQueryAsync("SELECT DISTINCT c._ts FROM c ORDER BY c._ts desc", - callback: (item) => _sortedTimespansIdPk.Add(item.RootElement.GetProperty("_ts").GetInt32())); + await executor.ExecuteQueryAsync( + "SELECT c._ts FROM c ORDER BY c._ts desc", + callback: (item) => _sortedTimespansIdPk.Add(item.RootElement.GetProperty("_ts").GetInt32())); // Insert additional items into the second container with a delay for unique timestamps and partitioned over name i.e planets name. // Number of partitions would be 9 as we have 9 unique names. @@ -71,9 +73,9 @@ await executor // Retrieve timestamps for the second container to use in validation. executor = new(_containerWithNamePk, new Mock().Object); - await executor - .ExecuteQueryAsync("SELECT DISTINCT c._ts FROM c ORDER BY c._ts desc", - callback: (item) => _sortedTimespansNamePk.Add(item.RootElement.GetProperty("_ts").GetInt32())); + await executor.ExecuteQueryAsync( + "SELECT c._ts FROM c ORDER BY c._ts desc", + callback: (item) => _sortedTimespansNamePk.Add(item.RootElement.GetProperty("_ts").GetInt32())); } /// @@ -118,7 +120,6 @@ public async Task TestTopNExtractor(int count, int? maxDays, int expectedCount) /// The path of the partition key to use for sampling. If null, partition key path is not considered. /// The number of records to retrieve per partition. Defaults to 5 if not specified. /// The maximum number of days to filter records within each partition. If null, no date-based filtering is applied. - /// The expected number of records returned by the sampler. /// /// This test case ensures that the EligibleDataSampler handles partition-based sampling correctly with various configurations. /// It verifies that the sampler correctly applies partition key paths, record limits per partition, and date-based filters as specified. @@ -198,31 +199,68 @@ public async Task TestGetPartitionInfoInEligibleDataSampler(string partitionKeyP /// The test cases also include scenarios where records are not evenly distributed across time-based groups. /// [TestMethod(displayName: "TimePartitionedSampler Scenarios")] - [DataRow(5, 1, 0, 5, DisplayName = "Retrieve 1 record, if it is allowed to fetch 1 item from a group and there are 5 groups (or time range)")] - [DataRow(1, 10, 0, 10, DisplayName = "Retrieve 10 records, if it is allowed to fetch 10 item from a group and there is only 1 group.")] - [DataRow(null, 1, 0, 10, DisplayName = "Retrieve 10 records, if 1 item is allowed to fetch from each group and number of groups is 10 (i.e default)")] - [DataRow(null, null, null, 10, DisplayName = "Retrieve 10 records i.e last 10 days data, based on default values when no specific limits are set.")] - [DataRow(5, 1, 4, 1, DisplayName = "Retrieve 1 record from a single group when records cannot be evenly divided into time-based groups.")] - public async Task TestTimePartitionedSampler(int? groupCount, int? numberOfRecordsPerGroup, int? maxDays, int expectedResultCount) + [DataRow(5, 1, 0, DisplayName = "Retrieve at most 1 record from each of 5 groups.")] + [DataRow(1, 10, 0, DisplayName = "Retrieve at most 10 records from a single group.")] + [DataRow(null, 1, 0, DisplayName = "Use the default group count and retrieve at most 1 record from each group.")] + [DataRow(null, null, null, DisplayName = "Use the default group, record, and day limits.")] + [DataRow(5, 1, 4, DisplayName = "Sample a short time range that cannot be evenly divided into 5 groups.")] + public async Task TestTimePartitionedSampler(int? groupCount, int? numberOfRecordsPerGroup, int? maxDays) { Mock timePartitionedSampler = new(_containerWithNamePk, groupCount, numberOfRecordsPerGroup, maxDays, _mockLogger.Object); - if (maxDays is null || maxDays == 0) + // Compress day-sized windows to seconds so this integration test does not take days to arrange its data. + // Cosmos writes can take longer than one second on hosted agents, so the observed timestamps may contain gaps. + int timeWindowInSeconds = maxDays ?? TimePartitionedSampler.MAX_DAYS; + if (timeWindowInSeconds > 0) { - maxDays = TimePartitionedSampler.MAX_DAYS; + timePartitionedSampler + .Setup(x => x.GetTimeStampThreshold()) + .Returns(_sortedTimespansNamePk[0] - timeWindowInSeconds); } - timePartitionedSampler - .Setup(x => x.GetTimeStampThreshold()) - .Returns((long)(_sortedTimespansNamePk[0] - maxDays)); - List result = await timePartitionedSampler.Object.GetSampleAsync(); + int expectedResultCount = CalculateExpectedTimePartitionedResultCount( + _sortedTimespansNamePk, + groupCount ?? DEFAULT_TIME_GROUP_COUNT, + numberOfRecordsPerGroup ?? DEFAULT_RECORDS_PER_TIME_GROUP, + timeWindowInSeconds); - // We're relying on a delay to create records with different timestamps. - // However, this can cause the actual result to intermittently vary by one record in some cases, particularly in pipelines. - // To prevent these tests from becoming flaky, the assertion has been adjusted. - Assert.IsTrue(expectedResultCount == result.Count || (expectedResultCount + 1) == result.Count || (expectedResultCount - 1) == result.Count, $"Expected result count is {expectedResultCount} and Actual result count is {result.Count}"); + Assert.AreEqual( + expectedResultCount, + result.Count, + $"The sampled result count should match the populated time groups. Timestamps: {string.Join(", ", _sortedTimespansNamePk)}"); + } + + private static int CalculateExpectedTimePartitionedResultCount( + IReadOnlyList timestamps, + int groupCount, + int numberOfRecordsPerGroup, + int timeWindowInSeconds) + { + long maxTimestamp = timestamps[0]; + long minTimestamp = timeWindowInSeconds > 0 ? maxTimestamp - timeWindowInSeconds : timestamps[^1]; + long rangeSize = (maxTimestamp - minTimestamp) / groupCount; + int expectedResultCount = 0; + + for (int group = 0; group < groupCount; group++) + { + long rangeStart = minTimestamp + (group * rangeSize); + long rangeEnd = group == groupCount - 1 ? maxTimestamp : rangeStart + rangeSize - 1; + int recordsInRange = 0; + + foreach (int timestamp in timestamps) + { + if (timestamp >= rangeStart && timestamp <= rangeEnd) + { + recordsInRange++; + } + } + + expectedResultCount += System.Math.Min(numberOfRecordsPerGroup, recordsInRange); + } + + return expectedResultCount; } /// diff --git a/src/Service.Tests/UnitTests/CosmosSamplerHelperTests.cs b/src/Service.Tests/UnitTests/CosmosSamplerHelperTests.cs index d168beb4c7..9618305eff 100644 --- a/src/Service.Tests/UnitTests/CosmosSamplerHelperTests.cs +++ b/src/Service.Tests/UnitTests/CosmosSamplerHelperTests.cs @@ -4,8 +4,11 @@ using System; using System.Collections.Generic; using System.IO; +using System.Linq; using System.Net; using System.Text; +using System.Text.Json; +using System.Text.RegularExpressions; using System.Threading; using System.Threading.Tasks; using Azure.DataApiBuilder.Core.Generator.Sampler; @@ -47,6 +50,31 @@ await Assert.ThrowsExceptionAsync(() => executor.ExecuteQueryAsync("SELECT * FROM c")); } + [TestMethod] + public async Task TimePartitionedSampler_GappedTimestampsReturnsOnlyPopulatedGroups() + { + int[] timestamps = { 100, 101, 103, 104, 106, 107, 109, 110 }; + Mock container = new(); + container.Setup(x => x.GetItemQueryStreamIterator( + It.IsAny(), + It.IsAny(), + It.IsAny())) + .Returns((QueryDefinition query, string _, QueryRequestOptions _) => + CreateIterator(CreateSamplerResponse(query.QueryText, timestamps))); + Mock sampler = new( + container.Object, + null, + null, + null, + Mock.Of()); + sampler.Setup(x => x.GetTimeStampThreshold()).Returns(100); + + List result = await sampler.Object.GetSampleAsync(); + + Assert.AreEqual(8, result.Count); + result.ForEach(document => document.Dispose()); + } + private static CosmosExecutor CreateExecutor(ResponseMessage response) { Mock iterator = new(); @@ -61,6 +89,36 @@ private static CosmosExecutor CreateExecutor(ResponseMessage response) return new CosmosExecutor(container.Object, Mock.Of()); } + private static FeedIterator CreateIterator(string content) + { + ResponseMessage response = new(HttpStatusCode.OK) + { + Content = new MemoryStream(Encoding.UTF8.GetBytes(content)) + }; + Mock iterator = new(); + iterator.SetupSequence(x => x.HasMoreResults).Returns(true).Returns(false); + iterator.Setup(x => x.ReadNextAsync(It.IsAny())).ReturnsAsync(response); + return iterator.Object; + } + + private static string CreateSamplerResponse(string query, IReadOnlyCollection timestamps) + { + if (query == "SELECT VALUE MAX(c._ts) FROM c") + { + return JsonSerializer.Serialize(new { Documents = new[] { timestamps.Max() } }); + } + + System.Text.RegularExpressions.Match range = Regex.Match(query, @"c\._ts >= (?\d+) AND c\._ts <= (?\d+)"); + Assert.IsTrue(range.Success, $"Unexpected sampler query: {query}"); + int start = int.Parse(range.Groups["start"].Value); + int end = int.Parse(range.Groups["end"].Value); + IEnumerable documents = timestamps + .Where(timestamp => timestamp >= start && timestamp <= end) + .Select(timestamp => new { timestamp }); + + return JsonSerializer.Serialize(new { Documents = documents }); + } + private sealed class SampleItem { public int Id { get; set; } From acf38d0e9b0d2739d91e25fc7faefa0717df371c Mon Sep 17 00:00:00 2001 From: Aaron Burtle Date: Sun, 30 Aug 2026 03:26:22 -0700 Subject: [PATCH 18/19] extend branch coverage --- .../AuthorizationResolverUnitTests.cs | 47 +++++ .../SchemaGeneratorFactoryTests.cs | 24 +++ ...aphQLStoredProcedureBuilderHelpersTests.cs | 60 ++++++ .../Sql/SchemaConverterTypeMappingTests.cs | 19 ++ .../Mcp/AggregateRecordsToolTests.cs | 32 +++- .../Mcp/DeleteRecordToolUnitTests.cs | 19 ++ .../Mcp/DynamicCustomToolTests.cs | 11 +- .../Mcp/ExecuteEntityToolTests.cs | 21 +++ src/Service.Tests/Mcp/McpJsonHelperTests.cs | 1 + .../AutoentityConverterCoverageTests.cs | 36 ++++ .../BaseSqlQueryStructureHelperTests.cs | 2 + .../ConfigObjectModelCoverageTests.cs | 21 +++ .../UnitTests/ConfigValidationUnitTests.cs | 96 ++++++++++ .../UnitTests/DmlToolsConfigConverterTests.cs | 75 +++++++- ...raphQLAuthorizationHandlerCoverageTests.cs | 172 ++++++++++++++++++ .../UnitTests/GraphQLFilterParserUnitTests.cs | 111 +++++++++++ .../MetadataProviderFactoryCoverageTests.cs | 82 +++++++++ .../UnitTests/MsSqlQueryBuilderHelperTests.cs | 62 +++++++ .../QueryManagerFactoryCoverageTests.cs | 93 ++++++++++ .../UnitTests/RuntimeConfigHelperTests.cs | 43 +++++ .../RuntimeOptionsConverterCoverageTests.cs | 71 ++++++++ 21 files changed, 1093 insertions(+), 5 deletions(-) create mode 100644 src/Service.Tests/UnitTests/GraphQLAuthorizationHandlerCoverageTests.cs create mode 100644 src/Service.Tests/UnitTests/MetadataProviderFactoryCoverageTests.cs create mode 100644 src/Service.Tests/UnitTests/QueryManagerFactoryCoverageTests.cs diff --git a/src/Service.Tests/Authorization/AuthorizationResolverUnitTests.cs b/src/Service.Tests/Authorization/AuthorizationResolverUnitTests.cs index 9dc9cfd717..81be4e2441 100644 --- a/src/Service.Tests/Authorization/AuthorizationResolverUnitTests.cs +++ b/src/Service.Tests/Authorization/AuthorizationResolverUnitTests.cs @@ -1744,6 +1744,53 @@ public async Task TestClaimsParsingToJson() Assert.AreEqual(expected: "", actual: claimsInRequestContext["nullValuedClaim"]); } + [TestMethod] + public void GetProcessedUserClaims_MultipleClaimsPreserveArrayValueTypes() + { + List claims = new() + { + new("booleans", "true", ClaimValueTypes.Boolean), + new("booleans", "false", ClaimValueTypes.Boolean), + new("integers", "-1", ClaimValueTypes.Integer), + new("integers", "2", ClaimValueTypes.Integer), + new("integer32s", "-3", ClaimValueTypes.Integer32), + new("integer32s", "4", ClaimValueTypes.Integer32), + new("uinteger32s", "5", ClaimValueTypes.UInteger32), + new("uinteger32s", "6", ClaimValueTypes.UInteger32), + new("integer64s", "-7", ClaimValueTypes.Integer64), + new("integer64s", "8", ClaimValueTypes.Integer64), + new("uinteger64s", "9", ClaimValueTypes.UInteger64), + new("uinteger64s", "10", ClaimValueTypes.UInteger64), + new("doubles", "11", ClaimValueTypes.Double), + new("doubles", "12", ClaimValueTypes.Double), + new("strings", "first", ClaimValueTypes.String), + new("strings", "second", ClaimValueTypes.String), + new("jsonNulls", "null", JsonClaimValueTypes.JsonNull), + new("jsonNulls", "null", JsonClaimValueTypes.JsonNull), + new("jsonObjects", "{\"id\":1}", JsonClaimValueTypes.Json), + new("jsonObjects", "{\"id\":2}", JsonClaimValueTypes.Json), + new("customs", "alpha", ClaimValueTypes.DateTime), + new("customs", "beta", ClaimValueTypes.DateTime) + }; + ClaimsIdentity identity = new(claims, TEST_AUTHENTICATION_TYPE, TEST_CLAIMTYPE_NAME, AuthenticationOptions.ROLE_CLAIM_TYPE); + DefaultHttpContext context = new() { User = new ClaimsPrincipal(identity) }; + + Dictionary processedClaims = AuthorizationResolver.GetProcessedUserClaims(context); + + Assert.AreEqual("[true,false]", processedClaims["booleans"]); + Assert.AreEqual("[-1,2]", processedClaims["integers"]); + Assert.AreEqual("[-3,4]", processedClaims["integer32s"]); + Assert.AreEqual("[5,6]", processedClaims["uinteger32s"]); + Assert.AreEqual("[-7,8]", processedClaims["integer64s"]); + Assert.AreEqual("[9,10]", processedClaims["uinteger64s"]); + Assert.AreEqual("[11,12]", processedClaims["doubles"]); + Assert.AreEqual("[\"first\",\"second\"]", processedClaims["strings"]); + Assert.AreEqual("[\"null\",\"null\"]", processedClaims["jsonNulls"]); + Assert.AreEqual("[\"{\\u0022id\\u0022:1}\",\"{\\u0022id\\u0022:2}\"]", processedClaims["jsonObjects"]); + Assert.AreEqual("[\"alpha\",\"beta\"]", processedClaims["customs"]); + Assert.AreEqual(0, AuthorizationResolver.GetProcessedUserClaims(null).Count); + } + /// /// JWT token JSON payloads may not be flat and may contain nested JSON objects or arrays. /// This test validates that when dotnet's JWT processing code flattens the JWT token payload diff --git a/src/Service.Tests/CosmosTests/SchemaGeneratorFactoryTests.cs b/src/Service.Tests/CosmosTests/SchemaGeneratorFactoryTests.cs index f10ac17354..b6af89236f 100644 --- a/src/Service.Tests/CosmosTests/SchemaGeneratorFactoryTests.cs +++ b/src/Service.Tests/CosmosTests/SchemaGeneratorFactoryTests.cs @@ -163,6 +163,30 @@ public async Task ExportGraphQLFromCosmosDB_GeneratesSchemaSuccessfully(string g } + [DataTestMethod] + [DataRow(0, null, null)] + [DataRow(null, 0, null)] + [DataRow(null, null, 0)] + [DataRow(1, 0, 0)] + [DataRow(1, 1, 0)] + public async Task Create_InvalidSamplingCountsThrow(int? days, int? groupCount, int? sampleCount) + { + RuntimeConfig runtimeConfig = new( + Schema: "schema", + DataSource: new DataSource(DatabaseType.CosmosDB_NoSQL, "noop", new Dictionary()), + Entities: new(new Dictionary())); + + await Assert.ThrowsExceptionAsync(() => SchemaGeneratorFactory.Create( + runtimeConfig, + SamplingModes.TopNExtractor.ToString(), + sampleCount, + partitionKeyPath: null, + days, + groupCount, + Mock.Of(), + Mock.Of())); + } + /// /// Creates a mock response message containing JSON documents. /// diff --git a/src/Service.Tests/GraphQLBuilder/Sql/GraphQLStoredProcedureBuilderHelpersTests.cs b/src/Service.Tests/GraphQLBuilder/Sql/GraphQLStoredProcedureBuilderHelpersTests.cs index f6c048f9a9..cdf1013512 100644 --- a/src/Service.Tests/GraphQLBuilder/Sql/GraphQLStoredProcedureBuilderHelpersTests.cs +++ b/src/Service.Tests/GraphQLBuilder/Sql/GraphQLStoredProcedureBuilderHelpersTests.cs @@ -1,9 +1,13 @@ // Copyright (c) Microsoft Corporation. // Licensed under the MIT License. +using System; using System.Collections.Generic; +using System.Reflection; using System.Text.Json; +using Azure.DataApiBuilder.Config.DatabasePrimitives; using Azure.DataApiBuilder.Service.GraphQLBuilder; +using Azure.DataApiBuilder.Service.Exceptions; using HotChocolate.Language; using Microsoft.VisualStudio.TestTools.UnitTesting; @@ -55,5 +59,61 @@ public void GetDefaultResultFieldForStoredProcedure_ReturnsResultStringField() Assert.IsInstanceOfType(field.Type, typeof(NamedTypeNode)); Assert.AreEqual("String", ((NamedTypeNode)field.Type).Name.Value); } + + [DataTestMethod] + [DataRow(typeof(Guid), "d2719f98-e062-4ae8-a786-4ea9c3524d7c", "UUID")] + [DataRow(typeof(byte), "255", "UnsignedByte")] + [DataRow(typeof(short), "-32768", "Short")] + [DataRow(typeof(int), "-2147483648", "Int")] + [DataRow(typeof(long), "9223372036854775807", "Long")] + [DataRow(typeof(float), "1.25", "Single")] + [DataRow(typeof(double), "2.5", "Float")] + [DataRow(typeof(decimal), "3.75", "Decimal")] + [DataRow(typeof(string), "text", "String")] + [DataRow(typeof(bool), "true", "Boolean")] + [DataRow(typeof(DateTime), "2025-01-02T03:04:05Z", "DateTime")] + [DataRow(typeof(byte[]), "AQID", "Base64String")] + [DataRow(typeof(TimeOnly), "12:34:56", "LocalTime")] + public void ConvertValueToGraphQLType_ConvertsEverySupportedScalar( + Type systemType, + string configuredValue, + string expectedGraphQLType) + { + Tuple result = InvokeConvertValueToGraphQLType(configuredValue, systemType); + + Assert.AreEqual(expectedGraphQLType, result.Item1); + Assert.IsNotNull(result.Item2); + } + + [DataTestMethod] + [DataRow("1", true)] + [DataRow("0", false)] + [DataRow("TrUe", true)] + [DataRow("FaLsE", false)] + public void ConvertValueToGraphQLType_ConvertsSupportedBooleanRepresentations(string configuredValue, bool expected) + { + Tuple result = InvokeConvertValueToGraphQLType(configuredValue, typeof(bool)); + + Assert.AreEqual(expected, ((BooleanValueNode)result.Item2).Value); + } + + [TestMethod] + public void ConvertValueToGraphQLType_InvalidValueWrapsConversionFailure() + { + TargetInvocationException exception = Assert.ThrowsException( + () => InvokeConvertValueToGraphQLType("not-a-boolean", typeof(bool))); + + Assert.IsInstanceOfType(exception.InnerException); + } + + private static Tuple InvokeConvertValueToGraphQLType(string configuredValue, Type systemType) + { + MethodInfo method = typeof(GraphQLStoredProcedureBuilder).GetMethod( + "ConvertValueToGraphQLType", + BindingFlags.NonPublic | BindingFlags.Static)!; + ParameterDefinition parameter = new() { SystemType = systemType }; + + return (Tuple)method.Invoke(null, new object[] { configuredValue, parameter })!; + } } } diff --git a/src/Service.Tests/GraphQLBuilder/Sql/SchemaConverterTypeMappingTests.cs b/src/Service.Tests/GraphQLBuilder/Sql/SchemaConverterTypeMappingTests.cs index ca6f030c10..a562ebcafc 100644 --- a/src/Service.Tests/GraphQLBuilder/Sql/SchemaConverterTypeMappingTests.cs +++ b/src/Service.Tests/GraphQLBuilder/Sql/SchemaConverterTypeMappingTests.cs @@ -80,6 +80,25 @@ public void CreateValueNodeFromDbObjectMetadata_ScalarTypes_WrapsInObjectFieldNo Assert.AreEqual(expectedTypeName, objectValueNode.Fields[0].Name.Value); } + [TestMethod] + public void CreateValueNodeFromDbObjectMetadata_RemainingScalarTypes_WrapInObjectFieldNode() + { + object[] values = + { + Guid.Parse("d2719f98-e062-4ae8-a786-4ea9c3524d7c"), + 1.5m, + new DateTimeOffset(2025, 1, 2, 3, 4, 5, TimeSpan.Zero), + new DateTime(2025, 1, 2, 3, 4, 5, DateTimeKind.Unspecified), + new DateTime(2025, 1, 2, 3, 4, 5, DateTimeKind.Utc), + new byte[] { 1, 2, 3 } + }; + + foreach (object value in values) + { + Assert.IsInstanceOfType(SchemaConverter.CreateValueNodeFromDbObjectMetadata(value)); + } + } + [TestMethod] public void CreateValueNodeFromDbObjectMetadata_UnsupportedType_Throws() { diff --git a/src/Service.Tests/Mcp/AggregateRecordsToolTests.cs b/src/Service.Tests/Mcp/AggregateRecordsToolTests.cs index 4d2ac2f9c2..5ffa6eaae7 100644 --- a/src/Service.Tests/Mcp/AggregateRecordsToolTests.cs +++ b/src/Service.Tests/Mcp/AggregateRecordsToolTests.cs @@ -235,7 +235,15 @@ public void BuildAggregationStructure_AddsGroupsAggregationHavingAndPaginationSt distinct: true, first: 2, groupby: new() { "category" }, - havingOperators: new() { ["gt"] = 10, ["lte"] = 100 }, + havingOperators: new() + { + ["eq"] = 10, + ["neq"] = 20, + ["gt"] = 30, + ["gte"] = 40, + ["lt"] = 50, + ["lte"] = 60 + }, havingInValues: new() { 20, 40 }); InvokePrivate( @@ -253,11 +261,31 @@ public void BuildAggregationStructure_AddsGroupsAggregationHavingAndPaginationSt Assert.AreEqual(1, structure.GroupByMetadata.Aggregations.Count); Assert.IsTrue(structure.GroupByMetadata.RequestedAggregations); Assert.IsNotNull(structure.GroupByMetadata.Aggregations[0].HavingPredicates); - Assert.AreEqual(4, structure.Parameters.Count); + Assert.AreEqual(8, structure.Parameters.Count); Assert.IsTrue(structure.IsListQuery); Assert.AreEqual(0, structure.OrderByColumns.Count); } + [TestMethod] + public void BuildAggregationStructure_InvalidHavingOperator_ThrowsArgumentException() + { + SqlQueryStructure structure = CreateUninitializedQueryStructure(); + AggregateRecordsTool.AggregateArguments args = CreateAggregateArguments( + havingOperators: new() { ["invalid"] = 10 }); + + TargetInvocationException exception = Assert.ThrowsException(() => InvokePrivate( + "BuildAggregationStructure", + args, + structure, + new DatabaseTable("dbo", "books"), + "book_price", + "sum_price", + "Book", + Mock.Of())); + + Assert.IsInstanceOfType(exception.InnerException); + } + [TestMethod] public void BuildAggregationStructure_InvalidGroupByMapping_ThrowsDataApiBuilderException() { diff --git a/src/Service.Tests/Mcp/DeleteRecordToolUnitTests.cs b/src/Service.Tests/Mcp/DeleteRecordToolUnitTests.cs index 7bfc3fd453..58c655afce 100644 --- a/src/Service.Tests/Mcp/DeleteRecordToolUnitTests.cs +++ b/src/Service.Tests/Mcp/DeleteRecordToolUnitTests.cs @@ -20,6 +20,7 @@ using Azure.DataApiBuilder.Core.Services.MetadataProviders; using Azure.DataApiBuilder.Mcp.BuiltInTools; using Azure.DataApiBuilder.Service.Exceptions; +using Azure.DataApiBuilder.Service.Tests.SqlTests; using Microsoft.AspNetCore.Http; using Microsoft.AspNetCore.Mvc; using Microsoft.Extensions.DependencyInjection; @@ -172,6 +173,24 @@ public async Task ExecuteAsync_Timeout_ReturnsTimeoutError() AssertErrorType(result, "TimeoutError"); } + [DataTestMethod] + [DataRow(547, "foreign key constraint")] + [DataRow(2627, "unique constraint")] + [DataRow(2601, "unique constraint")] + [DataRow(229, "Permission denied")] + [DataRow(262, "Permission denied")] + [DataRow(208, "not found")] + [DataRow(50000, "Database error")] + public async Task ExecuteAsync_SqlException_MapsErrorNumber(int errorNumber, string expectedMessage) + { + CallToolResult result = await ExecuteAsync( + "{\"entity\":\"Book\",\"keys\":{\"id\":1}}", + SqlTestHelper.CreateSqlException(errorNumber, "provider error")); + + AssertErrorType(result, "DatabaseError"); + StringAssert.Contains(GetText(result), expectedMessage); + } + [TestMethod] public async Task ExecuteAsync_InvalidPrimaryKey_ReturnsDataApiBuilderError() { diff --git a/src/Service.Tests/Mcp/DynamicCustomToolTests.cs b/src/Service.Tests/Mcp/DynamicCustomToolTests.cs index f955fe1111..d59f142776 100644 --- a/src/Service.Tests/Mcp/DynamicCustomToolTests.cs +++ b/src/Service.Tests/Mcp/DynamicCustomToolTests.cs @@ -445,6 +445,10 @@ public async Task ExecuteAsync_ConvertsAllJsonParameterKinds() { Dictionary dbParameters = new() { + ["text"] = new(), + ["integer"] = new(), + ["decimal"] = new(), + ["floatingPoint"] = new(), ["truth"] = new(), ["falsehood"] = new(), ["nothing"] = new(), @@ -456,12 +460,17 @@ public async Task ExecuteAsync_ConvertsAllJsonParameterKinds() context => capturedContext = context); DynamicCustomTool tool = new(TEST_ENTITY, CreateTestStoredProcedureEntity()); using JsonDocument arguments = JsonDocument.Parse( - "{\"truth\":true,\"falsehood\":false,\"nothing\":null,\"complex\":{\"x\":1}}"); + "{\"text\":\"value\",\"integer\":42,\"decimal\":1.5,\"floatingPoint\":1e40," + + "\"truth\":true,\"falsehood\":false,\"nothing\":null,\"complex\":{\"x\":1}}"); CallToolResult result = await tool.ExecuteAsync(arguments, serviceProvider, CancellationToken.None); AssertSuccess(result, "All JSON parameter kinds should be converted."); Assert.IsNotNull(capturedContext); + Assert.AreEqual("value", capturedContext.ResolvedParameters["text"]); + Assert.AreEqual(42L, capturedContext.ResolvedParameters["integer"]); + Assert.AreEqual(1.5m, capturedContext.ResolvedParameters["decimal"]); + Assert.AreEqual(1e40, capturedContext.ResolvedParameters["floatingPoint"]); Assert.AreEqual(true, capturedContext.ResolvedParameters["truth"]); Assert.AreEqual(false, capturedContext.ResolvedParameters["falsehood"]); Assert.IsNull(capturedContext.ResolvedParameters["nothing"]); diff --git a/src/Service.Tests/Mcp/ExecuteEntityToolTests.cs b/src/Service.Tests/Mcp/ExecuteEntityToolTests.cs index 379ba882c1..5500cc49af 100644 --- a/src/Service.Tests/Mcp/ExecuteEntityToolTests.cs +++ b/src/Service.Tests/Mcp/ExecuteEntityToolTests.cs @@ -20,6 +20,7 @@ using Azure.DataApiBuilder.Core.Services.MetadataProviders; using Azure.DataApiBuilder.Mcp.BuiltInTools; using Azure.DataApiBuilder.Service.Exceptions; +using Azure.DataApiBuilder.Service.Tests.SqlTests; using Microsoft.AspNetCore.Http; using Microsoft.AspNetCore.Mvc; using Microsoft.Extensions.DependencyInjection; @@ -451,6 +452,26 @@ public async Task ExecuteEntity_Timeout_ReturnsTimeoutError() AssertError(result, "TimeoutError"); } + [DataTestMethod] + [DataRow(2812, "not found")] + [DataRow(8144, "too many parameters")] + [DataRow(201, "were not supplied")] + [DataRow(245, "Type conversion failed")] + [DataRow(229, "Permission denied")] + [DataRow(262, "Permission denied")] + [DataRow(50000, "Database error")] + public async Task ExecuteEntity_SqlException_MapsErrorNumber(int errorNumber, string expectedMessage) + { + CallToolResult result = await ExecuteWithMockedEngineAsync( + TEST_ENTITY, + new(), + null, + queryOutcome: SqlTestHelper.CreateSqlException(errorNumber, "provider error")); + + AssertError(result, "DatabaseError"); + StringAssert.Contains(GetFirstText(result), expectedMessage); + } + [TestMethod] public async Task ExecuteEntity_BadRequestResult_ReturnsError() { diff --git a/src/Service.Tests/Mcp/McpJsonHelperTests.cs b/src/Service.Tests/Mcp/McpJsonHelperTests.cs index b7f503d807..8c070b8e07 100644 --- a/src/Service.Tests/Mcp/McpJsonHelperTests.cs +++ b/src/Service.Tests/Mcp/McpJsonHelperTests.cs @@ -29,6 +29,7 @@ public void GetJsonValue_Number_ReturnsDecimal() // McpJsonHelper prefers decimal for maximum precision. Assert.AreEqual(42m, McpJsonHelper.GetJsonValue(Value("42"))); Assert.AreEqual(3.14m, McpJsonHelper.GetJsonValue(Value("3.14"))); + Assert.AreEqual(1e40, McpJsonHelper.GetJsonValue(Value("1e40"))); } [TestMethod] diff --git a/src/Service.Tests/UnitTests/AutoentityConverterCoverageTests.cs b/src/Service.Tests/UnitTests/AutoentityConverterCoverageTests.cs index f8193642fb..96cf455459 100644 --- a/src/Service.Tests/UnitTests/AutoentityConverterCoverageTests.cs +++ b/src/Service.Tests/UnitTests/AutoentityConverterCoverageTests.cs @@ -60,5 +60,41 @@ public void AutoentityPatterns_NonArrayPatternThrows(string json) Assert.ThrowsException(() => JsonSerializer.Deserialize(json, Options)); } + + [TestMethod] + public void Autoentity_Write_EvaluatesEveryUserProvidedPatternFlag() + { + AutoentityPatterns[] patterns = + { + new(Include: new[] { "dbo.*" }), + new(Exclude: new[] { "dbo.internal_*" }), + new(Name: "generated_{object}") + }; + + foreach (AutoentityPatterns pattern in patterns) + { + string json = JsonSerializer.Serialize(new Autoentity(pattern, Template: null, Permissions: null), Options); + StringAssert.Contains(json, "\"patterns\""); + } + } + + [TestMethod] + public void Autoentity_Write_EvaluatesEveryUserProvidedTemplateFlag() + { + AutoentityTemplate[] templates = + { + new(Rest: new EntityRestOptions()), + new(GraphQL: new EntityGraphQLOptions("Book", "Books")), + new(Mcp: new EntityMcpOptions(customToolEnabled: true, dmlToolsEnabled: null)), + new(Health: new EntityHealthCheckConfig(enabled: true)), + new(Cache: new EntityCacheOptions(Enabled: true)) + }; + + foreach (AutoentityTemplate template in templates) + { + string json = JsonSerializer.Serialize(new Autoentity(Patterns: null, template, Permissions: null), Options); + StringAssert.Contains(json, "\"template\""); + } + } } } diff --git a/src/Service.Tests/UnitTests/BaseSqlQueryStructureHelperTests.cs b/src/Service.Tests/UnitTests/BaseSqlQueryStructureHelperTests.cs index 0d05eb43ef..592ac5b1a5 100644 --- a/src/Service.Tests/UnitTests/BaseSqlQueryStructureHelperTests.cs +++ b/src/Service.Tests/UnitTests/BaseSqlQueryStructureHelperTests.cs @@ -61,6 +61,8 @@ public void ParseParamAsSystemType_ParsesDatesAndArrays() DateTimeOffset offset = (DateTimeOffset)InvokeParse("2025-01-02T12:00:00+03:00", typeof(DateTimeOffset)); Assert.AreEqual(TimeSpan.FromHours(3), offset.Offset); + Assert.AreEqual(new TimeOnly(12, 34, 56), InvokeParse("12:34:56", typeof(TimeSpan))); + object[] values = (object[])InvokeParse("[1.5,2.25]", typeof(float[])); CollectionAssert.AreEqual(new object[] { 1.5f, 2.25f }, values); } diff --git a/src/Service.Tests/UnitTests/ConfigObjectModelCoverageTests.cs b/src/Service.Tests/UnitTests/ConfigObjectModelCoverageTests.cs index fef4479199..224ce98054 100644 --- a/src/Service.Tests/UnitTests/ConfigObjectModelCoverageTests.cs +++ b/src/Service.Tests/UnitTests/ConfigObjectModelCoverageTests.cs @@ -91,6 +91,27 @@ public void EmbeddingsOptions_NullOptionalFeaturesUseDocumentedFallbacks() Assert.IsFalse(options.IsLevel2CacheEnabled); } + [TestMethod] + public void EmbeddingsOptions_Level2CacheRequiresBothCacheLevels() + { + EmbeddingsOptions disabled = new(EmbeddingProviderType.OpenAI, "https://example.com", "key") + { + Cache = new EmbeddingsCacheOptions(Enabled: false, Level2: new EmbeddingsCacheLevel2Options(Enabled: true)) + }; + EmbeddingsOptions level2Disabled = new(EmbeddingProviderType.OpenAI, "https://example.com", "key") + { + Cache = new EmbeddingsCacheOptions(Enabled: true, Level2: new EmbeddingsCacheLevel2Options(Enabled: false)) + }; + EmbeddingsOptions enabled = new(EmbeddingProviderType.OpenAI, "https://example.com", "key") + { + Cache = new EmbeddingsCacheOptions(Enabled: true, Level2: new EmbeddingsCacheLevel2Options(Enabled: true)) + }; + + Assert.IsFalse(disabled.IsLevel2CacheEnabled); + Assert.IsFalse(level2Disabled.IsLevel2CacheEnabled); + Assert.IsTrue(enabled.IsLevel2CacheEnabled); + } + [TestMethod] public void EmbeddingsEndpointOptions_DefaultConstructorDisablesEndpoint() { diff --git a/src/Service.Tests/UnitTests/ConfigValidationUnitTests.cs b/src/Service.Tests/UnitTests/ConfigValidationUnitTests.cs index 680a354729..60eaa3b140 100644 --- a/src/Service.Tests/UnitTests/ConfigValidationUnitTests.cs +++ b/src/Service.Tests/UnitTests/ConfigValidationUnitTests.cs @@ -631,6 +631,92 @@ public void TestRelationshipWithNoLinkingObjectAndEitherSourceOrTargetFieldIsNul configValidator.ValidateRelationships(runtimeConfig, _metadataProviderFactory.Object); } + [DataTestMethod] + [DataRow(false, "explicit")] + [DataRow(false, "forward")] + [DataRow(false, "reverse")] + [DataRow(false, "none")] + [DataRow(true, "explicit")] + [DataRow(true, "forward")] + [DataRow(true, "none")] + public void ValidateRelationships_LoggingResolvesConfiguredAndInferredColumns(bool useLinkingObject, string resolution) + { + string[] sourceFields = resolution == "explicit" ? new[] { "source_id" } : null; + string[] targetFields = resolution == "explicit" ? new[] { "target_id" } : null; + string[] linkingSourceFields = useLinkingObject && resolution == "explicit" ? new[] { "link_source_id" } : null; + string[] linkingTargetFields = useLinkingObject && resolution == "explicit" ? new[] { "link_target_id" } : null; + string linkingObjectName = useLinkingObject ? "dbo.LINKING_TABLE" : null; + EntityRelationship relationship = new( + Cardinality: Cardinality.One, + TargetEntity: "Target", + SourceFields: sourceFields, + TargetFields: targetFields, + LinkingObject: linkingObjectName, + LinkingSourceFields: linkingSourceFields, + LinkingTargetFields: linkingTargetFields); + Dictionary entities = new() + { + ["Source"] = GetSampleEntityUsingSourceAndRelationshipMap( + "SOURCE_TABLE", + new Dictionary { ["relationship"] = relationship }, + new EntityGraphQLOptions("Source", "Sources", true)), + ["Target"] = GetSampleEntityUsingSourceAndRelationshipMap( + "TARGET_TABLE", + relationshipMap: null, + new EntityGraphQLOptions("Target", "Targets", true)) + }; + RuntimeConfig runtimeConfig = new( + Schema: "UnitTestSchema", + DataSource: new DataSource(DatabaseType.MSSQL, string.Empty), + Runtime: new RuntimeOptions(new(), new(), new(), new(null, null)), + Entities: new RuntimeEntities(entities)); + + DatabaseTable sourceTable = new("dbo", "SOURCE_TABLE"); + DatabaseTable targetTable = new("dbo", "TARGET_TABLE"); + DatabaseTable linkingTable = new("dbo", "LINKING_TABLE"); + RelationShipPair sourceTarget = new(sourceTable, targetTable); + RelationShipPair targetSource = new(targetTable, sourceTable); + RelationShipPair linkingSource = new(linkingTable, sourceTable); + RelationShipPair linkingTarget = new(linkingTable, targetTable); + Dictionary foreignKeys = new(); + if (resolution == "forward") + { + if (useLinkingObject) + { + foreignKeys[linkingSource] = CreateForeignKey(linkingSource); + foreignKeys[linkingTarget] = CreateForeignKey(linkingTarget); + } + else + { + foreignKeys[sourceTarget] = CreateForeignKey(sourceTarget); + } + } + else if (resolution == "reverse") + { + foreignKeys[targetSource] = CreateForeignKey(targetSource); + } + + Mock metadataProvider = new(); + metadataProvider.SetupGet(x => x.EntityToDatabaseObject).Returns(new Dictionary + { + ["Source"] = sourceTable, + ["Target"] = targetTable + }); + metadataProvider.SetupGet(x => x.PairToFkDefinition).Returns(foreignKeys); + metadataProvider.Setup(x => x.ParseSchemaAndDbTableName(linkingObjectName)).Returns(("dbo", "LINKING_TABLE")); + metadataProvider.Setup(x => x.VerifyForeignKeyExistsInDB(It.IsAny(), It.IsAny())).Returns(true); + string exposedField = string.Empty; + metadataProvider.Setup(x => x.TryGetExposedColumnName(It.IsAny(), It.IsAny(), out exposedField)).Returns(true); + Mock metadataProviderFactory = new(); + metadataProviderFactory.Setup(x => x.GetMetadataProvider(It.IsAny())).Returns(metadataProvider.Object); + + MockFileSystem fileSystem = new(); + RuntimeConfigProvider provider = new(new FileSystemRuntimeConfigLoader(fileSystem)); + RuntimeConfigValidator validator = new(provider, fileSystem, Mock.Of>()); + + validator.ValidateRelationships(runtimeConfig, metadataProviderFactory.Object); + } + /// /// Test method that ensures our validation code catches the cases where source and target fields do not match in some way /// and the linking object is null, indicating we have a one-many or many-one relationship. @@ -3780,6 +3866,16 @@ string[] linkingTargetFields }; return entityMap; } + + private static ForeignKeyDefinition CreateForeignKey(RelationShipPair pair) + { + return new ForeignKeyDefinition + { + Pair = pair, + ReferencingColumns = new() { "referencing_id" }, + ReferencedColumns = new() { "referenced_id" } + }; + } } } diff --git a/src/Service.Tests/UnitTests/DmlToolsConfigConverterTests.cs b/src/Service.Tests/UnitTests/DmlToolsConfigConverterTests.cs index 7887dca716..2b7ca85707 100644 --- a/src/Service.Tests/UnitTests/DmlToolsConfigConverterTests.cs +++ b/src/Service.Tests/UnitTests/DmlToolsConfigConverterTests.cs @@ -4,6 +4,7 @@ #nullable disable using System.IO; +using System.Linq; using System.Text; using System.Text.Json; using Azure.DataApiBuilder.Config; @@ -181,9 +182,44 @@ public void Deserialize_UnknownProperty_IsSkipped() } [TestMethod] - public void Deserialize_NonBooleanForKnownProperty_ThrowsJsonException() + public void Deserialize_UnknownNonBooleanProperty_IsSkipped() { - string json = @"{ ""create-record"": ""yes"" }"; + DmlToolsConfig config = JsonSerializer.Deserialize( + "{\"unknown-tool\":{\"nested\":true},\"read-records\":false}", + GetOptions()); + + Assert.IsNotNull(config); + Assert.IsFalse(config.ReadRecords); + } + + [DataTestMethod] + [DataRow("describe-entities")] + [DataRow("create-record")] + [DataRow("read-records")] + [DataRow("update-record")] + [DataRow("delete-record")] + [DataRow("execute-entity")] + public void Deserialize_EachBooleanToolProperty_AppliesOverride(string propertyName) + { + string json = $"{{\"{propertyName}\":false}}"; + + DmlToolsConfig config = JsonSerializer.Deserialize(json, GetOptions()); + JObject serialized = JObject.Parse(SerializeWithinObject(config)); + + Assert.IsNotNull(config); + Assert.IsFalse(serialized["dml-tools"][propertyName].Value()); + } + + [DataTestMethod] + [DataRow("describe-entities")] + [DataRow("create-record")] + [DataRow("read-records")] + [DataRow("update-record")] + [DataRow("delete-record")] + [DataRow("execute-entity")] + public void Deserialize_NonBooleanForEachKnownProperty_ThrowsJsonException(string propertyName) + { + string json = $"{{\"{propertyName}\":\"yes\"}}"; Assert.ThrowsException( () => JsonSerializer.Deserialize(json, GetOptions())); @@ -198,6 +234,41 @@ public void Deserialize_AggregateRecordsInvalidType_ThrowsJsonException() () => JsonSerializer.Deserialize(json, GetOptions())); } + [TestMethod] + public void Serialize_AllIndividualSettings_WritesEveryProperty() + { + DmlToolsConfig config = new( + describeEntities: false, + createRecord: true, + readRecords: false, + updateRecord: true, + deleteRecord: false, + executeEntity: true, + aggregateRecords: false); + + JObject dmlTools = JObject.Parse(SerializeWithinObject(config))["dml-tools"].Value(); + + Assert.AreEqual(7, dmlTools.Properties().Count()); + Assert.IsFalse(dmlTools["describe-entities"].Value()); + Assert.IsTrue(dmlTools["create-record"].Value()); + Assert.IsFalse(dmlTools["read-records"].Value()); + Assert.IsTrue(dmlTools["update-record"].Value()); + Assert.IsFalse(dmlTools["delete-record"].Value()); + Assert.IsTrue(dmlTools["execute-entity"].Value()); + Assert.IsFalse(dmlTools["aggregate-records"].Value()); + } + + [TestMethod] + public void Serialize_TimeoutWithoutAggregateSetting_OmitsEnabled() + { + DmlToolsConfig config = new(aggregateRecordsQueryTimeout: 90) { AggregateRecords = null }; + + JObject aggregate = JObject.Parse(SerializeWithinObject(config))["dml-tools"]["aggregate-records"].Value(); + + Assert.IsFalse(aggregate.ContainsKey("enabled")); + Assert.AreEqual(90, aggregate["query-timeout"].Value()); + } + [TestMethod] public void Serialize_FromBooleanTrue_WritesBooleanForm() { diff --git a/src/Service.Tests/UnitTests/GraphQLAuthorizationHandlerCoverageTests.cs b/src/Service.Tests/UnitTests/GraphQLAuthorizationHandlerCoverageTests.cs new file mode 100644 index 0000000000..f7499a5137 --- /dev/null +++ b/src/Service.Tests/UnitTests/GraphQLAuthorizationHandlerCoverageTests.cs @@ -0,0 +1,172 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +using System; +using System.Collections.Generic; +using System.Linq; +using System.Reflection; +using System.Security.Claims; +using Azure.DataApiBuilder.Auth; +using Azure.DataApiBuilder.Core.Authorization; +using HotChocolate.Authorization; +using HotChocolate.Resolvers; +using Microsoft.AspNetCore.Http; +using Microsoft.VisualStudio.TestTools.UnitTesting; +using Moq; + +namespace Azure.DataApiBuilder.Service.Tests.UnitTests +{ + [TestClass] + public class GraphQLAuthorizationHandlerCoverageTests + { + [TestMethod] + public void AuthorizeAsync_EvaluatesAuthenticationHeaderRoleAndPolicyBranches() + { + Mock resolver = new(); + resolver.Setup(x => x.IsRoleAllowedByDirective("reader", It.IsAny?>())) + .Returns(true); + GraphQLAuthorizationHandler handler = new(resolver.Object); + AuthorizeDirective roleDirective = CreateDirective(new[] { "reader" }, policy: null); + AuthorizeDirective policyDirective = CreateDirective(new[] { "reader" }, policy: "unsupported-policy"); + + Assert.AreEqual( + AuthorizeResult.NotAuthenticated, + handler.AuthorizeAsync(CreateContext(authenticated: false, includeHttpContext: false), roleDirective).Result); + Assert.AreEqual( + AuthorizeResult.NotAllowed, + handler.AuthorizeAsync(CreateContext(authenticated: true, includeHttpContext: false), roleDirective).Result); + Assert.AreEqual( + AuthorizeResult.NotAllowed, + handler.AuthorizeAsync(CreateContext(authenticated: true, includeHttpContext: true), roleDirective).Result); + Assert.AreEqual( + AuthorizeResult.Allowed, + handler.AuthorizeAsync(CreateContext(authenticated: true, includeHttpContext: true, role: "reader"), roleDirective).Result); + Assert.AreEqual( + AuthorizeResult.NotAllowed, + handler.AuthorizeAsync(CreateContext(authenticated: true, includeHttpContext: true, role: "reader"), policyDirective).Result); + + resolver.Setup(x => x.IsRoleAllowedByDirective("denied", It.IsAny?>())) + .Returns(false); + Assert.AreEqual( + AuthorizeResult.NotAllowed, + handler.AuthorizeAsync(CreateContext(authenticated: true, includeHttpContext: true, role: "denied"), roleDirective).Result); + } + + [TestMethod] + public void AuthorizeAsync_DirectiveListEvaluatesEveryBranch() + { + Mock resolver = new(); + resolver.Setup(x => x.IsRoleAllowedByDirective("reader", It.IsAny?>())) + .Returns(true); + resolver.Setup(x => x.IsRoleAllowedByDirective("denied", It.IsAny?>())) + .Returns(false); + GraphQLAuthorizationHandler handler = new(resolver.Object); + AuthorizeDirective roleDirective = CreateDirective(new[] { "reader" }, policy: null); + AuthorizeDirective policyDirective = CreateDirective(new[] { "reader" }, policy: "unsupported-policy"); + + Assert.AreEqual( + AuthorizeResult.NotAuthenticated, + handler.AuthorizeAsync( + CreateAuthorizationContext(authenticated: false, includeHttpContext: false), + new[] { roleDirective }).Result); + Assert.AreEqual( + AuthorizeResult.Allowed, + handler.AuthorizeAsync( + CreateAuthorizationContext(authenticated: true, includeHttpContext: false), + Array.Empty()).Result); + Assert.AreEqual( + AuthorizeResult.NotAllowed, + handler.AuthorizeAsync( + CreateAuthorizationContext(authenticated: true, includeHttpContext: false), + new[] { roleDirective }).Result); + Assert.AreEqual( + AuthorizeResult.NotAllowed, + handler.AuthorizeAsync( + CreateAuthorizationContext(authenticated: true, includeHttpContext: true, role: "denied"), + new[] { roleDirective }).Result); + Assert.AreEqual( + AuthorizeResult.NotAllowed, + handler.AuthorizeAsync( + CreateAuthorizationContext(authenticated: true, includeHttpContext: true, role: "reader"), + new[] { policyDirective }).Result); + Assert.AreEqual( + AuthorizeResult.Allowed, + handler.AuthorizeAsync( + CreateAuthorizationContext(authenticated: true, includeHttpContext: true, role: "reader"), + new[] { roleDirective, roleDirective }).Result); + } + + private static IMiddlewareContext CreateContext(bool authenticated, bool includeHttpContext, string? role = null) + { + Dictionary contextData = CreateContextData(authenticated, includeHttpContext, role); + Mock context = new(); + context.SetupGet(x => x.ContextData).Returns(contextData); + return context.Object; + } + + private static AuthorizationContext CreateAuthorizationContext(bool authenticated, bool includeHttpContext, string? role = null) + { + Dictionary contextData = CreateContextData(authenticated, includeHttpContext, role); + ConstructorInfo constructor = typeof(AuthorizationContext) + .GetConstructors(BindingFlags.Instance | BindingFlags.Public | BindingFlags.NonPublic) + .OrderBy(candidate => candidate.GetParameters().Length) + .First(); + object?[] arguments = constructor.GetParameters() + .Select(parameter => parameter.Name?.Contains("contextData", StringComparison.OrdinalIgnoreCase) == true || + parameter.ParameterType.IsAssignableFrom(contextData.GetType()) + ? contextData + : parameter.HasDefaultValue + ? parameter.DefaultValue + : parameter.ParameterType.IsValueType + ? Activator.CreateInstance(parameter.ParameterType) + : null) + .ToArray(); + + return (AuthorizationContext)constructor.Invoke(arguments); + } + + private static Dictionary CreateContextData(bool authenticated, bool includeHttpContext, string? role) + { + ClaimsIdentity identity = authenticated + ? new ClaimsIdentity(new[] { new Claim(ClaimTypes.Name, "test-user") }, "test-authentication") + : new ClaimsIdentity(); + Dictionary contextData = new() + { + [nameof(ClaimsPrincipal)] = new ClaimsPrincipal(identity) + }; + + if (includeHttpContext) + { + DefaultHttpContext httpContext = new(); + if (role is not null) + { + httpContext.Request.Headers[AuthorizationResolver.CLIENT_ROLE_HEADER] = role; + } + + contextData[nameof(HttpContext)] = httpContext; + } + + return contextData; + } + + private static AuthorizeDirective CreateDirective(IReadOnlyList roles, string? policy) + { + ConstructorInfo constructor = typeof(AuthorizeDirective) + .GetConstructors(BindingFlags.Instance | BindingFlags.Public | BindingFlags.NonPublic) + .OrderByDescending(candidate => candidate.GetParameters().Length) + .First(); + object?[] arguments = constructor.GetParameters() + .Select(parameter => parameter.Name?.ToLowerInvariant() switch + { + "roles" => roles, + "policy" => policy, + _ when parameter.HasDefaultValue => parameter.DefaultValue, + _ when parameter.ParameterType.IsValueType => Activator.CreateInstance(parameter.ParameterType), + _ => null + }) + .ToArray(); + + return (AuthorizeDirective)constructor.Invoke(arguments); + } + } +} diff --git a/src/Service.Tests/UnitTests/GraphQLFilterParserUnitTests.cs b/src/Service.Tests/UnitTests/GraphQLFilterParserUnitTests.cs index 46d31b98de..b4384ed295 100644 --- a/src/Service.Tests/UnitTests/GraphQLFilterParserUnitTests.cs +++ b/src/Service.Tests/UnitTests/GraphQLFilterParserUnitTests.cs @@ -3,12 +3,17 @@ using System.Collections.Generic; using System.Reflection; +using Azure.DataApiBuilder.Config.DatabasePrimitives; using Azure.DataApiBuilder.Config.ObjectModel; using Azure.DataApiBuilder.Core.Configurations; using Azure.DataApiBuilder.Core.Models; using Azure.DataApiBuilder.Core.Services.MetadataProviders; using Azure.DataApiBuilder.Service.Exceptions; +using HotChocolate; +using HotChocolate.Execution; using HotChocolate.Language; +using HotChocolate.Resolvers; +using HotChocolate.Types; using Microsoft.AspNetCore.Http; using Microsoft.VisualStudio.TestTools.UnitTesting; using Moq; @@ -147,6 +152,75 @@ public void PreprocessInOperatorValues_FiltersNullsAndReturnsNullForEmptyValues( Assert.IsNull(InvokePreprocessInOperatorValues(new List())); } + [DataTestMethod] + [DataRow("eq", PredicateOperation.Equal, false, false)] + [DataRow("neq", PredicateOperation.NotEqual, false, false)] + [DataRow("lt", PredicateOperation.LessThan, false, false)] + [DataRow("gt", PredicateOperation.GreaterThan, false, false)] + [DataRow("lte", PredicateOperation.LessThanOrEqual, false, false)] + [DataRow("gte", PredicateOperation.GreaterThanOrEqual, false, false)] + [DataRow("in", PredicateOperation.IN, false, false)] + [DataRow("contains", PredicateOperation.LIKE, false, false)] + [DataRow("contains", PredicateOperation.ARRAY_CONTAINS, true, false)] + [DataRow("notContains", PredicateOperation.NOT_LIKE, false, false)] + [DataRow("notContains", PredicateOperation.NOT_ARRAY_CONTAINS, true, false)] + [DataRow("startsWith", PredicateOperation.LIKE, false, false)] + [DataRow("endsWith", PredicateOperation.LIKE, false, false)] + [DataRow("isNull", PredicateOperation.IS, false, true)] + [DataRow("isNull", PredicateOperation.IS_NOT, false, false)] + public void FieldFilterParser_Parse_MapsEverySupportedOperator( + string operation, + PredicateOperation expectedOperation, + bool isListType, + bool booleanValue) + { + IValueNode value = operation switch + { + "in" => new ListValueNode(new List + { + new StringValueNode("first"), + NullValueNode.Default, + new StringValueNode("second") + }), + "isNull" => new BooleanValueNode(booleanValue), + _ => new StringValueNode(@"value%_[]\") + }; + + Predicate result = FieldFilterParser.Parse( + CreateMiddlewareContext(), + CreateFilterArgumentSchema(), + new Column("dbo", "books", "title"), + new List { new(operation, value) }, + (literal, columnName, lengthOverride) => $"{columnName}:{literal}:{lengthOverride}", + isListType); + + Assert.AreEqual(expectedOperation, result.Op); + } + + [TestMethod] + public void FieldFilterParser_Parse_NullValueIsIgnored() + { + Predicate result = FieldFilterParser.Parse( + CreateMiddlewareContext(), + CreateFilterArgumentSchema(), + new Column("dbo", "books", "title"), + new List { new("eq", NullValueNode.Default) }, + (literal, columnName, lengthOverride) => literal.ToString()!); + + Assert.IsNotNull(result); + } + + [TestMethod] + public void FieldFilterParser_Parse_UnsupportedOperatorThrows() + { + Assert.ThrowsException(() => FieldFilterParser.Parse( + CreateMiddlewareContext(), + CreateFilterArgumentSchema(), + new Column("dbo", "books", "title"), + new List { new("unsupported", new StringValueNode("value")) }, + (literal, columnName, lengthOverride) => literal.ToString()!)); + } + private static GQLFilterParser CreateParserWithDepthLimit(int? depthLimit) { RuntimeConfig config = new( @@ -170,5 +244,42 @@ private static GQLFilterParser CreateParserWithDepthLimit(int? depthLimit) "PreprocessInOperatorValues", BindingFlags.Static | BindingFlags.NonPublic)!.Invoke(null, new[] { value }); } + + private static IMiddlewareContext CreateMiddlewareContext() + { + Mock variables = new(); + Mock context = new(); + context.SetupGet(x => x.Variables).Returns(variables.Object); + return context.Object; + } + + private static IInputValueDefinition CreateFilterArgumentSchema() + { + return SchemaBuilder.New() + .AddDocumentFromString(""" + type Query { + test(filter: TestFilterInput): String + } + + input TestFilterInput { + eq: String + neq: String + lt: String + gt: String + lte: String + gte: String + in: [String] + contains: String + notContains: String + startsWith: String + endsWith: String + isNull: Boolean + unsupported: String + } + """) + .AddResolver("Query", "test", _ => string.Empty) + .Create() + .QueryType.Fields["test"].Arguments["filter"]; + } } } diff --git a/src/Service.Tests/UnitTests/MetadataProviderFactoryCoverageTests.cs b/src/Service.Tests/UnitTests/MetadataProviderFactoryCoverageTests.cs new file mode 100644 index 0000000000..7825829f6c --- /dev/null +++ b/src/Service.Tests/UnitTests/MetadataProviderFactoryCoverageTests.cs @@ -0,0 +1,82 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +using System; +using System.Collections.Generic; +using System.IO.Abstractions.TestingHelpers; +using System.Linq; +using Azure.DataApiBuilder.Config; +using Azure.DataApiBuilder.Config.ObjectModel; +using Azure.DataApiBuilder.Core.Configurations; +using Azure.DataApiBuilder.Core.Resolvers.Factories; +using Azure.DataApiBuilder.Core.Services; +using Azure.DataApiBuilder.Core.Services.MetadataProviders; +using Microsoft.Extensions.Logging; +using Microsoft.VisualStudio.TestTools.UnitTesting; +using Moq; + +namespace Azure.DataApiBuilder.Service.Tests.UnitTests +{ + [TestClass] + public class MetadataProviderFactoryCoverageTests + { + [TestMethod] + public void Constructor_CreatesProviderForEverySupportedSqlDatabaseType() + { + Dictionary dataSources = new() + { + ["mssql"] = new(DatabaseType.MSSQL, string.Empty), + ["dwsql"] = new(DatabaseType.DWSQL, string.Empty), + ["postgresql"] = new(DatabaseType.PostgreSQL, string.Empty), + ["mysql"] = new(DatabaseType.MySQL, string.Empty) + }; + + MetadataProviderFactory factory = CreateFactory(dataSources); + + Assert.AreEqual(dataSources.Count, factory.ListMetadataProviders().Count()); + Assert.AreEqual(2, factory.ListMetadataProviders().OfType().Count()); + Assert.AreEqual(1, factory.ListMetadataProviders().OfType().Count()); + Assert.AreEqual(1, factory.ListMetadataProviders().OfType().Count()); + } + + [TestMethod] + public void Constructor_RejectsUnsupportedDatabaseType() + { + Dictionary dataSources = new() + { + ["unsupported"] = new((DatabaseType)999, string.Empty) + }; + + Assert.ThrowsException(() => CreateFactory(dataSources)); + } + + private static MetadataProviderFactory CreateFactory(Dictionary dataSources) + { + DataSource defaultDataSource = dataSources.First().Value; + RuntimeConfig config = new( + Schema: string.Empty, + DataSource: defaultDataSource, + Runtime: new RuntimeOptions(null, null, null, null), + Entities: new RuntimeEntities(new Dictionary()), + DefaultDataSourceName: dataSources.First().Key, + DataSourceNameToDataSource: dataSources, + EntityNameToDataSourceName: new Dictionary()); + MockFileSystem fileSystem = new(); + FileSystemRuntimeConfigLoader loader = new(fileSystem) { RuntimeConfig = config }; + RuntimeConfigProvider provider = new(loader); + RuntimeConfigValidator validator = new( + provider, + fileSystem, + Mock.Of>()); + + return new MetadataProviderFactory( + provider, + validator, + Mock.Of(), + Mock.Of>(), + fileSystem, + handler: null, + isValidateOnly: true); + } + } +} diff --git a/src/Service.Tests/UnitTests/MsSqlQueryBuilderHelperTests.cs b/src/Service.Tests/UnitTests/MsSqlQueryBuilderHelperTests.cs index 52d77f3d5a..2cc9a5197f 100644 --- a/src/Service.Tests/UnitTests/MsSqlQueryBuilderHelperTests.cs +++ b/src/Service.Tests/UnitTests/MsSqlQueryBuilderHelperTests.cs @@ -4,6 +4,7 @@ using System; using System.Collections.Generic; using System.Data; +using System.Reflection; using Azure.DataApiBuilder.Auth; using Azure.DataApiBuilder.Config.DatabasePrimitives; using Azure.DataApiBuilder.Config.ObjectModel; @@ -16,6 +17,7 @@ using Microsoft.AspNetCore.Http; using Microsoft.VisualStudio.TestTools.UnitTesting; using Moq; +using static Azure.DataApiBuilder.Service.GraphQLBuilder.Sql.SchemaConverter; namespace Azure.DataApiBuilder.Service.Tests.UnitTests { @@ -88,6 +90,44 @@ public void BuildUpsert_AsymmetricTriggerConfigurationUsesCorrectOutputQualifier StringAssert.Contains(query, expectedInsertBranch); } + [DataTestMethod] + [DataRow("alias", "dbo", "[alias].[price]")] + [DataRow(null, "dbo", "[dbo].[books].[price]")] + [DataRow(null, "", "[books].[price]")] + public void BuildColumn_UsesAliasSchemaOrTableQualification(string? tableAlias, string tableSchema, string expected) + { + Column column = new(tableSchema, TABLE_NAME, "price", tableAlias); + + string sql = InvokeBuild(column); + + Assert.AreEqual(expected, sql); + } + + [DataTestMethod] + [DataRow("alias", "dbo", false, false, "sum([alias].[price])")] + [DataRow(null, "dbo", true, false, "sum(DISTINCT ([dbo].[books].[price]))")] + [DataRow(null, "", false, true, "AS [sum_price]")] + public void BuildAggregationColumn_UsesQualificationDistinctAndAlias( + string? tableAlias, + string tableSchema, + bool distinct, + bool useAlias, + string expectedFragment) + { + AggregationColumn column = new( + tableSchema, + TABLE_NAME, + "price", + AggregationType.sum, + "sum_price", + distinct, + tableAlias); + + string sql = InvokeBuild(column, useAlias); + + StringAssert.Contains(sql, expectedFragment); + } + private static SourceDefinition CreateSourceDefinition( bool isInsertTriggerEnabled, bool isUpdateTriggerEnabled, @@ -167,5 +207,27 @@ private static DefaultHttpContext CreateHttpContext() httpContext.Request.Headers[AuthorizationResolver.CLIENT_ROLE_HEADER] = "authenticated"; return httpContext; } + + private static string InvokeBuild(Column column) + { + MethodInfo method = typeof(BaseSqlQueryBuilder).GetMethod( + "Build", + BindingFlags.Instance | BindingFlags.NonPublic, + binder: null, + new[] { typeof(Column) }, + modifiers: null)!; + return (string)method.Invoke(new MsSqlQueryBuilder(), new object[] { column })!; + } + + private static string InvokeBuild(AggregationColumn column, bool useAlias) + { + MethodInfo method = typeof(BaseSqlQueryBuilder).GetMethod( + "Build", + BindingFlags.Instance | BindingFlags.NonPublic, + binder: null, + new[] { typeof(AggregationColumn), typeof(bool) }, + modifiers: null)!; + return (string)method.Invoke(new MsSqlQueryBuilder(), new object[] { column, useAlias })!; + } } } diff --git a/src/Service.Tests/UnitTests/QueryManagerFactoryCoverageTests.cs b/src/Service.Tests/UnitTests/QueryManagerFactoryCoverageTests.cs new file mode 100644 index 0000000000..4641378414 --- /dev/null +++ b/src/Service.Tests/UnitTests/QueryManagerFactoryCoverageTests.cs @@ -0,0 +1,93 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +using System; +using System.Collections.Generic; +using System.IO.Abstractions.TestingHelpers; +using System.Linq; +using Azure.DataApiBuilder.Config; +using Azure.DataApiBuilder.Config.ObjectModel; +using Azure.DataApiBuilder.Core.Configurations; +using Azure.DataApiBuilder.Core.Resolvers; +using Azure.DataApiBuilder.Core.Resolvers.Factories; +using Azure.DataApiBuilder.Service.Exceptions; +using Microsoft.AspNetCore.Http; +using Microsoft.Extensions.Logging; +using Microsoft.VisualStudio.TestTools.UnitTesting; +using Moq; + +namespace Azure.DataApiBuilder.Service.Tests.UnitTests +{ + [TestClass] + public class QueryManagerFactoryCoverageTests + { + [TestMethod] + public void Constructor_CreatesManagersForEverySupportedSqlDatabaseType() + { + Dictionary dataSources = new() + { + ["mssql"] = new(DatabaseType.MSSQL, string.Empty), + ["mssql-duplicate"] = new(DatabaseType.MSSQL, string.Empty), + ["dwsql"] = new(DatabaseType.DWSQL, string.Empty), + ["postgresql"] = new(DatabaseType.PostgreSQL, string.Empty), + ["mysql"] = new(DatabaseType.MySQL, string.Empty), + ["cosmos"] = new(DatabaseType.CosmosDB_NoSQL, string.Empty) + }; + + QueryManagerFactory factory = CreateFactory(dataSources); + + Assert.IsInstanceOfType(factory.GetQueryBuilder(DatabaseType.MSSQL)); + Assert.IsInstanceOfType(factory.GetQueryBuilder(DatabaseType.DWSQL)); + Assert.IsInstanceOfType(factory.GetQueryBuilder(DatabaseType.PostgreSQL)); + Assert.IsInstanceOfType(factory.GetQueryBuilder(DatabaseType.MySQL)); + Assert.IsInstanceOfType(factory.GetQueryExecutor(DatabaseType.MSSQL)); + Assert.IsInstanceOfType(factory.GetDbExceptionParser(DatabaseType.MSSQL)); + } + + [TestMethod] + public void Accessors_RejectUnconfiguredDatabaseType() + { + QueryManagerFactory factory = CreateFactory(new Dictionary + { + ["mssql"] = new(DatabaseType.MSSQL, string.Empty) + }); + DatabaseType missing = (DatabaseType)998; + + Assert.ThrowsException(() => factory.GetQueryBuilder(missing)); + Assert.ThrowsException(() => factory.GetQueryExecutor(missing)); + Assert.ThrowsException(() => factory.GetDbExceptionParser(missing)); + } + + [TestMethod] + public void Constructor_RejectsUnsupportedDatabaseType() + { + Dictionary dataSources = new() + { + ["unsupported"] = new((DatabaseType)999, string.Empty) + }; + + Assert.ThrowsException(() => CreateFactory(dataSources)); + } + + private static QueryManagerFactory CreateFactory(Dictionary dataSources) + { + DataSource defaultDataSource = dataSources.First().Value; + RuntimeConfig config = new( + Schema: string.Empty, + DataSource: defaultDataSource, + Runtime: new RuntimeOptions(null, null, null, null), + Entities: new RuntimeEntities(new Dictionary()), + DefaultDataSourceName: dataSources.First().Key, + DataSourceNameToDataSource: dataSources, + EntityNameToDataSourceName: new Dictionary()); + FileSystemRuntimeConfigLoader loader = new(new MockFileSystem()) { RuntimeConfig = config }; + RuntimeConfigProvider provider = new(loader); + + return new QueryManagerFactory( + provider, + Mock.Of>(), + new HttpContextAccessor(), + handler: null); + } + } +} diff --git a/src/Service.Tests/UnitTests/RuntimeConfigHelperTests.cs b/src/Service.Tests/UnitTests/RuntimeConfigHelperTests.cs index b720aab812..5d4a13e5fa 100644 --- a/src/Service.Tests/UnitTests/RuntimeConfigHelperTests.cs +++ b/src/Service.Tests/UnitTests/RuntimeConfigHelperTests.cs @@ -73,6 +73,49 @@ public void CosmosDisablesRestEvenWhenRuntimeEnablesIt() Assert.IsFalse(config.IsRestEnabled); } + [TestMethod] + public void RuntimeProperties_EvaluateConfiguredSectionsAndHealthValues() + { + HashSet roles = new() { "reader" }; + RuntimeOptions runtime = new( + Rest: null, + GraphQL: null, + Mcp: null, + Host: new HostOptions(null, null), + Health: new RuntimeHealthCheckConfig(enabled: true, roles, cacheTtlSeconds: 17)); + RuntimeConfig config = CreateConfig(runtime); + + Assert.IsTrue(config.IsRestEnabled); + Assert.IsTrue(config.IsGraphQLEnabled); + Assert.IsTrue(config.IsMcpEnabled); + Assert.IsTrue(config.IsHealthEnabled); + Assert.IsTrue(config.AllowedRolesForHealth.SetEquals(roles)); + Assert.AreEqual(17, config.CacheTtlSecondsForHealthReport); + } + + [TestMethod] + public void EnableDwNto1JoinOpt_EvaluatesEachNestedConfigurationState() + { + RuntimeConfig missingGraphQL = CreateConfig(new RuntimeOptions(null, null, null, null)); + RuntimeConfig missingFlags = CreateConfig(new RuntimeOptions( + null, + new GraphQLRuntimeOptions { FeatureFlags = null! }, + null, + null)); + RuntimeConfig enabled = CreateConfig(new RuntimeOptions( + null, + new GraphQLRuntimeOptions + { + FeatureFlags = new FeatureFlags { EnableDwNto1JoinQueryOptimization = true } + }, + null, + null)); + + Assert.IsFalse(missingGraphQL.EnableDwNto1JoinOpt); + Assert.IsFalse(missingFlags.EnableDwNto1JoinOpt); + Assert.IsTrue(enabled.EnableDwNto1JoinOpt); + } + [TestMethod] public void DataSourceAndEntityMaps_SupportLookupUpdateAndPathOperations() { diff --git a/src/Service.Tests/UnitTests/RuntimeOptionsConverterCoverageTests.cs b/src/Service.Tests/UnitTests/RuntimeOptionsConverterCoverageTests.cs index 129851b879..45bb7190a3 100644 --- a/src/Service.Tests/UnitTests/RuntimeOptionsConverterCoverageTests.cs +++ b/src/Service.Tests/UnitTests/RuntimeOptionsConverterCoverageTests.cs @@ -1,6 +1,7 @@ // Copyright (c) Microsoft Corporation. // Licensed under the MIT License. +using System.Collections.Generic; using System.Text.Json; using Azure.DataApiBuilder.Config; using Azure.DataApiBuilder.Config.ObjectModel; @@ -74,6 +75,45 @@ public void RestRuntimeOptions_BooleanShorthandDeserializes(string json, bool ex Assert.AreEqual(expectedEnabled, value.Enabled); } + [DataTestMethod] + [DataRow(null, true, DisplayName = "Missing sections default to enabled")] + [DataRow(true, true, DisplayName = "Explicitly enabled sections are enabled")] + [DataRow(false, false, DisplayName = "Explicitly disabled sections are disabled")] + public void RuntimeOptions_EnablementProperties_DefaultUnlessExplicitlyDisabled(bool? enabled, bool expected) + { + RuntimeOptions options = new( + Rest: enabled.HasValue ? new RestRuntimeOptions(Enabled: enabled.Value) : null, + GraphQL: enabled.HasValue ? new GraphQLRuntimeOptions(Enabled: enabled.Value) : null, + Mcp: enabled.HasValue ? new McpRuntimeOptions(Enabled: enabled.Value) : null, + Host: null, + Health: enabled.HasValue ? new RuntimeHealthCheckConfig(enabled.Value) : null); + + Assert.AreEqual(expected, options.IsRestEnabled); + Assert.AreEqual(expected, options.IsGraphQLEnabled); + Assert.AreEqual(expected, options.IsMcpEnabled); + Assert.AreEqual(expected, options.IsHealthCheckEnabled); + } + + [DataTestMethod] + [DataRow("rest")] + [DataRow("graphql")] + [DataRow("mcp")] + [DataRow("health")] + public void RuntimeOptions_EnablementProperties_EvaluateEachSectionIndependently(string disabledSection) + { + RuntimeOptions options = new( + Rest: new RestRuntimeOptions(Enabled: disabledSection != "rest"), + GraphQL: new GraphQLRuntimeOptions(Enabled: disabledSection != "graphql"), + Mcp: new McpRuntimeOptions(Enabled: disabledSection != "mcp"), + Host: null, + Health: new RuntimeHealthCheckConfig(enabled: disabledSection != "health")); + + Assert.AreEqual(disabledSection != "rest", options.IsRestEnabled); + Assert.AreEqual(disabledSection != "graphql", options.IsGraphQLEnabled); + Assert.AreEqual(disabledSection != "mcp", options.IsMcpEnabled); + Assert.AreEqual(disabledSection != "health", options.IsHealthCheckEnabled); + } + [DataTestMethod] [DataRow("{\"multiple-mutations\":{\"unknown\":true}}")] [DataRow("{\"multiple-mutations\":42}")] @@ -133,5 +173,36 @@ public void FileSinkOptions_InvalidValuesThrow(string json) { Assert.ThrowsException(() => JsonSerializer.Deserialize(json, Options)); } + + [TestMethod] + public void RuntimeHealthOptions_WriteConfiguredAndDefaultForms() + { + RuntimeHealthCheckConfig configured = new( + enabled: true, + roles: new HashSet { "reader" }, + cacheTtlSeconds: 12, + maxQueryParallelism: 3); + + string configuredJson = JsonSerializer.Serialize(configured, Options); + string defaultJson = JsonSerializer.Serialize(new RuntimeHealthCheckConfig(), Options); + + StringAssert.Contains(configuredJson, "\"enabled\": true"); + StringAssert.Contains(configuredJson, "\"cache-ttl-seconds\": 12"); + StringAssert.Contains(configuredJson, "\"roles\""); + StringAssert.Contains(configuredJson, "\"max-query-parallelism\": 3"); + Assert.AreEqual("null", defaultJson); + } + + [TestMethod] + public void DatasourceHealthOptions_WriteEachUserProvidedTrigger() + { + string enabled = JsonSerializer.Serialize(new DatasourceHealthCheckConfig(enabled: true), Options); + string named = JsonSerializer.Serialize(new DatasourceHealthCheckConfig(enabled: null, name: "primary"), Options); + string threshold = JsonSerializer.Serialize(new DatasourceHealthCheckConfig(enabled: null, thresholdMs: 42), Options); + + StringAssert.Contains(enabled, "\"enabled\": true"); + StringAssert.Contains(named, "\"name\": \"primary\""); + StringAssert.Contains(threshold, "\"threshold-ms\": 42"); + } } } From b7c8b1427994e059efa1298aae8412e0367e0b8d Mon Sep 17 00:00:00 2001 From: Aaron Burtle Date: Sun, 30 Aug 2026 04:13:38 -0700 Subject: [PATCH 19/19] format --- .../Sql/GraphQLStoredProcedureBuilderHelpersTests.cs | 2 +- src/Service.Tests/UnitTests/GraphQLFilterParserUnitTests.cs | 1 - 2 files changed, 1 insertion(+), 2 deletions(-) diff --git a/src/Service.Tests/GraphQLBuilder/Sql/GraphQLStoredProcedureBuilderHelpersTests.cs b/src/Service.Tests/GraphQLBuilder/Sql/GraphQLStoredProcedureBuilderHelpersTests.cs index cdf1013512..946af28bbf 100644 --- a/src/Service.Tests/GraphQLBuilder/Sql/GraphQLStoredProcedureBuilderHelpersTests.cs +++ b/src/Service.Tests/GraphQLBuilder/Sql/GraphQLStoredProcedureBuilderHelpersTests.cs @@ -6,8 +6,8 @@ using System.Reflection; using System.Text.Json; using Azure.DataApiBuilder.Config.DatabasePrimitives; -using Azure.DataApiBuilder.Service.GraphQLBuilder; using Azure.DataApiBuilder.Service.Exceptions; +using Azure.DataApiBuilder.Service.GraphQLBuilder; using HotChocolate.Language; using Microsoft.VisualStudio.TestTools.UnitTesting; diff --git a/src/Service.Tests/UnitTests/GraphQLFilterParserUnitTests.cs b/src/Service.Tests/UnitTests/GraphQLFilterParserUnitTests.cs index b4384ed295..9e6acc7612 100644 --- a/src/Service.Tests/UnitTests/GraphQLFilterParserUnitTests.cs +++ b/src/Service.Tests/UnitTests/GraphQLFilterParserUnitTests.cs @@ -3,7 +3,6 @@ using System.Collections.Generic; using System.Reflection; -using Azure.DataApiBuilder.Config.DatabasePrimitives; using Azure.DataApiBuilder.Config.ObjectModel; using Azure.DataApiBuilder.Core.Configurations; using Azure.DataApiBuilder.Core.Models;