Skip to content
This repository was archived by the owner on Sep 3, 2024. It is now read-only.

Commit d7bc9a2

Browse files
force query strings to be case sensitive
In order to maximise cache-hits on a reverse proxy e.g. CloudFront, we want to normalize all urls. However, aspnet core inexplicably treats query string keys as case-insensitive - ?d=b is the same as ?D=b.
1 parent 2e8947f commit d7bc9a2

6 files changed

Lines changed: 199 additions & 21 deletions

File tree

Lines changed: 64 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,64 @@
1+
namespace SqlStreamStore.HAL
2+
{
3+
using System;
4+
using System.Collections;
5+
using System.Collections.Generic;
6+
using Microsoft.AspNetCore.Http;
7+
using Microsoft.AspNetCore.Http.Features;
8+
using Microsoft.Extensions.Primitives;
9+
10+
internal class CaseSensitiveQueryFeature : IQueryFeature
11+
{
12+
private readonly IFeatureCollection _features;
13+
14+
public CaseSensitiveQueryFeature(IFeatureCollection features)
15+
{
16+
if(features == null)
17+
throw new ArgumentNullException(nameof(features));
18+
_features = features;
19+
}
20+
21+
public IQueryCollection Query
22+
{
23+
get;
24+
set;
25+
}
26+
}
27+
28+
internal class CaseSensitiveQueryCollection : IQueryCollection
29+
{
30+
private readonly QueryString _queryString;
31+
32+
private Dictionary<string, StringValues> _state;
33+
34+
public CaseSensitiveQueryCollection(QueryString queryString)
35+
{
36+
_queryString = queryString;
37+
}
38+
39+
private Dictionary<string, StringValues> GetState()
40+
=> _state
41+
?? (_state = QueryStringHelper.ParseQueryString(_queryString));
42+
43+
public IEnumerator<KeyValuePair<string, StringValues>> GetEnumerator()
44+
=> GetState().GetEnumerator();
45+
46+
IEnumerator IEnumerable.GetEnumerator()
47+
=> GetEnumerator();
48+
49+
public bool ContainsKey(string key)
50+
=> GetState().ContainsKey(key);
51+
52+
public bool TryGetValue(string key, out StringValues value)
53+
=> GetState().TryGetValue(key, out value);
54+
55+
public int Count
56+
=> GetState().Count;
57+
58+
public ICollection<string> Keys
59+
=> GetState().Keys;
60+
61+
public StringValues this[string key]
62+
=> GetState()[key];
63+
}
64+
}
Lines changed: 23 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,23 @@
1+
namespace SqlStreamStore.HAL
2+
{
3+
using System;
4+
using Microsoft.AspNetCore.Http;
5+
using Microsoft.Extensions.Primitives;
6+
7+
internal static class QueryCollectionExtensions
8+
{
9+
public static bool TryGetValueCaseInsensitive(this IQueryCollection query, char key, out StringValues values)
10+
{
11+
if(query == null)
12+
{
13+
throw new ArgumentNullException(nameof(query));
14+
}
15+
16+
return char.IsUpper(key)
17+
? query.TryGetValue(key.ToString(), out values)
18+
|| query.TryGetValue(char.ToLower(key).ToString(), out values)
19+
: query.TryGetValue(key.ToString(), out values)
20+
|| query.TryGetValue(char.ToUpper(key).ToString(), out values);
21+
}
22+
}
23+
}
Lines changed: 51 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,51 @@
1+
namespace SqlStreamStore.HAL
2+
{
3+
using System;
4+
using System.Collections.Generic;
5+
using Microsoft.AspNetCore.Http;
6+
using Microsoft.Extensions.Primitives;
7+
8+
internal static class QueryStringHelper
9+
{
10+
private static readonly char[] s_delimiter = { '=' };
11+
private static readonly QueryString s_almostEmpty = new QueryString("?");
12+
13+
public static Dictionary<string, StringValues> ParseQueryString(QueryString queryString)
14+
{
15+
var state = new Dictionary<string, StringValues>();
16+
17+
if(queryString == QueryString.Empty || queryString == s_almostEmpty)
18+
{
19+
return state;
20+
}
21+
22+
var qs = queryString.Value;
23+
24+
if(qs[0] == '?')
25+
{
26+
qs = qs.Substring(1);
27+
}
28+
29+
foreach(var pair in qs.Split('&'))
30+
{
31+
var parts = pair.Split(s_delimiter, 2);
32+
var key = Uri.UnescapeDataString(parts[0].Replace('+', ' '));
33+
34+
state.TryGetValue(key, out var values);
35+
36+
if(parts.Length == 1)
37+
{
38+
state[key] = values;
39+
}
40+
else
41+
{
42+
var value = Uri.UnescapeDataString(parts[1].Replace('+', ' '));
43+
44+
state[key] = StringValues.Concat(value, values);
45+
}
46+
}
47+
48+
return state;
49+
}
50+
}
51+
}

src/SqlStreamStore.HAL/Resources/ReadAllStreamOperation.cs

Lines changed: 21 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -12,21 +12,32 @@ internal class ReadAllStreamOperation : IStreamStoreOperation<ReadAllPage>
1212

