diff --git a/core/src/main/java/com/linecorp/armeria/internal/server/annotation/AnnotatedValueResolver.java b/core/src/main/java/com/linecorp/armeria/internal/server/annotation/AnnotatedValueResolver.java index 0e93c707a8e..77898ba1e02 100644 --- a/core/src/main/java/com/linecorp/armeria/internal/server/annotation/AnnotatedValueResolver.java +++ b/core/src/main/java/com/linecorp/armeria/internal/server/annotation/AnnotatedValueResolver.java @@ -64,6 +64,7 @@ import com.google.common.base.MoreObjects; import com.google.common.base.Splitter; import com.google.common.collect.ImmutableList; +import com.google.common.collect.ImmutableSet; import com.google.common.collect.Iterables; import com.google.common.primitives.Primitives; @@ -602,18 +603,60 @@ private static AnnotatedValueResolver ofQueryParamMap(String name, AnnotatedElement annotatedElement, AnnotatedElement typeElement, Class type, DescriptionInfo description) { + final Type valueType = ((ParameterizedType) ((Parameter) typeElement).getParameterizedType()) + .getActualTypeArguments()[1]; + final Class rawValueType = ClassUtil.typeToClass(valueType); + assert rawValueType != null; + + if (valueType instanceof ParameterizedType && !(List.class.isAssignableFrom(rawValueType) || + Set.class.isAssignableFrom(rawValueType))) { + throw new IllegalArgumentException( + "Invalid parameterized map value type: " + rawValueType + + " (expected List or Set)"); + } + + final BiFunction biFunction; + + if (Set.class.isAssignableFrom(rawValueType)) { + biFunction = (resolver, ctx) -> ctx.queryParams().stream() + .collect(toImmutableMap( + Entry::getKey, + e -> ImmutableSet.of(e.getValue()), + (existing, replacement) -> + ImmutableSet.builder() + .addAll(existing) + .addAll(replacement) + .build() + )); + } else if (List.class.isAssignableFrom(rawValueType) || + Collection.class.isAssignableFrom(rawValueType) || + Iterable.class.isAssignableFrom(rawValueType) + ) { + biFunction = (resolver, ctx) -> ctx.queryParams().stream() + .collect(toImmutableMap( + Entry::getKey, + e -> ImmutableList.of(e.getValue()), + (existing, replacement) -> + ImmutableList.builder() + .addAll(existing) + .addAll(replacement) + .build() + )); + } else { + biFunction = (resolver, ctx) -> ctx.queryParams().stream() + .collect(toImmutableMap( + Entry::getKey, + Entry::getValue, + (existing, replacement) -> replacement + )); + } return new Builder(annotatedElement, type, name) .annotationType(Param.class) .typeElement(typeElement) .description(description) .aggregation(AggregationStrategy.FOR_FORM_DATA) - .resolver((resolver, ctx) -> ctx.queryParams().stream() - .collect(toImmutableMap( - Entry::getKey, - Entry::getValue, - (existing, replacement) -> replacement - ))) + .resolver(biFunction) .build(); } diff --git a/core/src/test/java/com/linecorp/armeria/internal/server/annotation/AnnotatedServiceTest.java b/core/src/test/java/com/linecorp/armeria/internal/server/annotation/AnnotatedServiceTest.java index f32ff8afd5f..7f4de00d2f3 100644 --- a/core/src/test/java/com/linecorp/armeria/internal/server/annotation/AnnotatedServiceTest.java +++ b/core/src/test/java/com/linecorp/armeria/internal/server/annotation/AnnotatedServiceTest.java @@ -639,6 +639,22 @@ public String map(RequestContext ctx, @Param Map map) { .map(entry -> entry.getKey() + '=' + entry.getValue()) .collect(Collectors.joining(", ")); } + + @Get("/param/listMap") + public String listMap(RequestContext ctx, @Param Map> map) { + validateContext(ctx); + return map.isEmpty() ? "empty" : map.entrySet().stream() + .map(entry -> entry.getKey() + '=' + entry.getValue()) + .collect(Collectors.joining(", ")); + } + + @Get("/param/setMap") + public String setMap(RequestContext ctx, @Param Map> map) { + validateContext(ctx); + return map.isEmpty() ? "empty" : map.entrySet().stream() + .map(entry -> entry.getKey() + '=' + entry.getValue()) + .collect(Collectors.joining(", ")); + } } @ResponseConverter(UnformattedStringConverterFunction.class) @@ -1080,6 +1096,16 @@ void testParam() throws Exception { testBody(hc, get("/7/param/map?key1=value1&key2=value2"), "key1=value1, key2=value2"); testBody(hc, get("/7/param/map"), "empty"); + + // Case all query parameters test multi value map of List + testBody(hc, get("/7/param/listMap?key1=value1&key1=value2&key2=value1&key2=value2"), + "key1=[value1, value2], key2=[value1, value2]"); + testBody(hc, get("/7/param/listMap"), "empty"); + + // Case all query parameters test multi value map of Set + testBody(hc, get("/7/param/setMap?key1=value1&key1=value1&key2=value2&key2=value2"), + "key1=[value1], key2=[value2]"); + testBody(hc, get("/7/param/setMap"), "empty"); } } diff --git a/core/src/test/java/com/linecorp/armeria/internal/server/annotation/AnnotatedValueResolverTest.java b/core/src/test/java/com/linecorp/armeria/internal/server/annotation/AnnotatedValueResolverTest.java index d3acd4c0422..d59383060d5 100644 --- a/core/src/test/java/com/linecorp/armeria/internal/server/annotation/AnnotatedValueResolverTest.java +++ b/core/src/test/java/com/linecorp/armeria/internal/server/annotation/AnnotatedValueResolverTest.java @@ -108,11 +108,14 @@ class AnnotatedValueResolverTest { "value3", "value2"); + static final Set queryParamMaps = ImmutableSet.of("queryParamMap", + "queryParamListMap", + "queryParamSetMap"); + static final ResolverContext resolverContext; static final ServiceRequestContext context; static final HttpRequest request; static final RequestHeaders originalHeaders; - static final String QUERY_PARAM_MAP = "queryParamMap"; static Map> successExpectAttrKeys; static Map> failExpectAttrKeys; @@ -182,6 +185,15 @@ void ofMethods() { // Ignore this exception because MixedBean class has not annotated method. } }); + + // Validate that invalid multi-value map parameter types trigger an exception + getAllMethods(InvalidMultiValueMapService.class, + method -> !Modifier.isPrivate(method.getModifiers())).forEach( + method -> assertThatThrownBy(() -> AnnotatedValueResolver.ofServiceMethod( + method, pathParams, objectResolvers, false, noopDependencyInjector, null)) + .isInstanceOf(IllegalArgumentException.class) + .hasMessageContaining("Invalid parameterized map value type") + ); } @Test @@ -364,7 +376,7 @@ private static void testResolver(AnnotatedValueResolver resolver) { } } } else { - if (QUERY_PARAM_MAP.equals(resolver.httpElementName())) { + if (queryParamMaps.contains(resolver.httpElementName())) { assertThat(resolver.defaultValue()).isNull(); } else { assertThat(resolver.defaultValue()).isNotNull(); @@ -376,7 +388,7 @@ private static void testResolver(AnnotatedValueResolver resolver) { .isEqualTo(resolver.elementType()); } else if (resolver.shouldWrapValueAsOptional()) { assertThat(value).isEqualTo(Optional.of(resolver.defaultValue())); - } else if (QUERY_PARAM_MAP.equals(resolver.httpElementName())) { + } else if (queryParamMaps.contains(resolver.httpElementName())) { assertThat(value).isNotNull(); assertThat(value).isInstanceOf(Map.class); assertThat((Map) value).size() @@ -459,6 +471,8 @@ void method1(@Param String var1, @Param @Default List emptyParam3, @Param @Default List emptyParam4, @Param Map queryParamMap, + @Param Map> queryParamListMap, + @Param Map> queryParamSetMap, @Header List header1, @Header("header1") Optional> optionalHeader1, @Header String header2, @@ -519,7 +533,7 @@ void attributeTest( Queue successQueueToQueue, @Attribute("failCastListToSet") Set failCastListToSet - ) { } + ) {} void time(@Param @Default("PT20.345S") Duration duration, @Param @Default("2007-12-03T10:15:30.00Z") Instant instant, @@ -534,6 +548,10 @@ void time(@Param @Default("PT20.345S") Duration duration, @Param @Default("+01:00:00") ZoneOffset zoneOffset) {} } + static class InvalidMultiValueMapService { + void invalidParamWithMapOfMap(@Param Map> param) {} + } + private static Map> injectFailCaseOfAttrKeyToServiceContextForAttributeTest() { final ServiceRequestContext ctx = resolverContext.context(); final Map> expectFailAttrs = new HashMap<>();