33 Commits
Author SHA1 Message Date
GHXX 29eecc7887 fix incorrect check which might break when registering an endpoint class that contains no endpoints 2024-07-25 04:00:06 +02:00
GHXX a4ae359df0 cleanup some old auth stuff 2024-07-25 03:41:53 +02:00
GHXX 176c5e7197 fix required GET args being present not triggering a 400 when no GET args were passed at all 2024-07-21 06:45:13 +02:00
GHXX d7a934e25c cleanup 2024-07-20 08:02:06 +02:00
GHXX c75d29a1ba Switch over to a check-based system with multi attribute support 2024-07-19 03:31:04 +02:00
GHXX fa79134d02 Move file 2024-07-19 03:28:56 +02:00
GHXX cdab5151be fix parameter url decoding 2024-07-13 06:32:59 +02:00
GHXX b645d4d654 make dispose public 2024-07-07 19:11:55 +02:00
GHXX f20ba933dc threadsafety improvements 2024-05-25 23:54:57 +02:00
GHXX c91714a6af add svg staticserve content type handling 2024-02-02 05:07:20 +01:00
GHXX a9436bfda8 Revert "add ipv6 listener prefix"
This reverts commit f14294387e.
2024-01-31 05:14:03 +01:00
GHXX f14294387e add ipv6 listener prefix 2024-01-31 05:00:27 +01:00
GHXX 0bfe34ab6b improve static serve jail 2024-01-26 19:49:04 +01:00
GHXX 6cc849bf01 add ParsedParameters to RequestContext, cleanup error handling 2024-01-16 23:52:28 +01:00
GHXX 8cdff9268a fix some serving issues 2024-01-16 17:30:28 +01:00
GHXX 94b23cadc5 Add missing Close() statement 2024-01-15 21:09:17 +01:00
GHXX 0814bc6b2d Add static serving 2024-01-15 19:56:14 +01:00
GHXX 8545ed80e9 wip static serving 2024-01-15 04:28:17 +01:00
GHXX 4fad2d648e Replace argon2, add threadsafe SHA256 method, rename some variables 2024-01-14 23:15:12 +01:00
GHXX ea74cb899c move type to different folder 2024-01-14 21:50:49 +01:00
GHXX 92e472d526 Add tests for parameter optionality 2024-01-14 03:07:34 +01:00
GHXX 14ab546d4d Fix server producing a timeout on setting status code without returning a body 2024-01-14 02:58:01 +01:00
GHXX e1e1596e54 add query arg test 2024-01-14 02:31:45 +01:00
GHXX 6d74d659f6 fix query args 2024-01-14 02:31:31 +01:00
GHXX 0b4975e74b Improve error handling 2024-01-14 02:31:06 +01:00
GHXX d6190b024d Add missing string parameter converter 2024-01-14 02:30:20 +01:00
GHXX aa06679742 add tests for normalization and serving multiple pages 2024-01-14 00:56:56 +01:00
GHXX 020075ad54 normalize urls during registering and requesting so that they all start with a single slash 2024-01-14 00:55:43 +01:00
GHXX f0c9754fb2 fix spelling mistake 2024-01-14 00:55:11 +01:00
GHXX 7ad2b5185b add log output checking in tests 2024-01-13 19:07:00 +01:00
GHXX 08003d1fc3 Simplify tests 2024-01-13 03:14:22 +01:00
GHXX a2a70e8339 Catch exception on shutdown 2024-01-13 03:14:10 +01:00
GHXX 09fa3b8734 Huge refactor; Passing tests 2024-01-13 01:31:30 +01:00
13 changed files with 592 additions and 354 deletions
+2 -9
View File
@@ -1,22 +1,15 @@
using SimpleHttpServer.Internal; using SimpleHttpServer.Types;
namespace SimpleHttpServer; namespace SimpleHttpServer;
[AttributeUsage(AttributeTargets.Method, AllowMultiple = false)] [AttributeUsage(AttributeTargets.Method, AllowMultiple = false)]
public class HttpEndpointAttribute<T> : Attribute where T : IAuthorizer { public class HttpEndpointAttribute : Attribute {
public HttpRequestType RequestMethod { get; private set; } public HttpRequestType RequestMethod { get; private set; }
public string[] Locations { get; private set; } public string[] Locations { get; private set; }
public Type Authorizer { get; private set; }
public HttpEndpointAttribute(HttpRequestType requestMethod, params string[] locations) { public HttpEndpointAttribute(HttpRequestType requestMethod, params string[] locations) {
RequestMethod = requestMethod; RequestMethod = requestMethod;
Locations = locations; Locations = locations;
Authorizer = typeof(T);
} }
} }
[AttributeUsage(AttributeTargets.Method)]
public class HttpEndpointAttribute : HttpEndpointAttribute<DefaultAuthorizer> {
public HttpEndpointAttribute(HttpRequestType type, params string[] locations) : base(type, locations) { }
}
+153 -33
View File
@@ -4,6 +4,8 @@ using SimpleHttpServer.Types.ParameterConverters;
using System.Net; using System.Net;
using System.Numerics; using System.Numerics;
using System.Reflection; using System.Reflection;
using System.Text;
using static SimpleHttpServer.Types.EndpointInvocationInfo;
namespace SimpleHttpServer; namespace SimpleHttpServer;
@@ -13,7 +15,8 @@ public sealed class HttpServer {
private readonly HttpListener listener; private readonly HttpListener listener;
private Task? listenerTask; private Task? listenerTask;
private readonly Logger logger; private readonly Logger mainLogger;
private readonly Logger requestLogger;
private readonly SimpleHttpServerConfiguration conf; private readonly SimpleHttpServerConfiguration conf;
private bool shutdown = false; private bool shutdown = false;
@@ -22,19 +25,20 @@ public sealed class HttpServer {
conf = configuration; conf = configuration;
listener = new HttpListener(); listener = new HttpListener();
listener.Prefixes.Add($"http://localhost:{port}/"); listener.Prefixes.Add($"http://localhost:{port}/");
logger = new(LogOutputTopic.Main, conf); mainLogger = new(LogOutputTopic.Main, conf);
requestLogger = new(LogOutputTopic.Request, conf);
} }
public void Start() { public void Start() {
logger.Information($"Starting on port {Port}..."); mainLogger.Information($"Starting on port {Port}...");
Assert(listenerTask == null, "Server was already started!"); Assert(listenerTask == null, "Server was already started!");
listener.Start(); listener.Start();
listenerTask = Task.Run(GetContextLoopAsync); listenerTask = Task.Run(GetContextLoopAsync);
logger.Information($"Ready to handle requests!"); mainLogger.Information($"Ready to handle requests!");
} }
public async Task StopAsync(CancellationToken ctok) { public async Task StopAsync(CancellationToken ctok) {
logger.Information("Stopping server..."); mainLogger.Information("Stopping server...");
Assert(listenerTask != null, "Server was not started!"); Assert(listenerTask != null, "Server was not started!");
shutdown = true; shutdown = true;
listener.Stop(); listener.Stop();
@@ -46,8 +50,9 @@ public sealed class HttpServer {
try { try {
var ctx = await listener.GetContextAsync(); var ctx = await listener.GetContextAsync();
_ = ProcessRequestAsync(ctx); _ = ProcessRequestAsync(ctx);
} catch (HttpListenerException ex) when (ex.ErrorCode == 995) { //The I/O operation has been aborted because of either a thread exit or an application request
} catch (Exception ex) { } catch (Exception ex) {
logger.Fatal($"Caught otherwise uncaught exception in GetContextLoop:\n{ex}"); mainLogger.Fatal($"Caught otherwise uncaught exception in GetContextLoop:\n{ex}");
} }
} }
} }
@@ -56,6 +61,7 @@ public sealed class HttpServer {
void RegisterConverter<T>() where T : IParsable<T> { void RegisterConverter<T>() where T : IParsable<T> {
stringToTypeParameterConverters.Add(typeof(T), new ParsableParameterConverter<T>()); stringToTypeParameterConverters.Add(typeof(T), new ParsableParameterConverter<T>());
} }
stringToTypeParameterConverters.Add(typeof(string), new StringParameterConverter());
stringToTypeParameterConverters.Add(typeof(bool), new BoolParsableParameterConverter()); stringToTypeParameterConverters.Add(typeof(bool), new BoolParsableParameterConverter());
RegisterConverter<char>(); RegisterConverter<char>();
@@ -81,13 +87,13 @@ public sealed class HttpServer {
private readonly Dictionary<(string path, string rType), EndpointInvocationInfo> simpleEndpointMethodInfos = new(); private readonly Dictionary<(string path, string rType), EndpointInvocationInfo> simpleEndpointMethodInfos = new();
private static readonly Type[] expectedEndpointParameterTypes = new[] { typeof(RequestContext) }; private static readonly Type[] expectedEndpointParameterTypes = new[] { typeof(RequestContext) };
public void RegisterEndpointsFromType<T>() { public void RegisterEndpointsFromType<T>() {
if (simpleEndpointMethodInfos.Count == 0) if (stringToTypeParameterConverters.Count == 0)
RegisterDefaultConverters(); RegisterDefaultConverters();
var t = typeof(T); var t = typeof(T);
foreach (var (mi, attrib) in t.GetMethods() foreach (var (mi, attrib) in t.GetMethods()
.ToDictionary(x => x, x => x.GetCustomAttributes(typeof(HttpEndpointAttribute<>))) .ToDictionary(x => x, x => x.GetCustomAttributes<HttpEndpointAttribute>())
.Where(x => x.Value.Any()).ToDictionary(x => x.Key, x => (HttpEndpointAttribute) x.Value.Single())) { .Where(x => x.Value.Any()).ToDictionary(x => x.Key, x => x.Value.Single())) {
string GetFancyMethodName() => mi.DeclaringType!.FullName + "#" + mi.Name; string GetFancyMethodName() => mi.DeclaringType!.FullName + "#" + mi.Name;
@@ -104,43 +110,104 @@ public sealed class HttpServer {
Assert(mi.ReturnType == typeof(Task), $"Return type of {GetFancyMethodName()} is not {typeof(Task)}!"); Assert(mi.ReturnType == typeof(Task), $"Return type of {GetFancyMethodName()} is not {typeof(Task)}!");
var qparams = new List<(string, (Type type, bool isOptional))>(); var qparams = new List<QueryParameterInfo>();
for (int i = expectedEndpointParameterTypes.Length; i < methodParams.Length; i++) { for (int i = expectedEndpointParameterTypes.Length; i < methodParams.Length; i++) {
var par = methodParams[i]; var par = methodParams[i];
var attr = par.GetCustomAttribute<ParameterAttribute>(false); var attr = par.GetCustomAttribute<ParameterAttribute>(false);
qparams.Add((attr?.Name ?? par.Name ?? throw new ArgumentException($"C# variable name of parameter at index {i} of method {GetFancyMethodName()} is null!"), qparams.Add(new(
(par.GetType(), attr?.IsOptional ?? false))); attr?.Name ?? par.Name ?? throw new ArgumentException($"C# variable name of parameter at index {i} of method {GetFancyMethodName()} is null!"),
par.ParameterType,
attr?.IsOptional ?? false)
);
if (!stringToTypeParameterConverters.ContainsKey(par.ParameterType)) { if (!stringToTypeParameterConverters.ContainsKey(par.ParameterType)) {
throw new MissingParameterConverterException($"Parameter converter for type {par.ParameterType} has not been registered (yet)!"); throw new MissingParameterConverterException($"Parameter converter for type {par.ParameterType} has not been registered (yet)!");
} }
} }
// stores the check attributes that are defined on the method and on the containing class
var requiredChecks = mi.GetCustomAttributes<BaseEndpointCheckAttribute>(true).Concat(mi.DeclaringType?.GetCustomAttributes<BaseEndpointCheckAttribute>(true) ?? Enumerable.Empty<Attribute>())
.Where(a => a.GetType().IsAssignableTo(typeof(BaseEndpointCheckAttribute))).Cast<BaseEndpointCheckAttribute>().ToArray();
foreach (var location in attrib.Locations) { foreach (var location in attrib.Locations) {
int idx = location.IndexOf('{'); var normLocation = NormalizeUrlPath(location);
int idx = normLocation.IndexOf('{');
if (idx >= 0) { if (idx >= 0) {
// this path contains path parameters // this path contains path parameters
throw new NotImplementedException("Path parameters are not yet implemented!"); throw new NotImplementedException("Path parameters are not yet implemented!");
} }
var reqMethod = Enum.GetName(attrib.RequestMethod) ?? throw new ArgumentException("Request method was undefined"); var reqMethod = Enum.GetName(attrib.RequestMethod) ?? throw new ArgumentException("Request method was undefined");
simpleEndpointMethodInfos.Add((location, reqMethod), new EndpointInvocationInfo(mi, qparams)); mainLogger.Information($"Registered endpoint: '{reqMethod} {normLocation}'");
simpleEndpointMethodInfos.Add((normLocation, reqMethod), new EndpointInvocationInfo(mi, qparams, requiredChecks));
} }
} }
} }
/// <summary>
/// Serves all files located in <paramref name="filesystemDirectory"/> on a website path that is relative to <paramref name="requestPath"/>,
/// while restricting requests to inside the local filesystem directory. Static serving has a lower priority than registering an endpoint.
/// </summary>
/// <param name="requestPath"></param>
/// <param name="filesystemDirectory"></param>
public void RegisterStaticServePath(string requestPath, string filesystemDirectory) {
var absPath = Path.GetFullPath(filesystemDirectory);
string npath = NormalizeUrlPath(requestPath);
mainLogger.Information($"Registered static serve path: '{npath}' --> '{absPath}'");
staticServePaths.Add(npath, absPath);
}
private readonly Dictionary<string, string> staticServePaths = new();
private readonly Dictionary<Type, IParameterConverter> stringToTypeParameterConverters = new(); private readonly Dictionary<Type, IParameterConverter> stringToTypeParameterConverters = new();
private static string NormalizeUrlPath(string url) {
var fwdSlashUrl = url.Replace('\\', '/');
var segments = fwdSlashUrl.Trim('/').Split('/', StringSplitOptions.RemoveEmptyEntries).ToList();
List<string> simplifiedSegmentsReversed = new List<string>();
int doubleDotsEncountered = 0;
for (int i = segments.Count - 1; i >= 0; i--) {
var segment = segments[i];
if (segment == ".") {
continue; // remove single dot segments
}
if (segment == "..") {
doubleDotsEncountered++; // if we encounter a doubledot, keep track of that and dont add it to the output yet
continue;
}
// otherwise only keep the segment if doubleDotsEncountered > 0
if (doubleDotsEncountered > 0) {
doubleDotsEncountered--;
continue;
}
simplifiedSegmentsReversed.Add(segment);
}
var rv = new StringBuilder();
for (int i = 0; i < doubleDotsEncountered; i++) {
rv.Append("../");
}
rv.AppendJoin('/', simplifiedSegmentsReversed.Reverse<string>());
return '/' + (rv.ToString().TrimEnd('/') + (fwdSlashUrl.EndsWith('/') ? "/" : "")).TrimStart('/');
}
private async Task ProcessRequestAsync(HttpListenerContext ctx) { private async Task ProcessRequestAsync(HttpListenerContext ctx) {
using RequestContext rc = new RequestContext(ctx);
// TODO add path escape countermeasure-unittests
var splitted = (ctx.Request.RawUrl ?? "").Split('?', 2, StringSplitOptions.None);
var reqPath = NormalizeUrlPath(WebUtility.UrlDecode(splitted.First()));
string requestMethod = ctx.Request.HttpMethod.ToUpperInvariant();
bool wasStaticlyServed = false;
void LogRequest() {
requestLogger.Information($"{rc.ListenerContext.Response.StatusCode} {(wasStaticlyServed ? "static" : "endpnt")} {requestMethod} {ctx.Request.Url}");
}
try { try {
var decUri = WebUtility.UrlDecode(ctx.Request.RawUrl)!; // TODO add path escape countermeasures+unittests
var splitted = decUri.Split('?', 2, StringSplitOptions.None);
var path = WebUtility.UrlDecode(splitted.First());
if (simpleEndpointMethodInfos.TryGetValue((reqPath, requestMethod), out var endpointInvocationInfo)) {
using var rc = new RequestContext(ctx);
if (simpleEndpointMethodInfos.TryGetValue((decUri, ctx.Request.HttpMethod.ToUpperInvariant()), out var endpointInvocationInfo)) {
var mi = endpointInvocationInfo.methodInfo; var mi = endpointInvocationInfo.methodInfo;
var qparams = endpointInvocationInfo.queryParameters; var qparams = endpointInvocationInfo.queryParameters;
var args = splitted.Length == 2 ? splitted[1] : null; var args = splitted.Length == 2 ? splitted[1] : null;
@@ -148,63 +215,116 @@ public sealed class HttpServer {
var parsedQParams = new Dictionary<string, string>(); var parsedQParams = new Dictionary<string, string>();
var convertedQParamValues = new object[qparams.Count + 1]; var convertedQParamValues = new object[qparams.Count + 1];
// TODO add authcheck here // run the checks to see if the client is allowed to make this request
if (!endpointInvocationInfo.CheckAll(rc.ListenerContext.Request)) { // if any check failed return Forbidden
await HandleDefaultErrorPageAsync(rc, HttpStatusCode.Forbidden, "Client is not allowed to access this resource");
return;
}
if (args != null) { if (args != null) {
var queryStringArgs = args.Split('&', StringSplitOptions.None); var queryStringArgs = args.Split('&', StringSplitOptions.None);
foreach (var queryKV in queryStringArgs) { foreach (var queryKV in queryStringArgs) {
var queryKVSplitted = queryKV.Split('='); var queryKVSplitted = queryKV.Split('=');
if (queryKVSplitted.Length != 2) { if (queryKVSplitted.Length != 2) {
rc.SetStatusCodeAndDispose(HttpStatusCode.BadRequest, "Malformed request URL parameters"); await HandleDefaultErrorPageAsync(rc, HttpStatusCode.BadRequest, "Malformed request URL parameters");
return; return;
} }
if (!parsedQParams.TryAdd(WebUtility.UrlDecode(queryKVSplitted[0]), WebUtility.UrlDecode(queryKVSplitted[1]))) { if (!parsedQParams.TryAdd(WebUtility.UrlDecode(queryKVSplitted[0]), WebUtility.UrlDecode(queryKVSplitted[1]))) {
rc.SetStatusCodeAndDispose(HttpStatusCode.BadRequest, "Duplicate request URL parameters"); await HandleDefaultErrorPageAsync(rc, HttpStatusCode.BadRequest, "Duplicate request URL parameters");
return; return;
} }
} }
for (int i = 0; i < qparams.Count;) { for (int i = 0; i < qparams.Count;) {
var (qparamName, qparamInfo) = qparams[i]; var qparam = qparams[i];
i++; i++;
if (parsedQParams.TryGetValue(qparamName, out var qparamValue)) { if (parsedQParams.TryGetValue(qparam.Name, out var qparamValue)) {
if (stringToTypeParameterConverters[qparamInfo.type].TryConvertFromString(qparamValue, out object objRes)) { if (stringToTypeParameterConverters[qparam.Type].TryConvertFromString(qparamValue, out object objRes)) {
convertedQParamValues[i] = objRes; convertedQParamValues[i] = objRes;
} else { } else {
rc.SetStatusCodeAndDispose(HttpStatusCode.BadRequest); await HandleDefaultErrorPageAsync(rc, HttpStatusCode.BadRequest);
return; return;
} }
} else { } else {
if (qparamInfo.isOptional) { if (qparam.IsOptional) {
convertedQParamValues[i] = null!; convertedQParamValues[i] = null!;
} else { } else {
rc.SetStatusCodeAndDispose(HttpStatusCode.BadRequest, $"Missing required query parameter {qparamName}"); await HandleDefaultErrorPageAsync(rc, HttpStatusCode.BadRequest, $"Missing required query parameter {qparam.Name}");
return; return;
} }
} }
} }
} else {
var requiredParams = qparams.Where(x => !x.IsOptional).Select(x => $"'{x.Name}'").ToList();
if (requiredParams.Any()) {
await HandleDefaultErrorPageAsync(rc, HttpStatusCode.BadRequest, $"Missing required query parameter(s): {string.Join(",", requiredParams)}");
return;
}
} }
convertedQParamValues[0] = rc; convertedQParamValues[0] = rc;
rc.ParsedParameters = parsedQParams.AsReadOnly();
await (Task) (mi.Invoke(null, convertedQParamValues) ?? throw new NullReferenceException("Website func returned null unexpectedly")); await (Task) (mi.Invoke(null, convertedQParamValues) ?? throw new NullReferenceException("Website func returned null unexpectedly"));
} else { } else {
if (requestMethod == "GET")
foreach (var (k, v) in staticServePaths) {
if (reqPath.StartsWith(k)) { // do a static serve
wasStaticlyServed = true;
var relativeStaticReqPath = reqPath[k.Length..];
var staticResponsePath = Path.GetFullPath(Path.Join(v, relativeStaticReqPath.TrimStart('/')));
if (Path.GetRelativePath(v, staticResponsePath).Contains("..")) {
requestLogger.Warning($"Blocked GET request to {reqPath} as somehow the target file does not lie inside the static serve folder? Are you using symlinks?");
await HandleDefaultErrorPageAsync(rc, HttpStatusCode.NotFound);
return;
}
if (File.Exists(staticResponsePath)) {
rc.SetStatusCode(HttpStatusCode.OK);
if (staticResponsePath.EndsWith(".svg")) {
rc.ListenerContext.Response.AddHeader("Content-Type", "image/svg+xml");
}
using var f = File.OpenRead(staticResponsePath);
await f.CopyToAsync(rc.ListenerContext.Response.OutputStream);
} else {
await HandleDefaultErrorPageAsync(rc, HttpStatusCode.NotFound);
}
return;
}
}
// invoke 404 // invoke 404
await HandleDefaultErrorPageAsync(rc, 404); await HandleDefaultErrorPageAsync(rc, 404);
} }
} catch (Exception ex) { } catch (Exception ex) {
logger.Fatal($"Caught otherwise uncaught exception while ProcessingRequest:\n{ex}"); await HandleDefaultErrorPageAsync(rc, 500);
mainLogger.Fatal($"Caught otherwise uncaught exception while ProcessingRequest:\n{ex}");
} finally {
try { await rc.RespWriter.FlushAsync(); } catch (ObjectDisposedException) { }
rc.ListenerContext.Response.Close();
LogRequest();
} }
} }
private static async Task HandleDefaultErrorPageAsync(RequestContext ctx, HttpStatusCode errorCode, string? statusDescription = null) => await HandleDefaultErrorPageAsync(ctx, (int) errorCode, statusDescription);
private static async Task HandleDefaultErrorPageAsync(RequestContext ctx, int errorCode) { private static async Task HandleDefaultErrorPageAsync(RequestContext ctx, int errorCode, string? statusDescription = null) {
ctx.SetStatusCode(errorCode);
string desc = statusDescription != null ? $"\r\n{statusDescription}" : "";
await ctx.WriteLineToRespAsync($""" await ctx.WriteLineToRespAsync($"""
<body> <body>
<h1>Oh no, and error occurred!</h1> <h1>Oh no, an error occurred!</h1>
<p>Code: {errorCode}</p> <p>Code: {errorCode}</p>{desc}
</body> </body>
"""); """);
try {
if (statusDescription == null) {
await ctx.SetStatusCodeAndDisposeAsync(errorCode);
} else {
await ctx.SetStatusCodeAndDisposeAsync(errorCode, statusDescription);
}
} catch (ObjectDisposedException) { }
} }
} }
-7
View File
@@ -1,7 +0,0 @@
using System.Net;
namespace SimpleHttpServer;
public interface IAuthorizer {
public abstract (bool auth, object? data) IsAuthenticated(HttpListenerContext contect);
}
@@ -1,7 +0,0 @@
using System.Net;
namespace SimpleHttpServer.Internal;
public sealed class DefaultAuthorizer : IAuthorizer {
public (bool auth, object? data) IsAuthenticated(HttpListenerContext contect) => (true, null);
}
@@ -1,73 +1,73 @@
using Newtonsoft.Json; //using Newtonsoft.Json;
using System.Collections; //using System.Collections;
using System.Net; //using System.Net;
using System.Reflection; //using System.Reflection;
namespace SimpleHttpServer.Internal; //namespace SimpleHttpServer.Internal;
internal class HttpEndpointHandler { //internal class HttpEndpointHandler {
private static readonly DefaultAuthorizer defaultAuth = new(); // private static readonly DefaultAuthorizer defaultAuth = new();
private readonly IAuthorizer auth; // private readonly IAuthorizer auth;
private readonly MethodInfo handler; // private readonly MethodInfo handler;
private readonly Dictionary<string, (int pindex, Type type, int pparamIdx)> @params; // private readonly Dictionary<string, (int pindex, Type type, int pparamIdx)> @params;
private readonly Func<Exception, HttpResponseBuilder> errorPageBuilder; // private readonly Func<Exception, HttpResponseBuilder> errorPageBuilder;
public HttpEndpointHandler() { // public HttpEndpointHandler() {
auth = defaultAuth; // auth = defaultAuth;
} // }
public HttpEndpointHandler(IAuthorizer auth) { // public HttpEndpointHandler(IAuthorizer auth) {
} // }
public virtual void Handle(HttpListenerContext ctx) { // public virtual void Handle(HttpListenerContext ctx) {
try { // try {
var (isAuth, authData) = auth.IsAuthenticated(ctx); // var (isAuth, authData) = auth.IsAuthenticated(ctx);
if (!isAuth) { // if (!isAuth) {
throw new HttpHandlingException(401, "Authorization required!"); // throw new HttpHandlingException(401, "Authorization required!");
} // }
// collect parameters // // collect parameters
var invokeParams = new object?[@params.Count + 1]; // var invokeParams = new object?[@params.Count + 1];
var set = new BitArray(@params.Count); // var set = new BitArray(@params.Count);
invokeParams[0] = ctx; // invokeParams[0] = ctx;
// read pparams // // read pparams
// read qparams // // read qparams
var qst = ctx.Request.QueryString; // var qst = ctx.Request.QueryString;
foreach (var qelem in ctx.Request.QueryString.AllKeys) { // foreach (var qelem in ctx.Request.QueryString.AllKeys) {
if (@params.ContainsKey(qelem!)) { // if (@params.ContainsKey(qelem!)) {
var (pindex, type, isPParam) = @params[qelem!]; // var (pindex, type, isPParam) = @params[qelem!];
if (type == typeof(string)) { // if (type == typeof(string)) {
invokeParams[pindex] = ctx.Request.QueryString[qelem!]; // invokeParams[pindex] = ctx.Request.QueryString[qelem!];
set.Set(pindex - 1, true); // set.Set(pindex - 1, true);
} else { // } else {
var elem = JsonConvert.DeserializeObject(ctx.Request.QueryString[qelem!]!, type); // var elem = JsonConvert.DeserializeObject(ctx.Request.QueryString[qelem!]!, type);
if (elem != null) { // if (elem != null) {
invokeParams[pindex] = elem; // invokeParams[pindex] = elem;
set.Set(pindex - 1, true); // set.Set(pindex - 1, true);
} // }
} // }
} // }
} // }
// fill with defaults // // fill with defaults
foreach (var p in @params) { // foreach (var p in @params) {
if (!set.Get(p.Value.pindex)) { // if (!set.Get(p.Value.pindex)) {
invokeParams[p.Value.pindex] = p.Value.type.IsValueType ? Activator.CreateInstance(p.Value.type) : null; // invokeParams[p.Value.pindex] = p.Value.type.IsValueType ? Activator.CreateInstance(p.Value.type) : null;
} // }
} // }
var builder = handler.Invoke(null, invokeParams) as HttpResponseBuilder; // var builder = handler.Invoke(null, invokeParams) as HttpResponseBuilder;
builder!.SendResponse(ctx.Response); // builder!.SendResponse(ctx.Response);
} catch (Exception e) { // } catch (Exception e) {
if (e is TargetInvocationException tex) { // if (e is TargetInvocationException tex) {
e = tex.InnerException!; // e = tex.InnerException!;
} // }
errorPageBuilder(e).SendResponse(ctx.Response); // errorPageBuilder(e).SendResponse(ctx.Response);
} // }
} // }
} //}
+214 -212
View File
@@ -1,243 +1,245 @@
using Konscious.Security.Cryptography; //using Newtonsoft.Json;
using Newtonsoft.Json; //using System.Diagnostics.CodeAnalysis;
using System.Security.Cryptography; //using System.Security.Cryptography;
using System.Text; //using System.Text;
namespace SimpleHttpServer.Login; //namespace SimpleHttpServer.Login;
internal struct SerialLoginData { //internal struct SerialLoginData {
public string salt; // public string passwordSalt;
public string pwd; // public string extraDataSalt;
public string additionalData; // public string pwd;
// public string extraData;
public LoginData toPlainData() { // public LoginData ToPlainData() {
return new LoginData { // return new LoginData {
salt = Convert.FromBase64String(salt), // passwordSalt = Convert.FromBase64String(passwordSalt),
password = Convert.FromBase64String(pwd) // extraDataSalt = Convert.FromBase64String(extraDataSalt)
}; // };
} // }
} //}
internal struct LoginData { //internal struct LoginData {
public byte[] salt; // public byte[] passwordSalt;
public byte[] password; // public byte[] extraDataSalt;
public byte[] encryptedData; // public byte[] passwordHash;
// public byte[] encryptedExtraData;
public SerialLoginData toSerial() { // public SerialLoginData ToSerial() {
return new SerialLoginData { // return new SerialLoginData {
salt = Convert.ToBase64String(salt), // passwordSalt = Convert.ToBase64String(passwordSalt),
pwd = Convert.ToBase64String(password), // extraDataSalt = Convert.ToBase64String(extraDataSalt),
additionalData = Convert.ToBase64String(encryptedData) // pwd = Convert.ToBase64String(passwordHash),
}; // extraData = Convert.ToBase64String(encryptedExtraData)
} // };
} // }
//}
internal struct LoginDataProviderConfig { //internal struct LoginDataProviderConfig {
public int SALT_SIZE = 32; // /// <summary>
public int KEY_LENGTH = 256 / 8; // /// Size of the password salt and the extradata salt. So each salt will be of size <see cref="SALT_SIZE"/>.
public int A2_ITERATIONS = 5; // /// </summary>
public int A2_MEMORY_SIZE = 500_000; // public int SALT_SIZE = 32;
public int A2_PARALLELISM = 8; // public int KEY_LENGTH = 256 / 8;
public int A2_HASH_LENGTH = 256 / 8; // public int PBKDF2_ITERATIONS = 600_000;
public int A2_MAX_CONCURRENT = 4;
public int PBKDF2_ITERATIONS = 600_000;
public LoginDataProviderConfig() { } // public LoginDataProviderConfig() { }
} //}
public class LoginProvider<T> { //public class LoginProvider<TExtraData> {
private static readonly Func<T, byte[]> JsonSerialize = t => Encoding.UTF8.GetBytes(JsonConvert.SerializeObject(t)); // private static readonly Func<TExtraData, byte[]> JsonSerialize = t => Encoding.UTF8.GetBytes(JsonConvert.SerializeObject(t));
private static readonly Func<byte[], T> JsonDeserialize = b => JsonConvert.DeserializeObject<T>(Encoding.UTF8.GetString(b))!; // private static readonly Func<byte[], TExtraData> JsonDeserialize = b => JsonConvert.DeserializeObject<TExtraData>(Encoding.UTF8.GetString(b))!;
private readonly LoginDataProviderConfig config; // [ThreadStatic]
private readonly ReaderWriterLockSlim ldLock = new ReaderWriterLockSlim(LockRecursionPolicy.SupportsRecursion); // private static SHA256? _sha256PerThread;
private readonly string ldPath; // private static SHA256 Sha256PerThread { get => _sha256PerThread ??= SHA256.Create(); }
private readonly Dictionary<string, LoginData> loginData;
private readonly SemaphoreSlim argon2Limit;
private Func<T, byte[]> DataSerializer = JsonSerialize; // private readonly LoginDataProviderConfig config;
private Func<byte[], T> DataDeserializer = JsonDeserialize; // private readonly ReaderWriterLockSlim ldLock = new ReaderWriterLockSlim(LockRecursionPolicy.SupportsRecursion);
// private readonly string ldPath;
// private readonly Dictionary<string, LoginData> loginDatas;
public LoginProvider(string ldPath, string confPath) { // private Func<TExtraData, byte[]> DataSerializer = JsonSerialize;
this.ldPath = ldPath; // private Func<byte[], TExtraData> DataDeserializer = JsonDeserialize;
loginData = LoadLoginData(ldPath); // public void SetDataSerializers(Func<TExtraData, byte[]> serializer, Func<byte[], TExtraData> deserializer) {
config = LoadArgon2Config(confPath); // DataSerializer = serializer ?? JsonSerialize;
argon2Limit = new SemaphoreSlim(config.A2_MAX_CONCURRENT); // DataDeserializer = deserializer ?? JsonDeserialize;
} // }
private static Dictionary<string, LoginData> LoadLoginData(string path) {
Dictionary<string, SerialLoginData> tempData;
if (!File.Exists(path)) {
File.WriteAllText(path, "{}", Encoding.UTF8);
tempData = new();
} else {
tempData = JsonConvert.DeserializeObject<Dictionary<string, SerialLoginData>>(File.ReadAllText(path))!;
if (tempData == null) {
throw new InvalidDataException($"could not read login data from file {path}");
}
}
var ld = new Dictionary<string, LoginData>();
foreach (var pair in tempData!) {
ld.Add(pair.Key, pair.Value.toPlainData());
}
return ld;
}
private static LoginDataProviderConfig LoadArgon2Config(string path) { // public LoginProvider(string ldPath, string confPath) {
if (!File.Exists(path)) { // this.ldPath = ldPath;
var conf = new LoginDataProviderConfig(); // loginDatas = LoadLoginDatas(ldPath);
File.WriteAllText(path, JsonConvert.SerializeObject(conf)); // config = LoadLoginProviderConfig(confPath);
return conf; // }
}
return JsonConvert.DeserializeObject<LoginDataProviderConfig>(File.ReadAllText(path));
}
public void SetDataSerialization(Func<T, byte[]> serializer, Func<byte[], T> deserializer) { // private static Dictionary<string, LoginData> LoadLoginDatas(string path) {
DataSerializer = serializer ?? JsonSerialize; // Dictionary<string, SerialLoginData> tempData;
DataDeserializer = deserializer ?? JsonDeserialize; // if (!File.Exists(path)) {
} // File.WriteAllText(path, "{}", Encoding.UTF8);
// tempData = new();
// } else {
// tempData = JsonConvert.DeserializeObject<Dictionary<string, SerialLoginData>>(File.ReadAllText(path))!;
// if (tempData == null) {
// throw new InvalidDataException($"could not read login data from file {path}");
// }
// }
// var ld = new Dictionary<string, LoginData>();
// foreach (var pair in tempData) {
// ld.Add(pair.Key, pair.Value.ToPlainData());
// }
// return ld;
// }
private void StoreLoginData() { // private void SaveLoginData() {
var serial = new Dictionary<string, SerialLoginData>(); // var serial = new Dictionary<string, SerialLoginData>();
ldLock.EnterWriteLock(); // ldLock.EnterWriteLock();
try { // try {
foreach (var pair in loginData!) { // foreach (var pair in loginDatas) {
serial.Add(pair.Key, pair.Value.toSerial()); // serial.Add(pair.Key, pair.Value.ToSerial());
} // }
} finally { // } finally {
ldLock.ExitWriteLock(); // ldLock.ExitWriteLock();
} // }
File.WriteAllText(ldPath, JsonConvert.SerializeObject(serial)); // File.WriteAllText(ldPath, JsonConvert.SerializeObject(serial));
} // }
public bool AddUser(string username, string password, T additional) { // private static LoginDataProviderConfig LoadLoginProviderConfig(string path) {
ldLock.EnterWriteLock(); // if (!File.Exists(path)) {
try { // var conf = new LoginDataProviderConfig();
if (loginData.ContainsKey(username)) { // File.WriteAllText(path, JsonConvert.SerializeObject(conf));
return false; // return conf;
} // }
var salt = RandomNumberGenerator.GetBytes(config.SALT_SIZE); // return JsonConvert.DeserializeObject<LoginDataProviderConfig>(File.ReadAllText(path));
var pwdHash = HashPwd(password, salt); // }
LoginData ld = new LoginData() {
salt = salt,
password = pwdHash,
encryptedData = EncryptAdditionalData(password, salt, additional)
};
loginData.Add(username, ld);
StoreLoginData();
} finally {
ldLock.ExitWriteLock();
}
return true;
}
public bool RemoveUser(string username) { // public bool AddUser(string username, string password, TExtraData additional) {
ldLock.EnterWriteLock(); // ldLock.EnterWriteLock();
try { // try {
var removed = loginData.Remove(username); // if (loginDatas.ContainsKey(username)) {
if (removed) { // return false;
StoreLoginData(); // }
} // var passwordSalt = RandomNumberGenerator.GetBytes(config.SALT_SIZE);
return removed; // var extraDataSalt = RandomNumberGenerator.GetBytes(config.SALT_SIZE);
} finally { // LoginData ld = new LoginData() {
ldLock.ExitWriteLock(); // passwordSalt = passwordSalt,
} // extraDataSalt = extraDataSalt,
} // passwordHash = ComputeSaltedSha256Hash(password, passwordSalt),
// encryptedExtraData = EncryptExtraData(password, extraDataSalt, additional),
// };
// loginDatas.Add(username, ld);
// SaveLoginData();
// } finally {
// ldLock.ExitWriteLock();
// }
// return true;
// }
public bool ModifyUser(string username, string newPassword, T newAdditional) { // public bool RemoveUser(string username) {
ldLock.EnterWriteLock(); // ldLock.EnterWriteLock();
try { // try {
if (!loginData.ContainsKey(username)) { // var removed = loginDatas.Remove(username);
return false; // if (removed) {
} // SaveLoginData();
loginData.Remove(username, out var data); // }
data.password = HashPwd(newPassword, data.salt); // return removed;
data.encryptedData = EncryptAdditionalData(newPassword, data.salt, newAdditional); // } finally {
loginData.Add(username, data); // ldLock.ExitWriteLock();
StoreLoginData(); // }
} finally { // }
ldLock.ExitWriteLock();
}
return true;
}
public (bool, T) Authenticate(string username, string password) { // public bool ModifyUser(string username, string newPassword, TExtraData newExtraData) {
LoginData data; // ldLock.EnterWriteLock();
ldLock.EnterReadLock(); // try {
try { // if (!loginDatas.ContainsKey(username)) {
if (!loginData.TryGetValue(username, out data)) { // return false;
return (false, default(T)!); // }
} // loginDatas.Remove(username, out var data);
} finally { // data.passwordHash = ComputeSaltedSha256Hash(newPassword, data.passwordSalt);
ldLock.ExitReadLock(); // data.encryptedExtraData = EncryptExtraData(newPassword, data.extraDataSalt, newExtraData);
} // loginDatas.Add(username, data);
var hash = HashPwd(password, data.salt); // SaveLoginData();
if (!hash.SequenceEqual(data.password)) { // } finally {
return (false, default(T)!); // ldLock.ExitWriteLock();
} // }
return (true, DecryptAdditionalData(password, data.salt, data.encryptedData)); // return true;
} // }
private byte[] HashPwd(string pwd, byte[] salt) { // public bool TryAuthenticate(string username, string password, [MaybeNullWhen(false)] out TExtraData extraData) {
byte[] hash; // LoginData data;
argon2Limit.Wait(); // ldLock.EnterReadLock();
try { // try {
using (var argon2 = new Argon2id(Encoding.UTF8.GetBytes(pwd))) { // if (!loginDatas.TryGetValue(username, out data)) {
argon2.Iterations = config.A2_ITERATIONS; // extraData = default;
argon2.MemorySize = config.A2_MEMORY_SIZE; // return false;
argon2.DegreeOfParallelism = config.A2_PARALLELISM; // }
argon2.Salt = salt; // } finally {
hash = argon2.GetBytes(config.A2_HASH_LENGTH); // ldLock.ExitReadLock();
} // }
// force collection to reduce sustained memory usage if many hashes are done in close time proximity to each other // var hash = ComputeSaltedSha256Hash(password, data.passwordSalt);
GC.Collect(); // if (!hash.SequenceEqual(data.passwordHash)) {
} finally { // extraData = default;
argon2Limit.Release(); // return false;
} // }
return hash; // extraData = DecryptExtraData(password, data.extraDataSalt, data.encryptedExtraData);
} // return true;
// }
private byte[] EncryptAdditionalData(string pwd, byte[] salt, T data) { // /// <summary>
var pbkdf2 = new Rfc2898DeriveBytes(Encoding.UTF8.GetBytes(pwd), salt, config.PBKDF2_ITERATIONS, HashAlgorithmName.SHA256); // /// Threadsafe as the SHA256 instance (<see cref="Sha256PerThread"/>) is per thread.
var key = pbkdf2.GetBytes(config.KEY_LENGTH / 8); // /// </summary>
// /// <param name="data"></param>
// /// <param name="salt"></param>
// /// <returns></returns>
// private static byte[] ComputeSaltedSha256Hash(string data, byte[] salt) {
// var dataBytes = Encoding.UTF8.GetBytes(data);
// var buf = new byte[data.Length + salt.Length];
// Buffer.BlockCopy(dataBytes, 0, buf, 0, dataBytes.Length);
// Buffer.BlockCopy(salt, 0, buf, dataBytes.Length, salt.Length);
// return Sha256PerThread.ComputeHash(buf);
// }
var plainBytes = DataSerializer(data); // private byte[] EncryptExtraData(string pwd, byte[] salt, TExtraData extraData) {
using var aes = Aes.Create(); // var pbkdf2 = new Rfc2898DeriveBytes(Encoding.UTF8.GetBytes(pwd), salt, config.PBKDF2_ITERATIONS, HashAlgorithmName.SHA256);
aes.KeySize = config.KEY_LENGTH; // var key = pbkdf2.GetBytes(config.KEY_LENGTH / 8);
aes.Key = key;
aes.Mode = CipherMode.CBC;
aes.Padding = PaddingMode.PKCS7;
ICryptoTransform encryptor = aes.CreateEncryptor(aes.Key, aes.IV);
byte[] cipherBytes = encryptor.TransformFinalBlock(plainBytes, 0, plainBytes.Length);
var encryptedBytes = new byte[aes.IV.Length + cipherBytes.Length]; // var plainBytes = DataSerializer(extraData);
Array.Copy(aes.IV, 0, encryptedBytes, 0, aes.IV.Length); // using var aes = Aes.Create();
Array.Copy(cipherBytes, 0, encryptedBytes, aes.IV.Length, cipherBytes.Length); // aes.KeySize = config.KEY_LENGTH;
// aes.Key = key;
// aes.Mode = CipherMode.CBC;
// aes.Padding = PaddingMode.PKCS7;
// ICryptoTransform encryptor = aes.CreateEncryptor(aes.Key, aes.IV);
// byte[] cipherBytes = encryptor.TransformFinalBlock(plainBytes, 0, plainBytes.Length);
return encryptedBytes; // var encryptedBytes = new byte[aes.IV.Length + cipherBytes.Length];
} // Array.Copy(aes.IV, 0, encryptedBytes, 0, aes.IV.Length);
// Array.Copy(cipherBytes, 0, encryptedBytes, aes.IV.Length, cipherBytes.Length);
private T DecryptAdditionalData(string pwd, byte[] salt, byte[] encryptedData) { // return encryptedBytes;
var pbkdf2 = new Rfc2898DeriveBytes(Encoding.UTF8.GetBytes(pwd), salt, config.PBKDF2_ITERATIONS, HashAlgorithmName.SHA256); // }
var key = pbkdf2.GetBytes(config.KEY_LENGTH / 8);
using var aes = Aes.Create(); // private TExtraData DecryptExtraData(string pwd, byte[] salt, byte[] encryptedData) {
aes.KeySize = config.KEY_LENGTH; // var pbkdf2 = new Rfc2898DeriveBytes(Encoding.UTF8.GetBytes(pwd), salt, config.PBKDF2_ITERATIONS, HashAlgorithmName.SHA256);
aes.Key = key; // var key = pbkdf2.GetBytes(config.KEY_LENGTH / 8);
aes.Mode = CipherMode.CBC;
aes.Padding = PaddingMode.PKCS7;
var iv = new byte[aes.BlockSize / 8];
var cipherBytes = new byte[encryptedData.Length - iv.Length];
Array.Copy(encryptedData, 0, iv, 0, iv.Length); // using var aes = Aes.Create();
Array.Copy(encryptedData, iv.Length, cipherBytes, 0, cipherBytes.Length); // aes.KeySize = config.KEY_LENGTH;
// aes.Key = key;
// aes.Mode = CipherMode.CBC;
// aes.Padding = PaddingMode.PKCS7;
// var iv = new byte[aes.BlockSize / 8];
// var cipherBytes = new byte[encryptedData.Length - iv.Length];
aes.IV = iv; // Array.Copy(encryptedData, 0, iv, 0, iv.Length);
ICryptoTransform decryptor = aes.CreateDecryptor(aes.Key, aes.IV); // Array.Copy(encryptedData, iv.Length, cipherBytes, 0, cipherBytes.Length);
byte[] plainBytes = decryptor.TransformFinalBlock(cipherBytes, 0, cipherBytes.Length);
return DataDeserializer(plainBytes); // aes.IV = iv;
} // ICryptoTransform decryptor = aes.CreateDecryptor(aes.Key, aes.IV);
} // byte[] plainBytes = decryptor.TransformFinalBlock(cipherBytes, 0, cipherBytes.Length);
// return DataDeserializer(plainBytes);
// }
//}
-1
View File
@@ -7,7 +7,6 @@
</PropertyGroup> </PropertyGroup>
<ItemGroup> <ItemGroup>
<PackageReference Include="Konscious.Security.Cryptography.Argon2" Version="1.3.0" />
<PackageReference Include="Newtonsoft.Json" Version="13.0.3" /> <PackageReference Include="Newtonsoft.Json" Version="13.0.3" />
</ItemGroup> </ItemGroup>
@@ -0,0 +1,15 @@
using System.Net;
namespace SimpleHttpServer.Types;
[AttributeUsage(AttributeTargets.Method | AttributeTargets.Class, Inherited = true, AllowMultiple = true)]
public abstract class BaseEndpointCheckAttribute : Attribute {
public BaseEndpointCheckAttribute() { }
/// <summary>
/// Executed when the endpoint is invoked. The endpoint invocation is skipped if any of the checks fail.
/// </summary>
/// <returns>True to allow invocation, false to prevent.</returns>
public abstract bool Check(HttpListenerRequest req);
}
@@ -1,12 +1,19 @@
using System.Reflection; using System.Net;
using System.Reflection;
namespace SimpleHttpServer.Types; namespace SimpleHttpServer.Types;
internal struct EndpointInvocationInfo { internal readonly struct EndpointInvocationInfo {
internal readonly MethodInfo methodInfo; internal record struct QueryParameterInfo(string Name, Type Type, bool IsOptional);
internal readonly List<(string, (Type type, bool isOptional))> queryParameters;
public EndpointInvocationInfo(MethodInfo methodInfo, List<(string, (Type type, bool isOptional))> queryParameters) { internal readonly MethodInfo methodInfo;
internal readonly List<QueryParameterInfo> queryParameters;
internal readonly BaseEndpointCheckAttribute[] requiredChecks;
public EndpointInvocationInfo(MethodInfo methodInfo, List<QueryParameterInfo> queryParameters, BaseEndpointCheckAttribute[] requiredChecks) {
this.methodInfo = methodInfo ?? throw new ArgumentNullException(nameof(methodInfo)); this.methodInfo = methodInfo ?? throw new ArgumentNullException(nameof(methodInfo));
this.queryParameters = queryParameters ?? throw new ArgumentNullException(nameof(queryParameters)); this.queryParameters = queryParameters ?? throw new ArgumentNullException(nameof(queryParameters));
this.requiredChecks = requiredChecks;
} }
public readonly bool CheckAll(HttpListenerRequest req) => requiredChecks.All(x => x.Check(req));
} }
@@ -1,4 +1,4 @@
namespace SimpleHttpServer; namespace SimpleHttpServer.Types;
public enum HttpRequestType { public enum HttpRequestType {
GET, GET,
@@ -0,0 +1,7 @@
namespace SimpleHttpServer.Types.ParameterConverters;
internal class StringParameterConverter : IParameterConverter {
public bool TryConvertFromString(string value, out object result) {
result = value;
return true;
}
}
@@ -1,19 +1,29 @@
using System.Net; using System.Collections.ObjectModel;
using System.Net;
namespace SimpleHttpServer; namespace SimpleHttpServer.Types;
public class RequestContext : IDisposable { public class RequestContext : IDisposable {
public HttpListenerContext ListenerContext { get; } public HttpListenerContext ListenerContext { get; }
public ReadOnlyDictionary<string, string> ParsedParameters { get; internal set; }
private StreamReader? reqReader; private TextReader? reqReader;
public StreamReader ReqReader => reqReader ??= new(ListenerContext.Request.InputStream); /// <summary>
/// THREADSAFE
/// </summary>
public TextReader ReqReader => reqReader ??= TextReader.Synchronized(new StreamReader(ListenerContext.Request.InputStream));
private StreamWriter? respWriter; private TextWriter? respWriter;
public StreamWriter RespWriter => respWriter ??= new(ListenerContext.Response.OutputStream) { NewLine = "\n" }; /// <summary>
/// THREADSAFE
/// </summary>
public TextWriter RespWriter => respWriter ??= TextWriter.Synchronized(new StreamWriter(ListenerContext.Response.OutputStream) { NewLine = "\n" });
#pragma warning disable CS8618 // Non-nullable field must contain a non-null value when exiting constructor. Consider declaring as nullable.
public RequestContext(HttpListenerContext listenerContext) { public RequestContext(HttpListenerContext listenerContext) {
ListenerContext = listenerContext; ListenerContext = listenerContext;
} }
#pragma warning restore CS8618 // Non-nullable field must contain a non-null value when exiting constructor. Consider declaring as nullable.
public async Task WriteLineToRespAsync(string resp) => await RespWriter.WriteLineAsync(resp); public async Task WriteLineToRespAsync(string resp) => await RespWriter.WriteLineAsync(resp);
public async Task WriteToRespAsync(string resp) => await RespWriter.WriteAsync(resp); public async Task WriteToRespAsync(string resp) => await RespWriter.WriteAsync(resp);
@@ -25,27 +35,40 @@ public class RequestContext : IDisposable {
public void SetStatusCode(HttpStatusCode status) => SetStatusCode((int) status); public void SetStatusCode(HttpStatusCode status) => SetStatusCode((int) status);
public void SetStatusCodeAndDispose(int status) { public async Task SetStatusCodeAndDisposeAsync(int status) {
using (this) using (this) {
SetStatusCode(status); SetStatusCode(status);
await WriteToRespAsync("\n\n");
await RespWriter.FlushAsync();
}
} }
public void SetStatusCodeAndDispose(HttpStatusCode status) { public async Task SetStatusCodeAndDisposeAsync(HttpStatusCode status) {
using (this) using (this) {
SetStatusCode((int) status); SetStatusCode((int) status);
await WriteToRespAsync("\n\n");
await RespWriter.FlushAsync();
}
} }
public void SetStatusCodeAndDispose(int status, string description) { public async Task SetStatusCodeAndDisposeAsync(int status, string description) {
using (this) { using (this) {
ListenerContext.Response.StatusCode = status; ListenerContext.Response.StatusCode = status;
ListenerContext.Response.StatusDescription = description; ListenerContext.Response.StatusDescription = description;
await WriteToRespAsync("\n\n");
await RespWriter.FlushAsync();
} }
} }
public void SetStatusCodeAndDispose(HttpStatusCode status, string description) => SetStatusCodeAndDispose((int) status, description); public async Task SetStatusCodeAndDisposeAsync(HttpStatusCode status, string description) => await SetStatusCodeAndDisposeAsync((int) status, description);
void IDisposable.Dispose() { public async Task WriteRedirect302AndDisposeAsync(string url) {
ListenerContext.Response.AddHeader("Location", url);
await SetStatusCodeAndDisposeAsync(HttpStatusCode.Redirect);
}
public void Dispose() {
reqReader?.Dispose(); reqReader?.Dispose();
respWriter?.Dispose(); respWriter?.Dispose();
GC.SuppressFinalize(this); GC.SuppressFinalize(this);
+92 -6
View File
@@ -1,4 +1,6 @@
using SimpleHttpServer; using SimpleHttpServer;
using SimpleHttpServer.Types;
using System.Net;
namespace SimpleHttpServerTest; namespace SimpleHttpServerTest;
@@ -8,19 +10,36 @@ public class SimpleServerTest {
const int PORT = 8833; const int PORT = 8833;
private HttpServer? activeServer = null; private HttpServer? activeServer = null;
private HttpClient? activeHttpClient = null;
private bool failOnLogError = true;
private static string GetRequestPath(string url) => $"http://localhost:{PORT}/{url.TrimStart('/')}"; private static string GetRequestPath(string url) => $"http://localhost:{PORT}/{url.TrimStart('/')}";
private async Task RequestGetStringAsync(string path) => await activeHttpClient!.GetStringAsync(GetRequestPath(path));
private async Task<HttpResponseMessage> AssertGetStatusCodeAsync(string path, HttpStatusCode statusCode) {
var resp = await activeHttpClient!.GetAsync(GetRequestPath(path));
Assert.AreEqual(statusCode, resp.StatusCode);
return resp;
}
[TestInitialize] [TestInitialize]
public void Init() { public void Init() {
var conf = new SimpleHttpServerConfiguration(); var conf = new SimpleHttpServerConfiguration() {
DisableLogMessagePrinting = false,
LogMessageHandler = (LogOutputTopic topic, string message, LogOutputLevel logLevel) => {
if (failOnLogError && logLevel is LogOutputLevel.Error or LogOutputLevel.Fatal)
Assert.Fail($"An error was thrown in the log output:\n{topic} {message}");
}
};
if (activeServer != null) if (activeServer != null)
throw new InvalidOperationException("Tried to create another httpserver instance when an existing one was already running."); throw new InvalidOperationException("Tried to create another httpserver instance when an existing one was already running.");
Console.WriteLine("Starting server..."); Console.WriteLine("Starting server...");
failOnLogError = true;
activeServer = new HttpServer(PORT, conf); activeServer = new HttpServer(PORT, conf);
activeServer.RegisterEndpointsFromType<TestEndpoints>(); activeServer.RegisterEndpointsFromType<TestEndpoints>();
activeServer.Start(); activeServer.Start();
activeHttpClient = new HttpClient();
Console.WriteLine("Server started."); Console.WriteLine("Server started.");
} }
@@ -33,20 +52,87 @@ public class SimpleServerTest {
} }
await Console.Out.WriteLineAsync("Shutting down server..."); await Console.Out.WriteLineAsync("Shutting down server...");
await activeServer.StopAsync(ctokSrc.Token); await activeServer.StopAsync(ctokSrc.Token);
activeHttpClient?.Dispose();
activeHttpClient = null;
await Console.Out.WriteLineAsync("Shutdown finished."); await Console.Out.WriteLineAsync("Shutdown finished.");
} }
static string GetHttpPageContentFromPrefix(string page)
=> $"It works!!!!!!56sg5sdf46a4sd65a412f31sdfgdf89h74g9f8h4as56d4f56as2as1f3d24f87g9d87{page}";
[TestMethod] [TestMethod]
public async Task CheckSimpleServe() { public async Task CheckSimpleServe() {
using var hc = new HttpClient(); var resp = await AssertGetStatusCodeAsync("/", HttpStatusCode.OK);
await hc.GetStringAsync(GetRequestPath("/")); var str = await resp.Content.ReadAsStringAsync();
Assert.AreEqual("It works!", str);
}
[TestMethod]
public async Task CheckMultiServe() {
foreach (var item in "index2.html;testpage;testpage2;testpage3".Split(';')) {
await Console.Out.WriteLineAsync($"Checking page: /{item}");
var resp = await AssertGetStatusCodeAsync(item, HttpStatusCode.OK);
var str = await resp.Content.ReadAsStringAsync();
Assert.AreEqual(GetHttpPageContentFromPrefix(item), str);
}
}
[TestMethod]
public async Task CheckQueryArgs() {
foreach (var a1 in "test1;longstring2;something else with a space".Split(';')) {
foreach (var a2 in new[] { -10, 2, -2, 5, 0, 4 }) {
foreach (var a3 in new[] { -1, 9, 2, -20, 0 }) {
foreach (var a4 in new[] { -1, 9, 0 }) {
foreach (var page in "returnqueries;returnqueries2".Split(';')) {
var resp = await AssertGetStatusCodeAsync($"{page}?arg1={a1}&arg2={a2}&arg3={a3}&arg4={a4}", HttpStatusCode.OK);
var str = await resp.Content.ReadAsStringAsync();
Assert.AreEqual(TestEndpoints.GetReturnQueryPageResult(a1, a2, page == "returnqueries2" ? (a3 + a4) : a3), str);
}
}
}
}
}
} }
public class TestEndpoints { public class TestEndpoints {
[HttpEndpoint(HttpRequestType.GET, "/", "index.html")]
[HttpEndpoint(HttpRequestType.GET, "/", "index.html", "amogus.html")]
public static async Task Index(RequestContext req) { public static async Task Index(RequestContext req) {
await req.RespWriter.WriteLineAsync("It works!"); await req.RespWriter.WriteAsync("It works!");
}
[HttpEndpoint(HttpRequestType.GET, "index2.html")]
public static async Task Index2(RequestContext req) {
await req.RespWriter.WriteAsync(GetHttpPageContentFromPrefix("index2.html"));
}
[HttpEndpoint(HttpRequestType.GET, "/testpage")]
public static async Task TestPage(RequestContext req) {
await req.RespWriter.WriteAsync(GetHttpPageContentFromPrefix("testpage"));
}
[HttpEndpoint(HttpRequestType.GET, "testpage2")]
public static async Task TestPage2(RequestContext req) {
await req.RespWriter.WriteAsync(GetHttpPageContentFromPrefix("testpage2"));
}
[HttpEndpoint(HttpRequestType.GET, "/testpage3")]
public static async Task TestPage3(RequestContext req) {
await req.RespWriter.WriteAsync(GetHttpPageContentFromPrefix("testpage3"));
}
public static string GetReturnQueryPageResult(string arg1, int arg2, int arg3) => $"{arg1};{arg2 * 2 - arg3 * 5}";
[HttpEndpoint(HttpRequestType.GET, "/returnqueries")]
public static async Task ReturnQueriesPage(RequestContext req, string arg1, int arg2, int arg3) {
await req.RespWriter.WriteAsync(GetReturnQueryPageResult(arg1, arg2, arg3));
}
[HttpEndpoint(HttpRequestType.GET, "/returnqueries2")]
public static async Task ReturnQueriesPage2(RequestContext req,
[Parameter("arg2")] int arg1, [Parameter("arg1")] string arg2, int arg3, [Parameter("arg4", true)] int arg4) {
// arg4 should be equal to zero as it should get the deafult value because it is not passed to the server
await req.RespWriter.WriteAsync(GetReturnQueryPageResult(arg2, arg1, arg3 + arg4));
} }
} }
} }