1313
public ReadAllStreamOperation(HttpRequest request)
1414
{
15-
EmbedPayload = request.Query.TryGetValue("e", out _);
15+
EmbedPayload = request.Query.TryGetValueCaseInsensitive('e', out _);
1616

17-
ReadDirection = request.Query["d"] == "f"
17+
ReadDirection = request.Query.TryGetValueCaseInsensitive('d', out var readDirection)
18+
&& readDirection == "f" || readDirection == "F"
1819
? Constants.ReadDirection.Forwards
1920
: Constants.ReadDirection.Backwards;
2021

21-
if(!long.TryParse(request.Query["p"], out _fromPositionInclusive))
22-
{
23-
_fromPositionInclusive = ReadDirection > 0 ? Position.Start : Position.End;
24-
}
22+
_fromPositionInclusive = request.Query.TryGetValueCaseInsensitive('p', out var position)
23+
? (long.TryParse(position, out _fromPositionInclusive)
24+
? (_fromPositionInclusive < Position.End
25+
? Position.End
26+
: _fromPositionInclusive)
27+
: (ReadDirection == Constants.ReadDirection.Forwards
28+
? Position.Start
29+
: Position.End))
30+
: (ReadDirection == Constants.ReadDirection.Forwards
31+
? Position.Start
32+
: Position.End);
2533

26-
if(!int.TryParse(request.Query["m"], out _maxCount))
27-
{
28-
_maxCount = Constants.MaxCount;
29-
}
34+
_maxCount = request.Query.TryGetValueCaseInsensitive('m', out var maxCount)
35+
? (int.TryParse(maxCount, out _maxCount)
36+
? (_maxCount <= 0
37+
? Constants.MaxCount
38+
: _maxCount)
39+
: Constants.MaxCount)
40+
: Constants.MaxCount;
3041

3142
Self = ReadDirection == Constants.ReadDirection.Forwards
3243
? LinkFormatter.FormatForwardLink(

src/SqlStreamStore.HAL/Resources/ReadStreamOperation.cs

Lines changed: 24 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -14,23 +14,36 @@ public ReadStreamOperation(HttpRequest request)
1414
{
1515
StreamId = request.Path.Value.Remove(0, 1);
1616

17-
EmbedPayload = request.Query.TryGetValue("e", out _);
17+
EmbedPayload = request.Query.TryGetValueCaseInsensitive('e', out _);
1818

19-
ReadDirection = request.Query["d"] == "f"
19+
ReadDirection = request.Query.TryGetValueCaseInsensitive('d', out var readDirection)
20+
&& readDirection == "f" || readDirection == "F"
2021
? Constants.ReadDirection.Forwards
2122
: Constants.ReadDirection.Backwards;
2223

23-
if(!int.TryParse(request.Query["p"], out _fromVersionInclusive))
24-
{
25-
_fromVersionInclusive = ReadDirection == Constants.ReadDirection.Forwards
24+
_fromVersionInclusive = request.Query.TryGetValueCaseInsensitive('p', out var position)
25+
? (int.TryParse(position, out _fromVersionInclusive)
26+
? (ReadDirection == Constants.ReadDirection.Forwards
27+
? (_fromVersionInclusive < StreamVersion.Start
28+
? StreamVersion.Start
29+
: _fromVersionInclusive)
30+
: (_fromVersionInclusive < StreamVersion.End
31+
? StreamVersion.End
32+
: _fromVersionInclusive))
33+
: (ReadDirection == Constants.ReadDirection.Forwards
34+
? StreamVersion.Start
35+
: StreamVersion.End))
36+
: (ReadDirection == Constants.ReadDirection.Forwards
2637
? StreamVersion.Start
27-
: StreamVersion.End;
28-
}
38+
: StreamVersion.End);
2939

30-
if(!int.TryParse(request.Query["m"], out _maxCount))
31-
{
32-
_maxCount = Constants.MaxCount;
33-
}
40+
_maxCount = request.Query.TryGetValueCaseInsensitive('m', out var maxCount)
41+
? (int.TryParse(maxCount, out _maxCount)
42+
? (_maxCount <= 0
43+
? Constants.MaxCount
44+
: _maxCount)
45+
: Constants.MaxCount)
46+
: Constants.MaxCount;
3447

3548
Self = ReadDirection == Constants.ReadDirection.Forwards
3649
? LinkFormatter.FormatForwardLink(StreamId, MaxCount, FromVersionInclusive, EmbedPayload)

src/SqlStreamStore.HAL/SqlStreamStoreHalMiddleware.cs

Lines changed: 16 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -14,6 +14,21 @@
1414

1515
public static class SqlStreamStoreHalMiddleware
1616
{
17+
private static MidFunc CaseSensitiveQueryStrings => (context, next) =>
18+
{
19+
if(context.Request.QueryString != QueryString.Empty)
20+
{
21+
var queryString = context.Request.QueryString;
22+
context.Request.Query = new CaseSensitiveQueryCollection(queryString);
23+
// Setting context.Request.Query mutates context.Request.QueryString.
24+
// This has the unfortunate side effect of turning ?a=1&b into ?a=1&b=.
25+
// so, replace with original context.Request.QueryString and call it a day.
26+
context.Request.QueryString = queryString;
27+
}
28+
29+
return next();
30+
};
31+
1732
private static MidFunc MethodsNotAllowed(params string[] methods) => (context, next) =>
1833
{
1934
if(!methods.Contains(context.Request.Method))
@@ -68,6 +83,7 @@ public static IApplicationBuilder UseSqlStreamStoreHal(
6883

6984
return builder
7085
.Use(ExceptionHandlingMiddleware.HandleExceptions)
86+
.Use(CaseSensitiveQueryStrings)
7187
.Use(AcceptHalJson)
7288
.Use(Index)
7389
.Map("/stream", UseAllStream(streamStore))

0 commit comments

Comments
 (0)