using System.Text.Json; using Microsoft.AspNetCore.Http.Features; using Microsoft.Extensions.Primitives; using ReverseLlama.Protocol; namespace ReverseLlama.Server; internal static class ReverseProxyEndpoint { private const string UnauthorizedMessage = "Missing or invalid ReverseLlama token."; private static readonly HashSet HopByHopHeaders = new(StringComparer.OrdinalIgnoreCase) { "Connection", "Expect", "Keep-Alive", "Proxy-Authenticate", "Proxy-Authorization", "TE", "Trailer", "Transfer-Encoding", "Upgrade" }; private static readonly HashSet InternalHeaders = new(StringComparer.OrdinalIgnoreCase) { ProtocolConstants.TokenHeader }; public static async Task HandleRootAsync( HttpContext context, TunnelHub hub, ServerSettings settings, ILoggerFactory loggerFactory, EmbeddingCache embeddingCache, ManagementStore managementStore) { if (!TokenAuthentication.IsAuthorized(context.Request, settings, managementStore, allowQueryToken: false, allowPathToken: true)) { context.Response.StatusCode = StatusCodes.Status401Unauthorized; await context.Response.WriteAsync(UnauthorizedMessage, context.RequestAborted); return; } var pathTokenRemoved = TokenAuthentication.TryRemovePathToken(context.Request.Path, settings, managementStore, out var proxyPath); if (!pathTokenRemoved) { proxyPath = context.Request.Path; } if (pathTokenRemoved && HttpMethods.IsGet(context.Request.Method) && IsRootPath(proxyPath)) { await WriteRootStatusAsync(context, hub); return; } if (TryGetClientAddress(proxyPath, out var pathClientId, out var clientPath)) { await ForwardToClientAsync( context, pathClientId, clientPath, $"{clientPath}{context.Request.QueryString}", hub, settings, loggerFactory, embeddingCache, managementStore); return; } var embeddingRequest = await embeddingCache.TryReadRequestAsync(context.Request, proxyPath); if (embeddingRequest is not null && await embeddingCache.TryWriteCachedResponseAsync(context, embeddingRequest)) { return; } var requestedModel = embeddingRequest?.Model ?? await GetRequestedModelAsync(context.Request, proxyPath); var connection = hub.SelectBest( requestedModel, clientId => !managementStore.GetClientAccess(clientId).IsDisabled); if (connection is null) { if (!hub.HasClient) { context.Response.StatusCode = StatusCodes.Status503ServiceUnavailable; await context.Response.WriteAsync("No tunnel client is connected.", context.RequestAborted); return; } context.Response.StatusCode = StatusCodes.Status503ServiceUnavailable; await context.Response.WriteAsync(GetNoRouteMessage(requestedModel), context.RequestAborted); return; } var pathAndQuery = $"{proxyPath}{context.Request.QueryString}"; await ForwardAsync( context, connection, pathAndQuery, requestedModel, settings, loggerFactory, embeddingCache, embeddingRequest, managementStore); } public static async Task HandleClientAsync( HttpContext context, string clientId, string? path, TunnelHub hub, ServerSettings settings, ILoggerFactory loggerFactory, EmbeddingCache embeddingCache, ManagementStore managementStore) { if (!TokenAuthentication.IsAuthorized(context.Request, settings, managementStore, allowQueryToken: false, allowPathToken: true)) { context.Response.StatusCode = StatusCodes.Status401Unauthorized; await context.Response.WriteAsync(UnauthorizedMessage, context.RequestAborted); return; } var pathAndQuery = $"/{path}{context.Request.QueryString}"; var clientPath = new PathString($"/{path}"); await ForwardToClientAsync( context, clientId, clientPath, pathAndQuery, hub, settings, loggerFactory, embeddingCache, managementStore); } private static async Task ForwardToClientAsync( HttpContext context, string clientId, PathString clientPath, string pathAndQuery, TunnelHub hub, ServerSettings settings, ILoggerFactory loggerFactory, EmbeddingCache embeddingCache, ManagementStore managementStore) { var clientAccess = managementStore.GetClientAccess(clientId); if (clientAccess.IsDisabled) { context.Response.StatusCode = StatusCodes.Status403Forbidden; await context.Response.WriteAsync(GetClientDisabledMessage(clientId, clientAccess), context.RequestAborted); return; } var embeddingRequest = await embeddingCache.TryReadRequestAsync(context.Request, clientPath); if (embeddingRequest is not null && await embeddingCache.TryWriteCachedResponseAsync(context, embeddingRequest)) { return; } var connection = hub.Get(clientId); if (connection is null) { context.Response.StatusCode = StatusCodes.Status503ServiceUnavailable; await context.Response.WriteAsync($"No tunnel client with id '{clientId}' is connected.", context.RequestAborted); return; } var requestedModel = embeddingRequest?.Model ?? await GetRequestedModelAsync(context.Request, clientPath); await ForwardAsync( context, connection, pathAndQuery, requestedModel, settings, loggerFactory, embeddingCache, embeddingRequest, managementStore); } private static bool IsRootPath(PathString path) => string.IsNullOrEmpty(path.Value) || path.Value.Equals("/", StringComparison.Ordinal); private static Task WriteRootStatusAsync(HttpContext context, TunnelHub hub) => context.Response.WriteAsJsonAsync( new { status = "ok", connected = hub.HasClient, pendingRequests = hub.PendingRequestCount, clients = hub.ClientsSnapshot.Count }, context.RequestAborted); private static bool TryGetClientAddress(PathString path, out string clientId, out PathString clientPath) { clientId = ""; clientPath = PathString.Empty; if (!path.StartsWithSegments(new PathString("/clients"), out var pathAfterPrefix)) { return false; } var value = pathAfterPrefix.Value ?? ""; if (value.Length <= 1 || value[0] != '/') { return false; } var nextSlash = value.IndexOf('/', 1); clientId = nextSlash < 0 ? value[1..] : value[1..nextSlash]; if (string.IsNullOrWhiteSpace(clientId)) { return false; } clientPath = nextSlash < 0 ? new PathString("/") : new PathString(value[nextSlash..]); return true; } private static string GetNoRouteMessage(string? requestedModel) => string.IsNullOrWhiteSpace(requestedModel) ? "No tunnel client is available for this request." : $"No connected tunnel client reports model '{requestedModel}'. Check the status endpoint for connected client model lists."; private static string GetClientDisabledMessage(string clientId, ClientAccess access) { if (access.DisabledManually) { return $"Tunnel client '{clientId}' is disabled until it is enabled manually."; } return access.DisabledUntilUtc is { } disabledUntil ? $"Tunnel client '{clientId}' is disabled until {disabledUntil:O}." : $"Tunnel client '{clientId}' is disabled."; } private static async Task GetRequestedModelAsync(HttpRequest request, PathString proxyPath) { if (TryGetModelFromPath(proxyPath, out var pathModel)) { return pathModel; } if (request.Query.TryGetValue("model", out var queryValues)) { var queryModel = queryValues.FirstOrDefault(); if (!string.IsNullOrWhiteSpace(queryModel)) { return queryModel; } } if (!CanHaveBody(request) || (!IsJsonRequest(request) && !IsLikelyModelRequestPath(proxyPath))) { return null; } request.EnableBuffering(); try { using var document = await JsonDocument.ParseAsync( request.Body, cancellationToken: request.HttpContext.RequestAborted); return TryGetModelFromJson(document.RootElement, out var bodyModel) ? bodyModel : null; } catch (JsonException) { return null; } finally { if (request.Body.CanSeek) { request.Body.Position = 0; } } } private static bool TryGetModelFromPath(PathString path, out string model) { model = ""; const string openAiModelPrefix = "/v1/models/"; var value = path.Value ?? ""; if (!value.StartsWith(openAiModelPrefix, StringComparison.OrdinalIgnoreCase)) { return false; } var remaining = value[openAiModelPrefix.Length..]; var nextSlash = remaining.IndexOf('/'); model = Uri.UnescapeDataString(nextSlash < 0 ? remaining : remaining[..nextSlash]); return !string.IsNullOrWhiteSpace(model); } private static bool TryGetModelFromJson(JsonElement root, out string model) { model = ""; if (root.ValueKind != JsonValueKind.Object || !root.TryGetProperty("model", out var modelElement) || modelElement.ValueKind != JsonValueKind.String) { return false; } model = modelElement.GetString() ?? ""; return !string.IsNullOrWhiteSpace(model); } private static bool IsJsonRequest(HttpRequest request) { if (string.IsNullOrWhiteSpace(request.ContentType)) { return false; } var mediaType = request.ContentType.Split(';', 2)[0].Trim(); return mediaType.Equals("application/json", StringComparison.OrdinalIgnoreCase) || mediaType.EndsWith("+json", StringComparison.OrdinalIgnoreCase); } private static bool IsLikelyModelRequestPath(PathString path) { var value = path.Value ?? ""; return value.Equals("/api/generate", StringComparison.OrdinalIgnoreCase) || value.Equals("/api/chat", StringComparison.OrdinalIgnoreCase) || value.Equals("/api/embed", StringComparison.OrdinalIgnoreCase) || value.Equals("/api/embeddings", StringComparison.OrdinalIgnoreCase) || value.Equals("/api/show", StringComparison.OrdinalIgnoreCase) || value.Equals("/v1/chat/completions", StringComparison.OrdinalIgnoreCase) || value.Equals("/v1/completions", StringComparison.OrdinalIgnoreCase) || value.Equals("/v1/embeddings", StringComparison.OrdinalIgnoreCase) || value.Equals("/v1/responses", StringComparison.OrdinalIgnoreCase); } private static async Task ForwardAsync( HttpContext context, TunnelConnection connection, string pathAndQuery, string? requestedModel, ServerSettings settings, ILoggerFactory loggerFactory, EmbeddingCache embeddingCache, EmbeddingCacheRequest? embeddingRequest, ManagementStore managementStore) { var logger = loggerFactory.CreateLogger("ReverseLlama.Server.ReverseProxy"); var requestId = Guid.NewGuid().ToString("n"); var pending = connection.RegisterPending(requestId); var startedAt = DateTimeOffset.UtcNow; var tokenCounter = new ResponseTokenCounter(); int? statusCode = null; var responseCompleted = false; Task? requestBodyTask = null; try { var hasBody = CanHaveBody(context.Request); var requestMessage = new TunnelMessage { Type = TunnelMessageTypes.HttpRequest, RequestId = requestId, Method = context.Request.Method, PathAndQuery = pathAndQuery, HasBody = hasBody, Headers = CollectRequestHeaders(context.Request, settings, managementStore) }; await connection.SendAsync(requestMessage, context.RequestAborted); requestBodyTask = ForwardRequestBodyAsync(context.Request, connection, requestId, hasBody, settings, logger); _ = requestBodyTask.ContinueWith( task => pending.Fail(task.Exception!.GetBaseException()), CancellationToken.None, TaskContinuationOptions.OnlyOnFaulted, TaskScheduler.Default); var responseHeaders = await pending.WaitForHeadersAsync(context.RequestAborted); statusCode = responseHeaders.StatusCode; ApplyResponseHeaders(context.Response, responseHeaders); await context.Response.StartAsync(context.RequestAborted); if (embeddingRequest is not null) { var body = await ReadResponseBodyAsync(pending, context.RequestAborted); tokenCounter.Add(body); await context.Response.Body.WriteAsync(body, context.RequestAborted); await context.Response.Body.FlushAsync(context.RequestAborted); await embeddingCache.StoreResponseAsync( embeddingRequest, responseHeaders, body, CancellationToken.None); } else { await foreach (var chunk in pending.Body.ReadAllAsync(context.RequestAborted)) { tokenCounter.Add(chunk); await context.Response.Body.WriteAsync(chunk, context.RequestAborted); await context.Response.Body.FlushAsync(context.RequestAborted); } } responseCompleted = true; } catch (OperationCanceledException) when (context.RequestAborted.IsCancellationRequested) { logger.LogDebug("Proxy request {RequestId} was cancelled by the downstream caller.", requestId); } catch (Exception exception) { logger.LogWarning(exception, "Proxy request {RequestId} failed.", requestId); if (!context.Response.HasStarted) { context.Response.StatusCode = StatusCodes.Status502BadGateway; await context.Response.WriteAsync(exception.Message, CancellationToken.None); } else { context.Abort(); } } finally { connection.RemovePending(requestId); var completedAt = DateTimeOffset.UtcNow; managementStore.RecordRequest(new RequestMetric( connection.ClientId, requestedModel, context.Request.Method, pathAndQuery, statusCode ?? (context.Response.HasStarted ? context.Response.StatusCode : null), tokenCounter.CountTokens(), startedAt, completedAt, completedAt - startedAt)); if (!responseCompleted && connection.IsOpen) { try { await connection.SendAsync( new TunnelMessage { Type = TunnelMessageTypes.Cancel, RequestId = requestId }, CancellationToken.None); } catch { // The tunnel is already gone; nothing useful remains to notify. } } if (requestBodyTask is { IsCompleted: true }) { try { await requestBodyTask; } catch { // Already reflected through the proxy response path above. } } } } private static async Task ReadResponseBodyAsync(PendingProxyRequest pending, CancellationToken cancellationToken) { using var memory = new MemoryStream(); await foreach (var chunk in pending.Body.ReadAllAsync(cancellationToken)) { await memory.WriteAsync(chunk, cancellationToken); } return memory.ToArray(); } private static async Task ForwardRequestBodyAsync( HttpRequest request, TunnelConnection connection, string requestId, bool hasBody, ServerSettings settings, ILogger logger) { try { if (hasBody) { var buffer = new byte[settings.ChunkSize]; while (true) { var bytesRead = await request.Body.ReadAsync(buffer, request.HttpContext.RequestAborted); if (bytesRead == 0) { break; } await connection.SendAsync( new TunnelMessage { Type = TunnelMessageTypes.HttpRequestBody, RequestId = requestId, Body = buffer.AsSpan(0, bytesRead).ToArray() }, request.HttpContext.RequestAborted); } } await connection.SendAsync( new TunnelMessage { Type = TunnelMessageTypes.HttpRequestComplete, RequestId = requestId }, request.HttpContext.RequestAborted); } catch (Exception exception) { logger.LogDebug(exception, "Failed while forwarding request body {RequestId}.", requestId); throw; } } private static bool CanHaveBody(HttpRequest request) { var bodyDetection = request.HttpContext.Features.Get(); if (bodyDetection?.CanHaveBody is bool canHaveBody) { return canHaveBody; } return request.ContentLength is > 0 || request.Headers.ContainsKey("Transfer-Encoding"); } private static List CollectRequestHeaders( HttpRequest request, ServerSettings settings, ManagementStore managementStore) { var headers = new List(); var skip = HeadersToSkip(request.Headers); foreach (var header in request.Headers) { if (skip.Contains(header.Key) || InternalHeaders.Contains(header.Key)) { continue; } foreach (var value in header.Value) { if (IsOwnBearerToken(header.Key, value, settings, managementStore)) { continue; } headers.Add(new HeaderPair(header.Key, value ?? "")); } } return headers; } // Our token in Bearer form authenticates against the proxy and must not // leak upstream; any other Authorization header is forwarded untouched. private static bool IsOwnBearerToken( string headerName, string? value, ServerSettings settings, ManagementStore managementStore) => string.Equals(headerName, "Authorization", StringComparison.OrdinalIgnoreCase) && TokenAuthentication.IsOwnBearerValue(value, settings, managementStore); private static void ApplyResponseHeaders(HttpResponse response, TunnelMessage responseHeaders) { response.StatusCode = responseHeaders.StatusCode ?? StatusCodes.Status502BadGateway; foreach (var group in responseHeaders.Headers.GroupBy(header => header.Name, StringComparer.OrdinalIgnoreCase)) { if (ShouldSkipResponseHeader(group.Key)) { continue; } response.Headers[group.Key] = new StringValues(group.Select(header => header.Value).ToArray()); } } private static HashSet HeadersToSkip(IHeaderDictionary headers) { var skip = new HashSet(HopByHopHeaders, StringComparer.OrdinalIgnoreCase); if (headers.TryGetValue("Connection", out var connectionHeader)) { foreach (var value in connectionHeader) { foreach (var headerName in value?.Split(',', StringSplitOptions.RemoveEmptyEntries | StringSplitOptions.TrimEntries) ?? []) { skip.Add(headerName); } } } return skip; } private static bool ShouldSkipResponseHeader(string headerName) => HopByHopHeaders.Contains(headerName); }