From 96fde6a0997602895e863bdf582af849777ef54c Mon Sep 17 00:00:00 2001 From: LD-Reborn Date: Tue, 14 Jul 2026 19:09:58 +0200 Subject: [PATCH] Initial commit --- .github/workflows/deploy.yml | 48 ++ .gitignore | 11 + NuGet.Config | 8 + README.md | 106 +++ ReverseLlama.sln | 86 ++ deploy/install-client.sh | 291 +++++++ src/ReverseLlama.Client/ChannelHttpContent.cs | 35 + src/ReverseLlama.Client/ClientOptions.cs | 141 ++++ src/ReverseLlama.Client/Program.cs | 33 + .../Properties/AssemblyInfo.cs | 3 + .../ReverseLlama.Client.csproj | 18 + src/ReverseLlama.Client/TunnelClient.cs | 540 ++++++++++++ src/ReverseLlama.Client/TunnelWorker.cs | 30 + src/ReverseLlama.Client/UpstreamRequest.cs | 245 ++++++ src/ReverseLlama.Protocol/HeaderPair.cs | 18 + .../ProtocolConstants.cs | 10 + .../ReverseLlama.Protocol.csproj | 9 + src/ReverseLlama.Protocol/TunnelMessage.cs | 34 + .../TunnelMessageTypes.cs | 16 + .../WebSocketMessageTransport.cs | 57 ++ src/ReverseLlama.Server/AdminEndpoints.cs | 490 +++++++++++ src/ReverseLlama.Server/EmbeddingCache.cs | 566 +++++++++++++ src/ReverseLlama.Server/ManagementStore.cs | 700 ++++++++++++++++ src/ReverseLlama.Server/PendingCommand.cs | 17 + .../PendingProxyRequest.cs | 52 ++ src/ReverseLlama.Server/Program.cs | 171 ++++ .../Properties/launchSettings.json | 38 + .../ResponseTokenCounter.cs | 157 ++++ .../ReverseLlama.Server.csproj | 20 + .../ReverseProxyEndpoint.cs | 639 +++++++++++++++ src/ReverseLlama.Server/ServerSettings.cs | 98 +++ .../TokenAuthentication.cs | 133 +++ src/ReverseLlama.Server/TunnelConnection.cs | 333 ++++++++ src/ReverseLlama.Server/TunnelHub.cs | 148 ++++ .../appsettings.Example.json | 19 + src/ReverseLlama.Server/appsettings.json | 9 + src/ReverseLlama.Server/wwwroot/admin/app.css | 496 ++++++++++++ src/ReverseLlama.Server/wwwroot/admin/app.js | 766 ++++++++++++++++++ .../wwwroot/admin/index.html | 48 ++ .../PackageAuditTests.cs | 153 ++++ .../ReverseLlama.Client.Tests.csproj | 24 + .../TunnelClientModelTests.cs | 62 ++ .../UpstreamRequestTests.cs | 35 + 43 files changed, 6913 insertions(+) create mode 100644 .github/workflows/deploy.yml create mode 100644 .gitignore create mode 100644 NuGet.Config create mode 100644 README.md create mode 100644 ReverseLlama.sln create mode 100755 deploy/install-client.sh create mode 100644 src/ReverseLlama.Client/ChannelHttpContent.cs create mode 100644 src/ReverseLlama.Client/ClientOptions.cs create mode 100644 src/ReverseLlama.Client/Program.cs create mode 100644 src/ReverseLlama.Client/Properties/AssemblyInfo.cs create mode 100644 src/ReverseLlama.Client/ReverseLlama.Client.csproj create mode 100644 src/ReverseLlama.Client/TunnelClient.cs create mode 100644 src/ReverseLlama.Client/TunnelWorker.cs create mode 100644 src/ReverseLlama.Client/UpstreamRequest.cs create mode 100644 src/ReverseLlama.Protocol/HeaderPair.cs create mode 100644 src/ReverseLlama.Protocol/ProtocolConstants.cs create mode 100644 src/ReverseLlama.Protocol/ReverseLlama.Protocol.csproj create mode 100644 src/ReverseLlama.Protocol/TunnelMessage.cs create mode 100644 src/ReverseLlama.Protocol/TunnelMessageTypes.cs create mode 100644 src/ReverseLlama.Protocol/WebSocketMessageTransport.cs create mode 100644 src/ReverseLlama.Server/AdminEndpoints.cs create mode 100644 src/ReverseLlama.Server/EmbeddingCache.cs create mode 100644 src/ReverseLlama.Server/ManagementStore.cs create mode 100644 src/ReverseLlama.Server/PendingCommand.cs create mode 100644 src/ReverseLlama.Server/PendingProxyRequest.cs create mode 100644 src/ReverseLlama.Server/Program.cs create mode 100644 src/ReverseLlama.Server/Properties/launchSettings.json create mode 100644 src/ReverseLlama.Server/ResponseTokenCounter.cs create mode 100644 src/ReverseLlama.Server/ReverseLlama.Server.csproj create mode 100644 src/ReverseLlama.Server/ReverseProxyEndpoint.cs create mode 100644 src/ReverseLlama.Server/ServerSettings.cs create mode 100644 src/ReverseLlama.Server/TokenAuthentication.cs create mode 100644 src/ReverseLlama.Server/TunnelConnection.cs create mode 100644 src/ReverseLlama.Server/TunnelHub.cs create mode 100644 src/ReverseLlama.Server/appsettings.Example.json create mode 100644 src/ReverseLlama.Server/appsettings.json create mode 100644 src/ReverseLlama.Server/wwwroot/admin/app.css create mode 100644 src/ReverseLlama.Server/wwwroot/admin/app.js create mode 100644 src/ReverseLlama.Server/wwwroot/admin/index.html create mode 100644 tests/ReverseLlama.Client.Tests/PackageAuditTests.cs create mode 100644 tests/ReverseLlama.Client.Tests/ReverseLlama.Client.Tests.csproj create mode 100644 tests/ReverseLlama.Client.Tests/TunnelClientModelTests.cs create mode 100644 tests/ReverseLlama.Client.Tests/UpstreamRequestTests.cs diff --git a/.github/workflows/deploy.yml b/.github/workflows/deploy.yml new file mode 100644 index 0000000..0a9d696 --- /dev/null +++ b/.github/workflows/deploy.yml @@ -0,0 +1,48 @@ +name: Build & Deploy + +on: + push: + branches: [ main ] + +jobs: + build: + runs-on: ubuntu-latest + + steps: + - name: Checkout code + uses: actions/checkout@v4 + + - name: Setup .NET + uses: actions/setup-dotnet@v4 + with: + dotnet-version: '10.0.x' + + - name: Restore dependencies + run: dotnet restore src/ReverseLlama.Server + + - name: Build + run: dotnet build src/ReverseLlama.Server --configuration Release --no-restore + + - name: Publish + run: dotnet publish src/ReverseLlama.Server -c Release -o publish + + - name: Copy files to server + uses: appleboy/scp-action@v0.1.7 + with: + host: ${{ secrets.DEPLOY_HOST }} + username: ${{ secrets.DEPLOY_USER }} + password: ${{ secrets.DEPLOY_SSH_PASSWORD }} + port: ${{ secrets.DEPLOY_SSH_PORT }} + source: "publish/*" + target: "/var/www/ReverseLlama" + strip_components: 1 + + - name: Restart app + uses: appleboy/ssh-action@v1.2.0 + with: + host: ${{ secrets.DEPLOY_HOST }} + username: ${{ secrets.DEPLOY_USER }} + password: ${{ secrets.DEPLOY_SSH_PASSWORD }} + port: ${{ secrets.DEPLOY_SSH_PORT }} + script: | + sudo systemctl restart ReverseLlama.service \ No newline at end of file diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..a7779be --- /dev/null +++ b/.gitignore @@ -0,0 +1,11 @@ +**bin/ +**obj/ +*.user +*.suo +*.log +*.sqlite +*.sqlite-shm +*.sqlite-wal +.vs/ +**appsettings.Development.json +!**/appsettings.Example.json \ No newline at end of file diff --git a/NuGet.Config b/NuGet.Config new file mode 100644 index 0000000..2618b46 --- /dev/null +++ b/NuGet.Config @@ -0,0 +1,8 @@ + + + + + + + + diff --git a/README.md b/README.md new file mode 100644 index 0000000..284165b --- /dev/null +++ b/README.md @@ -0,0 +1,106 @@ +# ReverseLlama + +(Eigenentwicklung; weitesgehend vibecoded) + +ReverseLlama is a small outbound HTTP tunnel for testing Ollama on GPU workstations while exposing the API from a server that cannot reach those workstations directly. + +The client opens a WebSocket connection to the server. The server accepts normal HTTP requests and forwards them through that WebSocket to the client. The client then calls a local upstream such as `http://localhost:11434` and streams the response back. + +## Projects + +- `src/ReverseLlama.Server`: ASP.NET Core server. Exposes the public proxy endpoint and accepts the outbound client tunnel. +- `src/ReverseLlama.Client`: Console client. Runs on the GPU machine and forwards requests to local Ollama. +- `src/ReverseLlama.Protocol`: Shared tunnel message types. + +## Run + +Start the server: + +```powershell +dotnet run --project src/ReverseLlama.Server --urls http://0.0.0.0:5050 -- --token "change-me" +``` + +Start the client on the GPU workstation: + +```powershell +dotnet run --project src/ReverseLlama.Client -- --server http://your-server:5050 --upstream http://localhost:11434 --token "change-me" +``` + +Call Ollama through the server. Model-bearing requests on the root path are routed to a connected client that reports that model, preferring the client with the fewest in-flight requests. You can still address one client explicitly by id: + +```powershell +curl.exe -H "X-Reverse-Llama-Token: change-me" http://your-server:5050/api/tags +curl.exe -H "X-Reverse-Llama-Token: change-me" http://your-server:5050/clients/gpu-01/api/tags +curl.exe http://your-server:5050/token/change-me/api/tags +curl.exe http://your-server:5050/token/change-me/clients/gpu-01/api/tags +``` + +```powershell +curl.exe -H "X-Reverse-Llama-Token: change-me" ` + -H "Content-Type: application/json" ` + -d '{"model":"llama3.1","prompt":"hello"}' ` + http://your-server:5050/api/generate +``` + +## Configuration + +Server options: + +- `--token ` or `REVERSE_LLAMA_TOKEN`: optional shared token. If set, proxy calls must authenticate with `X-Reverse-Llama-Token`, `Authorization: Bearer `, or the `/token//...` path prefix. +- `--tunnel-path `: defaults to `/_reverse-llama/tunnel`. +- `--status-path `: defaults to `/_reverse-llama/status`. +- `--chunk-size ` or `REVERSE_LLAMA_CHUNK_SIZE`: defaults to `65536`. +- `--embedding-cache-path ` or `REVERSE_LLAMA_EMBEDDING_CACHE_PATH`: SQLite cache file for embedding vectors. Defaults to `App_Data\embedding-cache.sqlite` under the server app directory. +- `--management-database-path ` or `REVERSE_LLAMA_MANAGEMENT_DATABASE_PATH`: SQLite database for admin API keys, client disable state, and request/model metrics. Defaults to `App_Data\management.sqlite` under the server app directory. + +Admin UI: + +- `GET /admin` opens the Keycloak-protected management UI. +- The temporary Keycloak settings live under `Authentication:Keycloak` in `appsettings.json`. +- API keys created in the UI are accepted anywhere the shared token is accepted: `X-Reverse-Llama-Token`, `Authorization: Bearer `, `?token=...`, and `/token//...`. +- Model add/remove/load/unload commands are sent through the connected tunnel client to Ollama (`/api/pull`, `/api/delete`, `/api/generate`, and `/api/show`). + +Client options: + +- `--server ` or `REVERSE_LLAMA_SERVER`: server base URL, for example `http://your-server:5050`. +- `--upstream ` or `REVERSE_LLAMA_UPSTREAM`: local Ollama URL, defaults to `http://localhost:11434`. +- `--token ` or `REVERSE_LLAMA_TOKEN`: optional shared token. +- `--client-id ` or `REVERSE_LLAMA_CLIENT_ID`: identifies this machine on the server; defaults to the machine name. +- `--tunnel-path ` or `REVERSE_LLAMA_TUNNEL_PATH`: defaults to `/_reverse-llama/tunnel`. +- `--reconnect-delay ` or `REVERSE_LLAMA_RECONNECT_DELAY_SECONDS`: defaults to `5`. +- `--chunk-size ` or `REVERSE_LLAMA_CHUNK_SIZE`: defaults to `65536`. + +The token is accepted as `X-Reverse-Llama-Token`, as `Authorization: Bearer `, or as a path prefix like `/token//api/tags` or `/token//clients/{id}/v1`. The Bearer form lets OpenAI-compatible clients (e.g. n8n's OpenAI nodes pointed at `/clients/{id}/v1`) authenticate with their API-key field. The path-token form is useful for clients that cannot send custom headers. The server strips its own token header/Bearer value and removes the path prefix before forwarding; any other `Authorization` value is forwarded untouched. + +## Multiple clients + +Any number of machines can connect at the same time; each registers under its client id (machine name by default). + +- `GET/POST /clients/{client-id}/...` forwards to that specific machine. +- The client reports its local Ollama model list from `/api/tags` when it connects and refreshes it every minute. +- The plain root path (`/api/...`, `/v1/...`) routes model-bearing requests to a client that reports the requested model, preferring the lowest in-flight request count. Requests without a model are sent to the connected client with the fewest in-flight requests. +- The status endpoint lists all connected clients, their in-flight request counts, and their last reported model lists. +- If a client connects with an id that is already in use, the old connection is replaced and the replaced client exits instead of reconnecting. + +## Embedding cache + +The server keeps an in-memory KV cache for embedding vectors and persists it to SQLite. The cache key is the requested `model` plus the exact input text. It applies to `POST /api/embed`, `POST /api/embeddings`, and `POST /v1/embeddings`; cache hits return JSON in the same endpoint family shape and include `X-Reverse-Llama-Embedding-Cache: hit`. + +The authenticated status endpoint reports whether the cache is available, plus the cache count and database path. If SQLite cannot be initialized, proxy traffic continues without embedding-cache writes. + +## Client installer + +Linux: + +```bash +sudo bash deploy/install-client.sh --server http://your-server:5050 --token "change-me" +``` + +Options: `--server`, `--token` (required); `--client-id`, `--upstream`, `--install-dir`, `--service-name`, `--no-ollama` (optional). Missing required values are prompted interactively. + +The script ensures .NET 10 and Ollama are installed, builds the client self-contained, installs it to `/opt/reversellama-client`, and creates a systemd service (`reversellama-client`). Logs: `journalctl -u reversellama-client -f`. + +## Notes + +- Request and response bodies are streamed through the tunnel, which is important for Ollama streaming responses. +- Use HTTPS or a private network/VPN when exposing this outside a trusted network. The token is simple shared-secret protection, not a full access-control system. diff --git a/ReverseLlama.sln b/ReverseLlama.sln new file mode 100644 index 0000000..d526251 --- /dev/null +++ b/ReverseLlama.sln @@ -0,0 +1,86 @@ + +Microsoft Visual Studio Solution File, Format Version 12.00 +# Visual Studio Version 17 +VisualStudioVersion = 17.0.31903.59 +MinimumVisualStudioVersion = 10.0.40219.1 +Project("{2150E333-8FDC-42A3-9474-1A3956D46DE8}") = "src", "src", "{827E0CD3-B72D-47B6-A68D-7590B98EB39B}" +EndProject +Project("{FAE04EC0-301F-11D3-BF4B-00C04F79EFBC}") = "ReverseLlama.Protocol", "src\ReverseLlama.Protocol\ReverseLlama.Protocol.csproj", "{D85E1D3C-0AC3-4810-8285-16BE903EC8AB}" +EndProject +Project("{FAE04EC0-301F-11D3-BF4B-00C04F79EFBC}") = "ReverseLlama.Server", "src\ReverseLlama.Server\ReverseLlama.Server.csproj", "{FB853CB0-5AEF-4278-82E7-3C8E506F9DFE}" +EndProject +Project("{FAE04EC0-301F-11D3-BF4B-00C04F79EFBC}") = "ReverseLlama.Client", "src\ReverseLlama.Client\ReverseLlama.Client.csproj", "{0633EAE6-B82A-4A27-851C-B5D6E11CBE03}" +EndProject +Project("{2150E333-8FDC-42A3-9474-1A3956D46DE8}") = "tests", "tests", "{0AB3BF05-4346-4AA6-1389-037BE0695223}" +EndProject +Project("{FAE04EC0-301F-11D3-BF4B-00C04F79EFBC}") = "ReverseLlama.Client.Tests", "tests\ReverseLlama.Client.Tests\ReverseLlama.Client.Tests.csproj", "{B03E7794-2888-4B56-A30F-B6C25BCFB89A}" +EndProject +Global + GlobalSection(SolutionConfigurationPlatforms) = preSolution + Debug|Any CPU = Debug|Any CPU + Debug|x64 = Debug|x64 + Debug|x86 = Debug|x86 + Release|Any CPU = Release|Any CPU + Release|x64 = Release|x64 + Release|x86 = Release|x86 + EndGlobalSection + GlobalSection(ProjectConfigurationPlatforms) = postSolution + {D85E1D3C-0AC3-4810-8285-16BE903EC8AB}.Debug|Any CPU.ActiveCfg = Debug|Any CPU + {D85E1D3C-0AC3-4810-8285-16BE903EC8AB}.Debug|Any CPU.Build.0 = Debug|Any CPU + {D85E1D3C-0AC3-4810-8285-16BE903EC8AB}.Debug|x64.ActiveCfg = Debug|Any CPU + {D85E1D3C-0AC3-4810-8285-16BE903EC8AB}.Debug|x64.Build.0 = Debug|Any CPU + {D85E1D3C-0AC3-4810-8285-16BE903EC8AB}.Debug|x86.ActiveCfg = Debug|Any CPU + {D85E1D3C-0AC3-4810-8285-16BE903EC8AB}.Debug|x86.Build.0 = Debug|Any CPU + {D85E1D3C-0AC3-4810-8285-16BE903EC8AB}.Release|Any CPU.ActiveCfg = Release|Any CPU + {D85E1D3C-0AC3-4810-8285-16BE903EC8AB}.Release|Any CPU.Build.0 = Release|Any CPU + {D85E1D3C-0AC3-4810-8285-16BE903EC8AB}.Release|x64.ActiveCfg = Release|Any CPU + {D85E1D3C-0AC3-4810-8285-16BE903EC8AB}.Release|x64.Build.0 = Release|Any CPU + {D85E1D3C-0AC3-4810-8285-16BE903EC8AB}.Release|x86.ActiveCfg = Release|Any CPU + {D85E1D3C-0AC3-4810-8285-16BE903EC8AB}.Release|x86.Build.0 = Release|Any CPU + {FB853CB0-5AEF-4278-82E7-3C8E506F9DFE}.Debug|Any CPU.ActiveCfg = Debug|Any CPU + {FB853CB0-5AEF-4278-82E7-3C8E506F9DFE}.Debug|Any CPU.Build.0 = Debug|Any CPU + {FB853CB0-5AEF-4278-82E7-3C8E506F9DFE}.Debug|x64.ActiveCfg = Debug|Any CPU + {FB853CB0-5AEF-4278-82E7-3C8E506F9DFE}.Debug|x64.Build.0 = Debug|Any CPU + {FB853CB0-5AEF-4278-82E7-3C8E506F9DFE}.Debug|x86.ActiveCfg = Debug|Any CPU + {FB853CB0-5AEF-4278-82E7-3C8E506F9DFE}.Debug|x86.Build.0 = Debug|Any CPU + {FB853CB0-5AEF-4278-82E7-3C8E506F9DFE}.Release|Any CPU.ActiveCfg = Release|Any CPU + {FB853CB0-5AEF-4278-82E7-3C8E506F9DFE}.Release|Any CPU.Build.0 = Release|Any CPU + {FB853CB0-5AEF-4278-82E7-3C8E506F9DFE}.Release|x64.ActiveCfg = Release|Any CPU + {FB853CB0-5AEF-4278-82E7-3C8E506F9DFE}.Release|x64.Build.0 = Release|Any CPU + {FB853CB0-5AEF-4278-82E7-3C8E506F9DFE}.Release|x86.ActiveCfg = Release|Any CPU + {FB853CB0-5AEF-4278-82E7-3C8E506F9DFE}.Release|x86.Build.0 = Release|Any CPU + {0633EAE6-B82A-4A27-851C-B5D6E11CBE03}.Debug|Any CPU.ActiveCfg = Debug|Any CPU + {0633EAE6-B82A-4A27-851C-B5D6E11CBE03}.Debug|Any CPU.Build.0 = Debug|Any CPU + {0633EAE6-B82A-4A27-851C-B5D6E11CBE03}.Debug|x64.ActiveCfg = Debug|Any CPU + {0633EAE6-B82A-4A27-851C-B5D6E11CBE03}.Debug|x64.Build.0 = Debug|Any CPU + {0633EAE6-B82A-4A27-851C-B5D6E11CBE03}.Debug|x86.ActiveCfg = Debug|Any CPU + {0633EAE6-B82A-4A27-851C-B5D6E11CBE03}.Debug|x86.Build.0 = Debug|Any CPU + {0633EAE6-B82A-4A27-851C-B5D6E11CBE03}.Release|Any CPU.ActiveCfg = Release|Any CPU + {0633EAE6-B82A-4A27-851C-B5D6E11CBE03}.Release|Any CPU.Build.0 = Release|Any CPU + {0633EAE6-B82A-4A27-851C-B5D6E11CBE03}.Release|x64.ActiveCfg = Release|Any CPU + {0633EAE6-B82A-4A27-851C-B5D6E11CBE03}.Release|x64.Build.0 = Release|Any CPU + {0633EAE6-B82A-4A27-851C-B5D6E11CBE03}.Release|x86.ActiveCfg = Release|Any CPU + {0633EAE6-B82A-4A27-851C-B5D6E11CBE03}.Release|x86.Build.0 = Release|Any CPU + {B03E7794-2888-4B56-A30F-B6C25BCFB89A}.Debug|Any CPU.ActiveCfg = Debug|Any CPU + {B03E7794-2888-4B56-A30F-B6C25BCFB89A}.Debug|Any CPU.Build.0 = Debug|Any CPU + {B03E7794-2888-4B56-A30F-B6C25BCFB89A}.Debug|x64.ActiveCfg = Debug|Any CPU + {B03E7794-2888-4B56-A30F-B6C25BCFB89A}.Debug|x64.Build.0 = Debug|Any CPU + {B03E7794-2888-4B56-A30F-B6C25BCFB89A}.Debug|x86.ActiveCfg = Debug|Any CPU + {B03E7794-2888-4B56-A30F-B6C25BCFB89A}.Debug|x86.Build.0 = Debug|Any CPU + {B03E7794-2888-4B56-A30F-B6C25BCFB89A}.Release|Any CPU.ActiveCfg = Release|Any CPU + {B03E7794-2888-4B56-A30F-B6C25BCFB89A}.Release|Any CPU.Build.0 = Release|Any CPU + {B03E7794-2888-4B56-A30F-B6C25BCFB89A}.Release|x64.ActiveCfg = Release|Any CPU + {B03E7794-2888-4B56-A30F-B6C25BCFB89A}.Release|x64.Build.0 = Release|Any CPU + {B03E7794-2888-4B56-A30F-B6C25BCFB89A}.Release|x86.ActiveCfg = Release|Any CPU + {B03E7794-2888-4B56-A30F-B6C25BCFB89A}.Release|x86.Build.0 = Release|Any CPU + EndGlobalSection + GlobalSection(SolutionProperties) = preSolution + HideSolutionNode = FALSE + EndGlobalSection + GlobalSection(NestedProjects) = preSolution + {D85E1D3C-0AC3-4810-8285-16BE903EC8AB} = {827E0CD3-B72D-47B6-A68D-7590B98EB39B} + {FB853CB0-5AEF-4278-82E7-3C8E506F9DFE} = {827E0CD3-B72D-47B6-A68D-7590B98EB39B} + {0633EAE6-B82A-4A27-851C-B5D6E11CBE03} = {827E0CD3-B72D-47B6-A68D-7590B98EB39B} + {B03E7794-2888-4B56-A30F-B6C25BCFB89A} = {0AB3BF05-4346-4AA6-1389-037BE0695223} + EndGlobalSection +EndGlobal diff --git a/deploy/install-client.sh b/deploy/install-client.sh new file mode 100755 index 0000000..edfa29e --- /dev/null +++ b/deploy/install-client.sh @@ -0,0 +1,291 @@ +#!/usr/bin/env bash +set -euo pipefail + +# ── Defaults ────────────────────────────────────────────────────────────────── +DEFAULT_INSTALL_DIR="/opt/reversellama-client" +DEFAULT_SERVICE_NAME="reversellama-client" +DEFAULT_UPSTREAM="http://localhost:11434" + +SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)" +REPO_ROOT="$(cd "$SCRIPT_DIR/.." && pwd)" + +INSTALL_DIR="$DEFAULT_INSTALL_DIR" +SERVICE_NAME="$DEFAULT_SERVICE_NAME" +SERVER_URL="" +TOKEN="" +CLIENT_ID="$(hostname -s 2>/dev/null || echo "linux-client")" +UPSTREAM="$DEFAULT_UPSTREAM" +SKIP_OLLAMA=false + +# ── Colors ──────────────────────────────────────────────────────────────────── +RED='\033[0;31m' +GREEN='\033[0;32m' +YELLOW='\033[0;33m' +NC='\033[0m' + +info() { echo -e "${GREEN}[INFO]${NC} $*"; } +warn() { echo -e "${YELLOW}[WARN]${NC} $*"; } +err() { echo -e "${RED}[ERROR]${NC} $*" >&2; } +die() { err "$@"; exit 1; } + +# ── Usage ───────────────────────────────────────────────────────────────────── +usage() { + cat < ReverseLlama server URL (e.g. http://my-server:5050) + --token Shared secret token for the server + +Optional: + --client-id Client identifier; defaults to hostname + --upstream Local Ollama URL; defaults to $DEFAULT_UPSTREAM + --install-dir Install directory; defaults to $DEFAULT_INSTALL_DIR + --service-name systemd service name; defaults to $DEFAULT_SERVICE_NAME + --no-ollama Skip Ollama installation and status check + -h, --help Show this help message + +Examples: + $0 --server http://gpu-server:5050 --token "my-secret" + $0 --server http://gpu-server:5050 --token "my-secret" --no-ollama +EOF + exit 0 +} + +# ── Argument parsing ────────────────────────────────────────────────────────── +while [[ $# -gt 0 ]]; do + case "$1" in + --server) SERVER_URL="$2"; shift 2 ;; + --token) TOKEN="$2"; shift 2 ;; + --client-id) CLIENT_ID="$2"; shift 2 ;; + --upstream) UPSTREAM="$2"; shift 2 ;; + --install-dir) INSTALL_DIR="$2"; shift 2 ;; + --service-name) SERVICE_NAME="$2"; shift 2 ;; + --no-ollama) SKIP_OLLAMA=true; shift ;; + -h|--help) usage ;; + *) die "Unknown option: $1" ;; + esac +done + +# ── Prompt for missing required values ─────────────────────────────────────── +if [[ -z "$SERVER_URL" ]]; then + read -rp "ReverseLlama server URL (e.g. http://my-server:5050): " SERVER_URL +fi +if [[ -z "$SERVER_URL" ]]; then + die "Server URL is required." +fi + +if [[ -z "$TOKEN" ]]; then + read -rsp "Server token: " TOKEN + echo +fi +if [[ -z "$TOKEN" ]]; then + die "Token is required." +fi + +# ── Root check ──────────────────────────────────────────────────────────────── +if [[ $EUID -ne 0 ]]; then + die "This script must be run as root (or with sudo)." +fi + +# ── Check / install .NET 10 SDK ────────────────────────────────────────────── +install_dotnet() { + info "Installing .NET 10 SDK..." + + if command -v apt-get &>/dev/null; then + apt-get update -qq + apt-get install -y -qq wget apt-transport-https + + # Detect distro codename for the Microsoft package repo + if grep -qi "ubuntu" /etc/os-release 2>/dev/null; then + DISTRO="ubuntu" + CODENAME="$(. /etc/os-release && echo "$VERSION_CODENAME")" + case "$CODENAME" in + noble|jammy|focal) ;; # supported + *) + warn "Ubuntu $CODENAME may not have .NET 10 packages yet; using latest available." + CODENAME="noble" + ;; + esac + elif grep -qi "debian" /etc/os-release 2>/dev/null; then + DISTRO="debian" + CODENAME="$(. /etc/os-release && echo "$VERSION_CODENAME")" + case "$CODENAME" in + trixie|bookworm|bullseye) ;; + *) + warn "Debian $CODENAME may not have .NET 10 packages; using latest available." + CODENAME="trixie" + ;; + esac + else + die "Unsupported distro. Install .NET 10 SDK manually: https://dotnet.microsoft.com/download" + fi + + # Add Microsoft package repository GPG key and repo + curl -fsSL "https://packages.microsoft.com/config/$DISTRO/$CODENAME/packages-microsoft-prod.deb" \ + -o /tmp/packages-microsoft-prod.deb + dpkg -i /tmp/packages-microsoft-prod.deb + rm -f /tmp/packages-microsoft-prod.deb + + apt-get update -qq + apt-get install -y -qq dotnet-sdk-10.0 + + elif command -v dnf &>/dev/null || command -v yum &>/dev/null; then + PKG_MGR="dnf" + command -v dnf &>/dev/null || PKG_MGR="yum" + + # Add Microsoft package repository for RHEL/CentOS/Fedora + cat > /etc/yum.repos.d/microsoft-prod.repo <<'REPO' +[microsoft-prod] +name=Microsoft Production Repository +baseurl=https://packages.microsoft.com/rhel/9/prod/ +enabled=1 +gpgcheck=1 +gpgkey=https://packages.microsoft.com/keys/microsoft.asc +REPO + + $PKG_MGR install -y dotnet-sdk-10.0 + + elif command -v pacman &>/dev/null; then + warn "Arch Linux detected. .NET 10 may need to be installed from the AUR or manually." + warn "See: https://dotnet.microsoft.com/download" + die "Cannot auto-install .NET SDK on Arch. Please install dotnet-sdk-10 manually." + + else + die "Unsupported package manager. Install .NET 10 SDK manually: https://dotnet.microsoft.com/download" + fi + + DOTNET_CMD="$(command -v dotnet 2>/dev/null || echo /usr/share/dotnet/dotnet)" + if [[ ! -x "$DOTNET_CMD" ]]; then + DOTNET_CMD="/usr/bin/dotnet" + fi + if [[ ! -x "$DOTNET_CMD" ]]; then + die ".NET SDK installation succeeded but dotnet binary not found. Please add it to PATH." + fi + + DOTNET_VER="$("$DOTNET_CMD" --version 2>/dev/null || true)" + info ".NET SDK installed: $DOTNET_CMD ($DOTNET_VER)" +} + +DOTNET_CMD="" +if command -v dotnet &>/dev/null; then + DOTNET_CMD="$(command -v dotnet)" + DOTNET_VER="$("$DOTNET_CMD" --version 2>/dev/null || true)" + if [[ "$DOTNET_VER" == 10.* ]]; then + info "dotnet 10 is already installed: $DOTNET_CMD ($DOTNET_VER)" + else + warn "dotnet is installed ($DOTNET_VER) but version 10.x is required." + install_dotnet + fi +else + install_dotnet +fi + +# ── Check / install Ollama ─────────────────────────────────────────────────── +if [[ "$SKIP_OLLAMA" == "false" ]]; then + if command -v ollama &>/dev/null; then + info "Ollama is installed: $(command -v ollama)" + if systemctl is-active --quiet ollama 2>/dev/null; then + info "Ollama service is running." + else + warn "Ollama is installed but not running. Starting..." + systemctl start ollama + systemctl enable ollama 2>/dev/null || true + info "Ollama started." + fi + else + info "Ollama is not installed. Installing..." + curl -fsSL https://ollama.com/install.sh | sh + info "Ollama installed." + systemctl start ollama + systemctl enable ollama 2>/dev/null || true + info "Ollama started." + fi +else + info "Skipping Ollama check (--no-ollama)." +fi + +# ── Build client from source ───────────────────────────────────────────────── +CLIENT_SRC="$REPO_ROOT/src/ReverseLlama.Client" +if [[ ! -d "$CLIENT_SRC" ]]; then + die "Client source not found at $CLIENT_SRC. Run this script from the repository or pass --install-dir." +fi + +info "Building ReverseLlama client (self-contained, linux-x64)..." +BUILD_DIR="$(mktemp -d /tmp/reversellama-build.XXXXXX)" +trap 'rm -rf "$BUILD_DIR"' EXIT + +"$DOTNET_CMD" publish "$CLIENT_SRC/ReverseLlama.Client.csproj" \ + -c Release \ + -r linux-x64 \ + --self-contained true \ + -o "$BUILD_DIR" + +if [[ ! -f "$BUILD_DIR/ReverseLlama.Client" ]]; then + die "Build failed. ReverseLlama.Client binary not found in output." +fi + +info "Build successful." + +# ── Install ─────────────────────────────────────────────────────────────────── +info "Installing to $INSTALL_DIR..." +mkdir -p "$INSTALL_DIR" +cp -a "$BUILD_DIR"/. "$INSTALL_DIR/" +chmod +x "$INSTALL_DIR/ReverseLlama.Client" + +info "Client installed to $INSTALL_DIR." + +# ── Create systemd service ─────────────────────────────────────────────────── +SERVICE_FILE="/etc/systemd/system/${SERVICE_NAME}.service" + +if systemctl list-unit-files "$SERVICE_NAME.service" &>/dev/null 2>&1; then + info "Stopping existing service $SERVICE_NAME..." + systemctl stop "$SERVICE_NAME" 2>/dev/null || true +fi + +cat > "$SERVICE_FILE" < _reader; + + public ChannelHttpContent(ChannelReader reader, CancellationToken cancellationToken) + { + _reader = reader; + _cancellationToken = cancellationToken; + } + + protected override Task SerializeToStreamAsync(Stream stream, TransportContext? context) => + SerializeToStreamAsync(stream, context, _cancellationToken); + + protected override async Task SerializeToStreamAsync(Stream stream, TransportContext? context, CancellationToken cancellationToken) + { + using var linkedCts = CancellationTokenSource.CreateLinkedTokenSource(_cancellationToken, cancellationToken); + + await foreach (var chunk in _reader.ReadAllAsync(linkedCts.Token)) + { + await stream.WriteAsync(chunk, linkedCts.Token); + } + } + + protected override bool TryComputeLength(out long length) + { + length = -1; + return false; + } +} diff --git a/src/ReverseLlama.Client/ClientOptions.cs b/src/ReverseLlama.Client/ClientOptions.cs new file mode 100644 index 0000000..ed241c4 --- /dev/null +++ b/src/ReverseLlama.Client/ClientOptions.cs @@ -0,0 +1,141 @@ +using ReverseLlama.Protocol; + +namespace ReverseLlama.Client; + +internal sealed class ClientOptions +{ + public Uri Server { get; init; } = new("http://localhost:5001"); + + public Uri Upstream { get; init; } = new("http://localhost:11434"); + + public string TunnelPath { get; init; } = ProtocolConstants.DefaultTunnelPath; + + public string? Token { get; init; } + + public string ClientId { get; init; } = Environment.MachineName.ToLowerInvariant(); + + public TimeSpan ReconnectDelay { get; init; } = TimeSpan.FromSeconds(5); + + public int ChunkSize { get; init; } = 64 * 1024; + + public Uri TunnelUri + { + get + { + var builder = new UriBuilder(Server); + builder.Scheme = builder.Scheme.ToLowerInvariant() switch + { + "http" => "ws", + "https" => "wss", + "ws" => "ws", + "wss" => "wss", + var unsupported => throw new InvalidOperationException($"Unsupported server URI scheme '{unsupported}'. Use http, https, ws, or wss.") + }; + + if (string.IsNullOrWhiteSpace(builder.Path) || builder.Path == "/") + { + builder.Path = NormalizePath(TunnelPath); + } + + return builder.Uri; + } + } + + public static ClientOptions Parse(string[] args) + { + var values = ParseArgs(args); + + return new ClientOptions + { + Server = ReadUri(values, "server", "REVERSE_LLAMA_SERVER", "http://localhost:5001"), + Upstream = ReadUri(values, "upstream", "REVERSE_LLAMA_UPSTREAM", "http://localhost:11434"), + TunnelPath = NormalizePath(Read(values, "tunnel-path", "REVERSE_LLAMA_TUNNEL_PATH") ?? ProtocolConstants.DefaultTunnelPath), + Token = Read(values, "token", "REVERSE_LLAMA_TOKEN"), + ClientId = Read(values, "client-id", "REVERSE_LLAMA_CLIENT_ID") ?? Environment.MachineName.ToLowerInvariant(), + ReconnectDelay = TimeSpan.FromSeconds(ReadInt(values, 5, "reconnect-delay", "REVERSE_LLAMA_RECONNECT_DELAY_SECONDS")), + ChunkSize = ReadInt(values, 64 * 1024, "chunk-size", "REVERSE_LLAMA_CHUNK_SIZE") + }; + } + + public static string Usage => + """ + ReverseLlama.Client options: + --server Server base URL, e.g. http://my-server:5050 + --upstream Local upstream URL, e.g. http://localhost:11434 + --token Optional token matching the server + --client-id Identifies this machine on the server; defaults to the machine name + --tunnel-path Defaults to /_reverse-llama/tunnel + --reconnect-delay Defaults to 5 + --chunk-size Defaults to 65536 + """; + + private static Dictionary ParseArgs(string[] args) + { + var values = new Dictionary(StringComparer.OrdinalIgnoreCase); + + for (var i = 0; i < args.Length; i++) + { + var arg = args[i]; + if (!arg.StartsWith("--", StringComparison.Ordinal)) + { + throw new ArgumentException($"Unexpected argument '{arg}'."); + } + + var keyValue = arg[2..].Split('=', 2); + if (keyValue.Length == 2) + { + values[keyValue[0]] = keyValue[1]; + continue; + } + + if (i + 1 >= args.Length || args[i + 1].StartsWith("--", StringComparison.Ordinal)) + { + throw new ArgumentException($"Missing value for '{arg}'."); + } + + values[keyValue[0]] = args[++i]; + } + + return values; + } + + private static string? Read(Dictionary values, params string[] keys) + { + foreach (var key in keys) + { + if (values.TryGetValue(key, out var value) && !string.IsNullOrWhiteSpace(value)) + { + return value; + } + + value = Environment.GetEnvironmentVariable(key); + if (!string.IsNullOrWhiteSpace(value)) + { + return value; + } + } + + return null; + } + + private static int ReadInt(Dictionary values, int fallback, params string[] keys) + { + var value = Read(values, keys); + return int.TryParse(value, out var parsed) && parsed > 0 ? parsed : fallback; + } + + private static Uri ReadUri(Dictionary values, string key, string envKey, string fallback) + { + var value = Read(values, key, envKey) ?? fallback; + + if (!Uri.TryCreate(value, UriKind.Absolute, out var uri)) + { + throw new ArgumentException($"'{value}' is not an absolute URI."); + } + + return uri; + } + + private static string NormalizePath(string path) => + path.StartsWith('/') ? path : $"/{path}"; +} diff --git a/src/ReverseLlama.Client/Program.cs b/src/ReverseLlama.Client/Program.cs new file mode 100644 index 0000000..e295caf --- /dev/null +++ b/src/ReverseLlama.Client/Program.cs @@ -0,0 +1,33 @@ +using Microsoft.Extensions.DependencyInjection; +using Microsoft.Extensions.Hosting; +using Microsoft.Extensions.Logging; +using ReverseLlama.Client; + +try +{ + var options = ClientOptions.Parse(args); + + Console.WriteLine("ReverseLlama client"); + Console.WriteLine($" client id: {options.ClientId}"); + Console.WriteLine($" server tunnel: {options.TunnelUri}"); + Console.WriteLine($" local upstream: {options.Upstream}"); + + // Args are parsed by ClientOptions; keep them away from the host configuration. + var builder = Host.CreateApplicationBuilder(new HostApplicationBuilderSettings { Args = [] }); + builder.Services.AddSingleton(options); + builder.Services.AddSingleton(); + builder.Services.AddHostedService(); + builder.Services.AddWindowsService(service => service.ServiceName = "ReverseLlamaClient"); + // The EventLog provider defaults to Warning; connection state is worth seeing there. + builder.Logging.AddFilter("ReverseLlama.Client", LogLevel.Information); + + await builder.Build().RunAsync(); + return 0; +} +catch (Exception exception) +{ + Console.Error.WriteLine(exception.Message); + Console.Error.WriteLine(); + Console.Error.WriteLine(ClientOptions.Usage); + return 1; +} diff --git a/src/ReverseLlama.Client/Properties/AssemblyInfo.cs b/src/ReverseLlama.Client/Properties/AssemblyInfo.cs new file mode 100644 index 0000000..4458d19 --- /dev/null +++ b/src/ReverseLlama.Client/Properties/AssemblyInfo.cs @@ -0,0 +1,3 @@ +using System.Runtime.CompilerServices; + +[assembly: InternalsVisibleTo("ReverseLlama.Client.Tests")] diff --git a/src/ReverseLlama.Client/ReverseLlama.Client.csproj b/src/ReverseLlama.Client/ReverseLlama.Client.csproj new file mode 100644 index 0000000..27d71c3 --- /dev/null +++ b/src/ReverseLlama.Client/ReverseLlama.Client.csproj @@ -0,0 +1,18 @@ + + + + + + + + + + + + Exe + net10.0 + enable + enable + + + diff --git a/src/ReverseLlama.Client/TunnelClient.cs b/src/ReverseLlama.Client/TunnelClient.cs new file mode 100644 index 0000000..7a72cc1 --- /dev/null +++ b/src/ReverseLlama.Client/TunnelClient.cs @@ -0,0 +1,540 @@ +using System.Collections.Concurrent; +using System.Net.WebSockets; +using System.Text; +using System.Text.Json; +using Microsoft.Extensions.Logging; +using Microsoft.Extensions.Logging.Abstractions; +using ReverseLlama.Protocol; + +namespace ReverseLlama.Client; + +internal sealed class TunnelClient +{ + private static readonly TimeSpan ModelRefreshInterval = TimeSpan.FromSeconds(15); + private static readonly TimeSpan ModelRefreshTimeout = TimeSpan.FromSeconds(10); + private static readonly JsonSerializerOptions JsonOptions = new(JsonSerializerDefaults.Web); + private const string EmbeddingWarmupInput = "ReverseLlama warmup"; + + private readonly ConcurrentDictionary _activeRequests = new(); + private readonly HttpClient _httpClient; + private readonly ClientOptions _options; + private readonly ILogger _logger; + private readonly object _modelSnapshotLock = new(); + private readonly SemaphoreSlim _sendLock = new(1, 1); + private List _lastActiveModels = []; + private List _lastModels = []; + + public TunnelClient(ClientOptions options, ILogger? logger = null) + { + _options = options; + _logger = logger ?? NullLogger.Instance; + _httpClient = new HttpClient + { + Timeout = Timeout.InfiniteTimeSpan + }; + } + + public async Task RunAsync(CancellationToken cancellationToken) + { + while (!cancellationToken.IsCancellationRequested) + { + using var socket = new ClientWebSocket(); + socket.Options.KeepAliveInterval = TimeSpan.FromSeconds(30); + + if (!string.IsNullOrWhiteSpace(_options.Token)) + { + socket.Options.SetRequestHeader(ProtocolConstants.TokenHeader, _options.Token); + } + + socket.Options.SetRequestHeader(ProtocolConstants.ClientIdHeader, _options.ClientId); + + using var connectionCts = CancellationTokenSource.CreateLinkedTokenSource(cancellationToken); + Task? modelRefreshTask = null; + + try + { + _logger.LogInformation("Connecting to {TunnelUri}...", _options.TunnelUri); + await socket.ConnectAsync(_options.TunnelUri, cancellationToken); + _logger.LogInformation("Tunnel connected."); + + modelRefreshTask = RefreshModelsLoopAsync(socket, connectionCts.Token); + await ReceiveLoopAsync(socket, connectionCts.Token); + + if (socket.CloseStatusDescription == ProtocolConstants.ReplacedCloseDescription) + { + _logger.LogWarning("This client was replaced by a newer tunnel client. Exiting."); + return; + } + + _logger.LogInformation("Tunnel closed by server ({Reason}).", socket.CloseStatusDescription ?? "no reason given"); + } + catch (OperationCanceledException) when (cancellationToken.IsCancellationRequested) + { + break; + } + catch (Exception exception) + { + _logger.LogWarning("Tunnel disconnected: {Message}", exception.Message); + } + finally + { + connectionCts.Cancel(); + if (modelRefreshTask is not null) + { + try + { + await modelRefreshTask; + } + catch (OperationCanceledException) + { + } + } + + CancelAllActiveRequests(); + } + + if (!cancellationToken.IsCancellationRequested) + { + _logger.LogInformation("Reconnecting in {Seconds:0.#} seconds...", _options.ReconnectDelay.TotalSeconds); + await Task.Delay(_options.ReconnectDelay, cancellationToken); + } + } + } + + private async Task RefreshModelsLoopAsync(ClientWebSocket socket, CancellationToken cancellationToken) + { + while (socket.State == WebSocketState.Open && !cancellationToken.IsCancellationRequested) + { + try + { + await RefreshModelsOnceAsync(socket, cancellationToken); + } + catch (OperationCanceledException) when (cancellationToken.IsCancellationRequested) + { + break; + } + catch (Exception exception) + { + _logger.LogWarning("Failed to report upstream model list: {Message}", exception.Message); + } + + await Task.Delay(ModelRefreshInterval, cancellationToken); + } + } + + private async Task RefreshModelsOnceAsync(ClientWebSocket socket, CancellationToken cancellationToken) + { + var modelsTask = TryRefreshModelListAsync(GetUpstreamModelsAsync, "listed", cancellationToken); + var activeModelsTask = TryRefreshModelListAsync(GetActiveUpstreamModelsAsync, "active", cancellationToken); + + await Task.WhenAll(modelsTask, activeModelsTask); + + var snapshot = UpdateCachedModelSnapshot(modelsTask.Result, activeModelsTask.Result); + + await SendAsync( + socket, + new TunnelMessage + { + Type = TunnelMessageTypes.ModelSnapshot, + Models = snapshot.Models, + ActiveModels = snapshot.ActiveModels + }, + cancellationToken); + + _logger.LogInformation( + "Reported {ModelCount} listed and {ActiveModelCount} active upstream model(s).", + snapshot.Models.Count, + snapshot.ActiveModels.Count); + } + + private async Task> GetUpstreamModelsAsync(CancellationToken cancellationToken) + { + using var request = new HttpRequestMessage(HttpMethod.Get, new Uri(_options.Upstream, "/api/tags")); + using var response = await _httpClient.SendAsync(request, HttpCompletionOption.ResponseHeadersRead, cancellationToken); + response.EnsureSuccessStatusCode(); + + await using var stream = await response.Content.ReadAsStreamAsync(cancellationToken); + using var document = await JsonDocument.ParseAsync(stream, cancellationToken: cancellationToken); + + return ExtractModelNames(document.RootElement); + } + + private async Task> GetActiveUpstreamModelsAsync(CancellationToken cancellationToken) + { + using var request = new HttpRequestMessage(HttpMethod.Get, new Uri(_options.Upstream, "/api/ps")); + using var response = await _httpClient.SendAsync(request, HttpCompletionOption.ResponseHeadersRead, cancellationToken); + + if (!response.IsSuccessStatusCode) + { + _logger.LogDebug("Ollama /api/ps returned {StatusCode}; active model list will be empty.", response.StatusCode); + return []; + } + + await using var stream = await response.Content.ReadAsStreamAsync(cancellationToken); + using var document = await JsonDocument.ParseAsync(stream, cancellationToken: cancellationToken); + + return ExtractModelNames(document.RootElement); + } + + internal static List ExtractModelNames(JsonElement root) + { + var models = new List(); + + if (root.ValueKind == JsonValueKind.Object + && root.TryGetProperty("models", out var ollamaModels) + && ollamaModels.ValueKind == JsonValueKind.Array) + { + AddModelNames(models, ollamaModels, "name"); + AddModelNames(models, ollamaModels, "model"); + } + + return NormalizeModelNames(models); + } + + private static void AddModelNames(List models, JsonElement array, string propertyName) + { + foreach (var item in array.EnumerateArray()) + { + if (item.ValueKind == JsonValueKind.Object + && item.TryGetProperty(propertyName, out var model) + && model.ValueKind == JsonValueKind.String + && !string.IsNullOrWhiteSpace(model.GetString())) + { + models.Add(model.GetString()!); + } + } + } + + private async Task?> TryRefreshModelListAsync( + Func>> refresh, + string listName, + CancellationToken cancellationToken) + { + try + { + using var refreshCts = CancellationTokenSource.CreateLinkedTokenSource(cancellationToken); + refreshCts.CancelAfter(ModelRefreshTimeout); + + return await refresh(refreshCts.Token); + } + catch (OperationCanceledException) when (cancellationToken.IsCancellationRequested) + { + throw; + } + catch (OperationCanceledException) + { + _logger.LogWarning("Timed out refreshing {ModelListName} upstream model list.", listName); + return null; + } + catch (Exception exception) + { + _logger.LogWarning( + "Failed to refresh {ModelListName} upstream model list: {Message}", + listName, + exception.Message); + return null; + } + } + + private (List Models, List ActiveModels) UpdateCachedModelSnapshot( + List? models, + List? activeModels) + { + lock (_modelSnapshotLock) + { + if (models is not null) + { + _lastModels = models; + } + + if (activeModels is not null) + { + _lastActiveModels = activeModels; + } + + return ( + [.. _lastModels], + [.. _lastActiveModels]); + } + } + + private static List NormalizeModelNames(IEnumerable models) => + models + .Where(model => !string.IsNullOrWhiteSpace(model)) + .Select(model => model.Trim()) + .Distinct(StringComparer.OrdinalIgnoreCase) + .OrderBy(model => model, StringComparer.OrdinalIgnoreCase) + .ToList(); + + private async Task ReceiveLoopAsync(ClientWebSocket socket, CancellationToken cancellationToken) + { + while (socket.State == WebSocketState.Open && !cancellationToken.IsCancellationRequested) + { + var message = await WebSocketMessageTransport.ReceiveAsync(socket, cancellationToken); + if (message is null) + { + break; + } + + await DispatchAsync(socket, message, cancellationToken); + } + } + + private Task DispatchAsync(ClientWebSocket socket, TunnelMessage message, CancellationToken cancellationToken) + { + switch (message.Type) + { + case TunnelMessageTypes.HttpRequest: + StartRequest(socket, message, cancellationToken); + break; + + case TunnelMessageTypes.HttpRequestBody: + if (_activeRequests.TryGetValue(message.RequestId, out var requestWithBody)) + { + requestWithBody.AddBody(message.Body ?? []); + } + break; + + case TunnelMessageTypes.HttpRequestComplete: + if (_activeRequests.TryGetValue(message.RequestId, out var completedRequest)) + { + completedRequest.CompleteBody(); + } + break; + + case TunnelMessageTypes.Cancel: + if (_activeRequests.TryRemove(message.RequestId, out var cancelledRequest)) + { + cancelledRequest.Cancel(); + } + break; + + case TunnelMessageTypes.ModelCommand: + _ = Task.Run(() => RunModelCommandAsync(socket, message, cancellationToken), cancellationToken); + break; + } + + return Task.CompletedTask; + } + + private async Task RunModelCommandAsync(ClientWebSocket socket, TunnelMessage message, CancellationToken cancellationToken) + { + try + { + var response = await ExecuteModelCommandAsync(message, cancellationToken); + if (response.StatusCode is >= 200 and < 300) + { + await RefreshModelsOnceAsync(socket, cancellationToken); + } + + await SendAsync(socket, response, cancellationToken); + } + catch (OperationCanceledException) when (cancellationToken.IsCancellationRequested) + { + } + catch (Exception exception) + { + await SendAsync( + socket, + new TunnelMessage + { + Type = TunnelMessageTypes.ModelCommandResult, + RequestId = message.RequestId, + Error = exception.Message + }, + CancellationToken.None); + } + } + + private async Task ExecuteModelCommandAsync(TunnelMessage message, CancellationToken cancellationToken) + { + if (string.IsNullOrWhiteSpace(message.RequestId)) + { + throw new InvalidOperationException("Model command is missing a request id."); + } + + if (string.IsNullOrWhiteSpace(message.Model)) + { + throw new InvalidOperationException("Model command is missing a model name."); + } + + using var request = BuildModelCommandRequest(message); + using var response = await _httpClient.SendAsync(request, HttpCompletionOption.ResponseHeadersRead, cancellationToken); + var body = await response.Content.ReadAsByteArrayAsync(cancellationToken); + + if (ShouldRetryModelCommandWithEmbedding(message.Command, response, body)) + { + using var embeddingRequest = BuildEmbeddingModelCommandRequest(_options.Upstream, message.Command, message.Model); + using var embeddingResponse = await _httpClient.SendAsync(embeddingRequest, HttpCompletionOption.ResponseHeadersRead, cancellationToken); + var embeddingBody = await embeddingResponse.Content.ReadAsByteArrayAsync(cancellationToken); + + return BuildModelCommandResult(message.RequestId, embeddingResponse, embeddingBody); + } + + return BuildModelCommandResult(message.RequestId, response, body); + } + + private static TunnelMessage BuildModelCommandResult( + string requestId, + HttpResponseMessage response, + byte[] body) + { + return new TunnelMessage + { + Type = TunnelMessageTypes.ModelCommandResult, + RequestId = requestId, + StatusCode = (int)response.StatusCode, + ReasonPhrase = response.ReasonPhrase, + Body = body + }; + } + + private static bool ShouldRetryModelCommandWithEmbedding( + string? command, + HttpResponseMessage response, + byte[] body) + { + var normalizedCommand = NormalizeModelCommand(command); + if (normalizedCommand is not ("load" or "unload") + || response.IsSuccessStatusCode) + { + return false; + } + + if (response.StatusCode == System.Net.HttpStatusCode.BadRequest) + { + return true; + } + + var responseText = body.Length > 0 + ? Encoding.UTF8.GetString(body) + : ""; + + return responseText.Contains("does not support generate", StringComparison.OrdinalIgnoreCase); + } + + private HttpRequestMessage BuildModelCommandRequest(TunnelMessage message) + { + return BuildModelCommandRequest(_options.Upstream, message.Command, message.Model); + } + + internal static HttpRequestMessage BuildModelCommandRequest(Uri upstream, string? command, string? modelName) + { + if (string.IsNullOrWhiteSpace(modelName)) + { + throw new InvalidOperationException("Model command is missing a model name."); + } + + var model = modelName.Trim(); + var normalizedCommand = NormalizeModelCommand(command); + + return normalizedCommand switch + { + "pull" => new HttpRequestMessage(HttpMethod.Post, new Uri(upstream, "/api/pull")) + { + Content = JsonContent(new { model, stream = false }) + }, + "delete" => new HttpRequestMessage(HttpMethod.Delete, new Uri(upstream, "/api/delete")) + { + Content = JsonContent(new { model }) + }, + "load" => new HttpRequestMessage(HttpMethod.Post, new Uri(upstream, "/api/generate")) + { + Content = JsonContent(new { model, stream = false, keep_alive = -1 }) + }, + "unload" => new HttpRequestMessage(HttpMethod.Post, new Uri(upstream, "/api/generate")) + { + Content = JsonContent(new { model, stream = false, keep_alive = 0 }) + }, + "show" => new HttpRequestMessage(HttpMethod.Post, new Uri(upstream, "/api/show")) + { + Content = JsonContent(new { model }) + }, + _ => throw new InvalidOperationException($"Unsupported model command '{command}'.") + }; + } + + internal static HttpRequestMessage BuildEmbeddingModelCommandRequest(Uri upstream, string? command, string? modelName) + { + if (string.IsNullOrWhiteSpace(modelName)) + { + throw new InvalidOperationException("Model command is missing a model name."); + } + + var model = modelName.Trim(); + var normalizedCommand = NormalizeModelCommand(command); + + return normalizedCommand switch + { + "load" => new HttpRequestMessage(HttpMethod.Post, new Uri(upstream, "/api/embed")) + { + Content = JsonContent(new { model, input = EmbeddingWarmupInput, keep_alive = -1 }) + }, + "unload" => new HttpRequestMessage(HttpMethod.Post, new Uri(upstream, "/api/embed")) + { + Content = JsonContent(new { model, input = EmbeddingWarmupInput, keep_alive = 0 }) + }, + _ => throw new InvalidOperationException($"Unsupported embedding model command '{command}'.") + }; + } + + private static string NormalizeModelCommand(string? command) => + (command ?? "").Trim().ToLowerInvariant(); + + private static StringContent JsonContent(T value) => + new(JsonSerializer.Serialize(value, JsonOptions), Encoding.UTF8, "application/json"); + + private void StartRequest(ClientWebSocket socket, TunnelMessage message, CancellationToken cancellationToken) + { + var request = new UpstreamRequest( + _options, + _httpClient, + message, + (response, token) => SendAsync(socket, response, token), + requestId => _activeRequests.TryRemove(requestId, out _), + cancellationToken); + + if (!_activeRequests.TryAdd(message.RequestId, request)) + { + _ = SendAsync( + socket, + new TunnelMessage + { + Type = TunnelMessageTypes.Error, + RequestId = message.RequestId, + Error = "Duplicate request id." + }, + cancellationToken); + return; + } + + _ = Task.Run(request.RunAsync, cancellationToken); + } + + private async Task SendAsync(ClientWebSocket socket, TunnelMessage message, CancellationToken cancellationToken) + { + await _sendLock.WaitAsync(cancellationToken); + try + { + if (socket.State == WebSocketState.Open) + { + await WebSocketMessageTransport.SendAsync(socket, message, cancellationToken); + } + } + finally + { + _sendLock.Release(); + } + } + + private void CancelAllActiveRequests() + { + foreach (var pair in _activeRequests.ToArray()) + { + if (_activeRequests.TryRemove(pair.Key, out var request)) + { + request.Cancel(); + } + } + } +} diff --git a/src/ReverseLlama.Client/TunnelWorker.cs b/src/ReverseLlama.Client/TunnelWorker.cs new file mode 100644 index 0000000..9e1b156 --- /dev/null +++ b/src/ReverseLlama.Client/TunnelWorker.cs @@ -0,0 +1,30 @@ +using Microsoft.Extensions.Hosting; + +namespace ReverseLlama.Client; + +internal sealed class TunnelWorker : BackgroundService +{ + private readonly TunnelClient _client; + private readonly IHostApplicationLifetime _lifetime; + + public TunnelWorker(TunnelClient client, IHostApplicationLifetime lifetime) + { + _client = client; + _lifetime = lifetime; + } + + protected override async Task ExecuteAsync(CancellationToken stoppingToken) + { + try + { + await _client.RunAsync(stoppingToken); + } + finally + { + // RunAsync only returns when cancelled or replaced by a newer client. + // Stop gracefully (exit 0) so service recovery does not restart us + // into a reconnect fight with the replacement. + _lifetime.StopApplication(); + } + } +} diff --git a/src/ReverseLlama.Client/UpstreamRequest.cs b/src/ReverseLlama.Client/UpstreamRequest.cs new file mode 100644 index 0000000..3d4358e --- /dev/null +++ b/src/ReverseLlama.Client/UpstreamRequest.cs @@ -0,0 +1,245 @@ +using System.Net.Http.Headers; +using System.Threading.Channels; +using ReverseLlama.Protocol; + +namespace ReverseLlama.Client; + +internal sealed class UpstreamRequest +{ + private static readonly HashSet HeadersToSkip = new(StringComparer.OrdinalIgnoreCase) + { + "Connection", + "Content-Length", + "Expect", + "Host", + "Keep-Alive", + "Proxy-Authenticate", + "Proxy-Authorization", + "TE", + "Trailer", + "Transfer-Encoding", + "Upgrade", + ProtocolConstants.TokenHeader + }; + + private readonly CancellationTokenSource _cancellationTokenSource; + private readonly Channel _requestBody = Channel.CreateUnbounded( + new UnboundedChannelOptions + { + SingleReader = true, + SingleWriter = false + }); + + private readonly HttpClient _httpClient; + private readonly TunnelMessage _initialMessage; + private readonly Action _onComplete; + private readonly ClientOptions _options; + private readonly Func _sendAsync; + + public UpstreamRequest( + ClientOptions options, + HttpClient httpClient, + TunnelMessage initialMessage, + Func sendAsync, + Action onComplete, + CancellationToken cancellationToken) + { + _options = options; + _httpClient = httpClient; + _initialMessage = initialMessage; + _sendAsync = sendAsync; + _onComplete = onComplete; + _cancellationTokenSource = CancellationTokenSource.CreateLinkedTokenSource(cancellationToken); + + if (!initialMessage.HasBody) + { + _requestBody.Writer.TryComplete(); + } + } + + public void AddBody(byte[] body) + { + if (body.Length > 0) + { + _requestBody.Writer.TryWrite(body); + } + } + + public void CompleteBody() => + _requestBody.Writer.TryComplete(); + + public void Cancel() + { + _requestBody.Writer.TryComplete(); + _cancellationTokenSource.Cancel(); + } + + public async Task RunAsync() + { + try + { + using var request = BuildHttpRequest(); + using var response = await _httpClient.SendAsync( + request, + HttpCompletionOption.ResponseHeadersRead, + _cancellationTokenSource.Token); + + await SendResponseHeadersAsync(response); + await SendResponseBodyAsync(response); + + await _sendAsync( + new TunnelMessage + { + Type = TunnelMessageTypes.HttpResponseComplete, + RequestId = _initialMessage.RequestId + }, + _cancellationTokenSource.Token); + } + catch (OperationCanceledException) when (_cancellationTokenSource.IsCancellationRequested) + { + } + catch (Exception exception) + { + await SendErrorAsync(exception); + } + finally + { + _requestBody.Writer.TryComplete(); + _onComplete(_initialMessage.RequestId); + _cancellationTokenSource.Dispose(); + } + } + + private HttpRequestMessage BuildHttpRequest() + { + var method = new HttpMethod(_initialMessage.Method ?? HttpMethod.Get.Method); + var request = new HttpRequestMessage(method, BuildUpstreamUri(_options.Upstream, _initialMessage.PathAndQuery)); + + if (_initialMessage.HasBody) + { + request.Content = new ChannelHttpContent(_requestBody.Reader, _cancellationTokenSource.Token); + } + + foreach (var header in _initialMessage.Headers) + { + if (HeadersToSkip.Contains(header.Name)) + { + continue; + } + + if (!request.Headers.TryAddWithoutValidation(header.Name, header.Value) + && request.Content is not null) + { + request.Content.Headers.TryAddWithoutValidation(header.Name, header.Value); + } + } + + return request; + } + + internal static Uri BuildUpstreamUri(Uri upstream, string? pathAndQuery) + { + if (string.IsNullOrWhiteSpace(pathAndQuery)) + { + return upstream; + } + + if (!IsOriginPathAndQuery(pathAndQuery)) + { + throw new InvalidOperationException("Tunnel request path must be an origin-form path."); + } + + var uri = new Uri(upstream, pathAndQuery); + if (!HasSameOrigin(uri, upstream)) + { + throw new InvalidOperationException("Tunnel request path resolved outside the configured upstream origin."); + } + + return uri; + } + + private static bool IsOriginPathAndQuery(string pathAndQuery) => + pathAndQuery.StartsWith("/", StringComparison.Ordinal) + && !pathAndQuery.StartsWith("//", StringComparison.Ordinal) + && !pathAndQuery.Contains('\\'); + + private static bool HasSameOrigin(Uri uri, Uri upstream) => + string.Equals(uri.Scheme, upstream.Scheme, StringComparison.OrdinalIgnoreCase) + && string.Equals(uri.IdnHost, upstream.IdnHost, StringComparison.OrdinalIgnoreCase) + && uri.Port == upstream.Port; + + private async Task SendResponseHeadersAsync(HttpResponseMessage response) + { + await _sendAsync( + new TunnelMessage + { + Type = TunnelMessageTypes.HttpResponseHeaders, + RequestId = _initialMessage.RequestId, + StatusCode = (int)response.StatusCode, + ReasonPhrase = response.ReasonPhrase, + Headers = CollectResponseHeaders(response) + }, + _cancellationTokenSource.Token); + } + + private async Task SendResponseBodyAsync(HttpResponseMessage response) + { + await using var stream = await response.Content.ReadAsStreamAsync(_cancellationTokenSource.Token); + var buffer = new byte[_options.ChunkSize]; + + while (true) + { + var bytesRead = await stream.ReadAsync(buffer, _cancellationTokenSource.Token); + if (bytesRead == 0) + { + break; + } + + await _sendAsync( + new TunnelMessage + { + Type = TunnelMessageTypes.HttpResponseBody, + RequestId = _initialMessage.RequestId, + Body = buffer.AsSpan(0, bytesRead).ToArray() + }, + _cancellationTokenSource.Token); + } + } + + private async Task SendErrorAsync(Exception exception) + { + try + { + await _sendAsync( + new TunnelMessage + { + Type = TunnelMessageTypes.Error, + RequestId = _initialMessage.RequestId, + Error = exception.Message + }, + CancellationToken.None); + } + catch + { + } + } + + private static List CollectResponseHeaders(HttpResponseMessage response) + { + var headers = new List(); + AddHeaders(headers, response.Headers); + AddHeaders(headers, response.Content.Headers); + return headers; + } + + private static void AddHeaders(List target, HttpHeaders headers) + { + foreach (var header in headers) + { + foreach (var value in header.Value) + { + target.Add(new HeaderPair(header.Key, value)); + } + } + } +} diff --git a/src/ReverseLlama.Protocol/HeaderPair.cs b/src/ReverseLlama.Protocol/HeaderPair.cs new file mode 100644 index 0000000..4a06fac --- /dev/null +++ b/src/ReverseLlama.Protocol/HeaderPair.cs @@ -0,0 +1,18 @@ +namespace ReverseLlama.Protocol; + +public sealed class HeaderPair +{ + public HeaderPair() + { + } + + public HeaderPair(string name, string value) + { + Name = name; + Value = value; + } + + public string Name { get; set; } = ""; + + public string Value { get; set; } = ""; +} diff --git a/src/ReverseLlama.Protocol/ProtocolConstants.cs b/src/ReverseLlama.Protocol/ProtocolConstants.cs new file mode 100644 index 0000000..36235d9 --- /dev/null +++ b/src/ReverseLlama.Protocol/ProtocolConstants.cs @@ -0,0 +1,10 @@ +namespace ReverseLlama.Protocol; + +public static class ProtocolConstants +{ + public const string DefaultStatusPath = "/_reverse-llama/status"; + public const string DefaultTunnelPath = "/_reverse-llama/tunnel"; + public const string TokenHeader = "X-Reverse-Llama-Token"; + public const string ClientIdHeader = "X-Reverse-Llama-Client-Id"; + public const string ReplacedCloseDescription = "reverse-llama-replaced"; +} diff --git a/src/ReverseLlama.Protocol/ReverseLlama.Protocol.csproj b/src/ReverseLlama.Protocol/ReverseLlama.Protocol.csproj new file mode 100644 index 0000000..b760144 --- /dev/null +++ b/src/ReverseLlama.Protocol/ReverseLlama.Protocol.csproj @@ -0,0 +1,9 @@ + + + + net10.0 + enable + enable + + + diff --git a/src/ReverseLlama.Protocol/TunnelMessage.cs b/src/ReverseLlama.Protocol/TunnelMessage.cs new file mode 100644 index 0000000..1828527 --- /dev/null +++ b/src/ReverseLlama.Protocol/TunnelMessage.cs @@ -0,0 +1,34 @@ +namespace ReverseLlama.Protocol; + +public sealed class TunnelMessage +{ + public string Type { get; set; } = ""; + + public string RequestId { get; set; } = ""; + + public string? Method { get; set; } + + public string? PathAndQuery { get; set; } + + public bool HasBody { get; set; } + + public List Headers { get; set; } = []; + + public int? StatusCode { get; set; } + + public string? ReasonPhrase { get; set; } + + public byte[]? Body { get; set; } + + public List Models { get; set; } = []; + + public List ActiveModels { get; set; } = []; + + public string? Command { get; set; } + + public string? Model { get; set; } + + public string? PayloadJson { get; set; } + + public string? Error { get; set; } +} diff --git a/src/ReverseLlama.Protocol/TunnelMessageTypes.cs b/src/ReverseLlama.Protocol/TunnelMessageTypes.cs new file mode 100644 index 0000000..375c2e1 --- /dev/null +++ b/src/ReverseLlama.Protocol/TunnelMessageTypes.cs @@ -0,0 +1,16 @@ +namespace ReverseLlama.Protocol; + +public static class TunnelMessageTypes +{ + public const string HttpRequest = "http.request"; + public const string HttpRequestBody = "http.request.body"; + public const string HttpRequestComplete = "http.request.complete"; + public const string HttpResponseHeaders = "http.response.headers"; + public const string HttpResponseBody = "http.response.body"; + public const string HttpResponseComplete = "http.response.complete"; + public const string ModelSnapshot = "models.snapshot"; + public const string ModelCommand = "models.command"; + public const string ModelCommandResult = "models.command.result"; + public const string Cancel = "cancel"; + public const string Error = "error"; +} diff --git a/src/ReverseLlama.Protocol/WebSocketMessageTransport.cs b/src/ReverseLlama.Protocol/WebSocketMessageTransport.cs new file mode 100644 index 0000000..2a24a47 --- /dev/null +++ b/src/ReverseLlama.Protocol/WebSocketMessageTransport.cs @@ -0,0 +1,57 @@ +using System.Net.WebSockets; +using System.Text.Json; + +namespace ReverseLlama.Protocol; + +public static class WebSocketMessageTransport +{ + private const int BufferSize = 64 * 1024; + private const int MaxMessageSize = 128 * 1024 * 1024; + + private static readonly JsonSerializerOptions JsonOptions = new(JsonSerializerDefaults.Web); + + public static async Task SendAsync(WebSocket socket, TunnelMessage message, CancellationToken cancellationToken) + { + var payload = JsonSerializer.SerializeToUtf8Bytes(message, JsonOptions); + await socket.SendAsync(payload, WebSocketMessageType.Text, true, cancellationToken); + } + + public static async Task ReceiveAsync(WebSocket socket, CancellationToken cancellationToken) + { + var buffer = new byte[BufferSize]; + using var payload = new MemoryStream(); + + while (true) + { + var result = await socket.ReceiveAsync(buffer, cancellationToken); + if (result.MessageType == WebSocketMessageType.Close) + { + return null; + } + + if (result.MessageType != WebSocketMessageType.Text) + { + throw new InvalidOperationException("Only text WebSocket messages are supported."); + } + + if (result.Count > 0) + { + payload.Write(buffer.AsSpan(0, result.Count)); + } + + if (payload.Length > MaxMessageSize) + { + throw new InvalidOperationException($"Tunnel message exceeded {MaxMessageSize} bytes."); + } + + if (result.EndOfMessage) + { + break; + } + } + + payload.Position = 0; + var message = await JsonSerializer.DeserializeAsync(payload, JsonOptions, cancellationToken); + return message ?? throw new InvalidOperationException("Received an empty tunnel message."); + } +} diff --git a/src/ReverseLlama.Server/AdminEndpoints.cs b/src/ReverseLlama.Server/AdminEndpoints.cs new file mode 100644 index 0000000..c02da4e --- /dev/null +++ b/src/ReverseLlama.Server/AdminEndpoints.cs @@ -0,0 +1,490 @@ +using System.Security.Claims; +using System.Text; +using System.Text.Json; +using Microsoft.AspNetCore.Authentication; +using Microsoft.AspNetCore.Authentication.Cookies; +using Microsoft.AspNetCore.Authentication.OpenIdConnect; +using Microsoft.AspNetCore.StaticFiles; + +namespace ReverseLlama.Server; + +internal static class AdminEndpoints +{ + private static readonly FileExtensionContentTypeProvider ContentTypes = new(); + + public static void MapAdminEndpoints(this WebApplication app, ServerSettings settings) + { + if (settings.Keycloak.IsConfigured) + { + app.MapGet("/admin/login", (string? returnUrl) => + Results.Challenge( + new AuthenticationProperties { RedirectUri = NormalizeLocalReturnUrl(returnUrl) }, + [OpenIdConnectDefaults.AuthenticationScheme])) + .AllowAnonymous(); + + app.MapPost("/admin/logout", () => + Results.SignOut( + new AuthenticationProperties { RedirectUri = "/admin" }, + [CookieAuthenticationDefaults.AuthenticationScheme, OpenIdConnectDefaults.AuthenticationScheme])) + .RequireAuthorization(); + } + + app.MapGet("/admin/auth-error", () => + Results.Text( + "Login failed while processing the Keycloak callback. The exception was written to ELMAH.", + "text/plain")) + .AllowAnonymous(); + + var api = app.MapGroup("/api/admin"); + + if (settings.Keycloak.IsConfigured) + { + api.RequireAuthorization(); + } + + api.MapGet("/summary", (HttpContext context, TunnelHub hub, ManagementStore store) => + Results.Json(BuildSummary(context.User, hub, store, settings))); + + api.MapGet("/me", (HttpContext context, ManagementStore store) => + Results.Json(new + { + authenticated = context.User.Identity?.IsAuthenticated ?? false, + name = GetUserName(context.User), + keycloakConfigured = settings.Keycloak.IsConfigured, + management = new + { + available = store.IsAvailable, + databasePath = store.DatabasePath, + lastError = store.LastError + } + })); + + api.MapPost("/clients/{clientId}/disable", (string clientId, DisableClientRequest request, ManagementStore store) => + { + try + { + var manual = string.Equals(request.Mode, "manual", StringComparison.OrdinalIgnoreCase); + TimeSpan? duration = manual + ? null + : TimeSpan.FromMinutes(Math.Clamp(request.DurationMinutes ?? 60, 1, 24 * 60)); + + store.DisableClient(clientId, duration, manual, request.Reason); + return Results.Ok(new { clientId, disabled = true }); + } + catch (Exception exception) + { + return Results.Problem(exception.Message, statusCode: StatusCodes.Status400BadRequest); + } + }); + + api.MapPost("/clients/{clientId}/enable", (string clientId, ManagementStore store) => + { + try + { + store.EnableClient(clientId); + return Results.Ok(new { clientId, disabled = false }); + } + catch (Exception exception) + { + return Results.Problem(exception.Message, statusCode: StatusCodes.Status400BadRequest); + } + }); + + api.MapGet("/models/detail", async ( + HttpContext context, + string model, + string? clientId, + TunnelHub hub, + ManagementStore store) => + { + if (string.IsNullOrWhiteSpace(model)) + { + return Results.BadRequest(new { error = "Model is required." }); + } + + var modelSummary = BuildModelSummaries(hub, store) + .FirstOrDefault(item => item.Name.Equals(model, StringComparison.OrdinalIgnoreCase)); + var selectedClientId = ResolveModelClientId(hub, modelSummary, model, clientId); + object? show = null; + + if (!string.IsNullOrWhiteSpace(selectedClientId)) + { + var connection = hub.Get(selectedClientId); + if (connection is not null) + { + show = await SendModelCommandForApiAsync( + connection, + "show", + model, + TimeSpan.FromSeconds(60), + context.RequestAborted); + } + } + + return Results.Json(new + { + model, + listedClients = modelSummary?.ListedClients ?? [], + activeClients = modelSummary?.ActiveClients ?? [], + metrics = modelSummary?.Metrics ?? EmptyModelMetrics(), + selectedClientId, + show + }); + }); + + api.MapPost("/models/actions", async ( + HttpContext context, + ModelActionRequest request, + TunnelHub hub) => + { + if (string.IsNullOrWhiteSpace(request.ClientId) + || string.IsNullOrWhiteSpace(request.Model) + || string.IsNullOrWhiteSpace(request.Action)) + { + return Results.BadRequest(new { error = "Client id, model, and action are required." }); + } + + if (!TryMapModelAction(request.Action, out var command, out var timeout)) + { + return Results.BadRequest(new { error = $"Unsupported action '{request.Action}'." }); + } + + var connection = hub.Get(request.ClientId); + if (connection is null) + { + return Results.NotFound(new { error = $"Client '{request.ClientId}' is not connected." }); + } + + var result = await SendModelCommandForApiAsync( + connection, + command, + request.Model, + timeout, + context.RequestAborted); + + return Results.Json(result); + }); + + api.MapGet("/api-keys", (ManagementStore store) => + Results.Json(store.ListApiKeys())); + + api.MapPost("/api-keys", (CreateApiKeyRequest request, ManagementStore store) => + { + try + { + return Results.Json(store.CreateApiKey(request.Name)); + } + catch (Exception exception) + { + return Results.Problem(exception.Message, statusCode: StatusCodes.Status400BadRequest); + } + }); + + api.MapDelete("/api-keys/{id}", (string id, ManagementStore store) => + store.DeleteApiKey(id) + ? Results.NoContent() + : Results.NotFound(new { error = $"API key '{id}' was not found." })); + + var adminHome = app.MapGet("/admin", (IWebHostEnvironment environment) => + ServeAdminAsset(environment, null)); + var adminAssets = app.MapGet("/admin/{**assetPath}", (IWebHostEnvironment environment, string? assetPath) => + ServeAdminAsset(environment, assetPath)); + + if (settings.Keycloak.IsConfigured) + { + adminHome.RequireAuthorization(); + adminAssets.RequireAuthorization(); + } + } + + private static object BuildSummary( + ClaimsPrincipal user, + TunnelHub hub, + ManagementStore store, + ServerSettings settings) => + new + { + generatedAtUtc = DateTimeOffset.UtcNow, + user = new + { + name = GetUserName(user), + authenticated = user.Identity?.IsAuthenticated ?? false + }, + auth = new + { + keycloakConfigured = settings.Keycloak.IsConfigured, + sharedTokenConfigured = !string.IsNullOrWhiteSpace(settings.Token), + apiKeysConfigured = store.HasApiKeys + }, + management = new + { + available = store.IsAvailable, + databasePath = store.DatabasePath, + lastError = store.LastError + }, + clients = BuildClientSummaries(hub, store), + models = BuildModelSummaries(hub, store), + apiKeys = store.ListApiKeys() + }; + + private static IReadOnlyList BuildClientSummaries(TunnelHub hub, ManagementStore store) + { + var connected = hub.ClientSnapshots.ToDictionary(client => client.Id, StringComparer.OrdinalIgnoreCase); + var controls = store.ListClientControls(); + var stats = store.GetClientRequestStats(); + var clientIds = connected.Keys + .Concat(controls.Keys) + .Concat(stats.Keys) + .Distinct(StringComparer.OrdinalIgnoreCase) + .OrderBy(clientId => clientId, StringComparer.OrdinalIgnoreCase); + var result = new List(); + + foreach (var clientId in clientIds) + { + connected.TryGetValue(clientId, out var snapshot); + controls.TryGetValue(clientId, out var access); + stats.TryGetValue(clientId, out var requestStats); + access ??= ClientAccess.Enabled; + + result.Add(new ClientSummary( + clientId, + snapshot is not null, + snapshot?.PendingRequests ?? 0, + snapshot?.Models ?? [], + snapshot?.ActiveModels ?? [], + snapshot?.ModelsUpdatedAt, + access.IsDisabled, + access.DisabledUntilUtc, + access.DisabledManually, + access.DisabledReason, + requestStats ?? new ClientRequestStats(0, 0, 0))); + } + + return result; + } + + private static IReadOnlyList BuildModelSummaries(TunnelHub hub, ManagementStore store) + { + var listedClients = new Dictionary>(StringComparer.OrdinalIgnoreCase); + var activeClients = new Dictionary>(StringComparer.OrdinalIgnoreCase); + + foreach (var client in hub.ClientSnapshots) + { + AddModelClients(listedClients, client.Models, client.Id); + AddModelClients(activeClients, client.ActiveModels, client.Id); + } + + var metrics = store.GetModelUsageStats(); + var modelNames = listedClients.Keys + .Concat(activeClients.Keys) + .Concat(metrics.Keys) + .Distinct(StringComparer.OrdinalIgnoreCase) + .OrderBy(model => model, StringComparer.OrdinalIgnoreCase); + var result = new List(); + + foreach (var model in modelNames) + { + metrics.TryGetValue(model, out var modelMetrics); + + result.Add(new ModelSummary( + model, + listedClients.TryGetValue(model, out var listed) ? listed.ToArray() : [], + activeClients.TryGetValue(model, out var active) ? active.ToArray() : [], + modelMetrics ?? EmptyModelMetrics())); + } + + return result; + } + + private static void AddModelClients( + Dictionary> target, + IEnumerable models, + string clientId) + { + foreach (var model in models) + { + if (!target.TryGetValue(model, out var clients)) + { + clients = new SortedSet(StringComparer.OrdinalIgnoreCase); + target[model] = clients; + } + + clients.Add(clientId); + } + } + + private static string? ResolveModelClientId( + TunnelHub hub, + ModelSummary? modelSummary, + string model, + string? requestedClientId) + { + if (!string.IsNullOrWhiteSpace(requestedClientId) + && hub.Get(requestedClientId) is not null) + { + return requestedClientId; + } + + return modelSummary?.ActiveClients.FirstOrDefault(clientId => hub.Get(clientId) is not null) + ?? modelSummary?.ListedClients.FirstOrDefault(clientId => hub.Get(clientId) is not null) + ?? hub.SelectBest(model)?.ClientId; + } + + private static ModelUsageStats EmptyModelMetrics() => + new(0, 0, 0, 0, 0); + + private static async Task SendModelCommandForApiAsync( + TunnelConnection connection, + string command, + string model, + TimeSpan timeout, + CancellationToken cancellationToken) + { + try + { + var response = await connection.SendModelCommandAsync( + command, + model, + payloadJson: null, + timeout, + cancellationToken); + var body = response.Body is { Length: > 0 } + ? Encoding.UTF8.GetString(response.Body) + : ""; + + return new + { + ok = response.StatusCode is >= 200 and < 300, + statusCode = response.StatusCode, + reasonPhrase = response.ReasonPhrase, + body = ParseJsonOrText(body) + }; + } + catch (OperationCanceledException) + { + return new + { + ok = false, + statusCode = StatusCodes.Status504GatewayTimeout, + reasonPhrase = "Timed out", + body = "The model command timed out." + }; + } + catch (Exception exception) + { + return new + { + ok = false, + statusCode = StatusCodes.Status502BadGateway, + reasonPhrase = "Command failed", + body = exception.Message + }; + } + } + + private static object? ParseJsonOrText(string body) + { + if (string.IsNullOrWhiteSpace(body)) + { + return null; + } + + try + { + using var document = JsonDocument.Parse(body); + return document.RootElement.Clone(); + } + catch (JsonException) + { + return body.Length <= 100_000 ? body : body[..100_000]; + } + } + + private static bool TryMapModelAction(string action, out string command, out TimeSpan timeout) + { + command = action.Trim().ToLowerInvariant() switch + { + "add" or "pull" => "pull", + "remove" or "delete" => "delete", + "load" => "load", + "unload" => "unload", + _ => "" + }; + + timeout = command == "pull" ? TimeSpan.FromMinutes(30) : TimeSpan.FromMinutes(2); + return command.Length > 0; + } + + private static IResult ServeAdminAsset(IWebHostEnvironment environment, string? assetPath) + { + var path = string.IsNullOrWhiteSpace(assetPath) ? "index.html" : assetPath; + + if (path.Contains("..", StringComparison.Ordinal) + || path.Contains('\\')) + { + return Results.BadRequest(); + } + + var file = environment.WebRootFileProvider.GetFileInfo($"admin/{path}"); + if (!file.Exists && !Path.HasExtension(path)) + { + file = environment.WebRootFileProvider.GetFileInfo("admin/index.html"); + } + + if (!file.Exists) + { + return Results.NotFound(); + } + + ContentTypes.TryGetContentType(file.Name, out var contentType); + return Results.Stream(file.CreateReadStream(), contentType ?? "application/octet-stream"); + } + + private static string NormalizeLocalReturnUrl(string? returnUrl) + { + if (string.IsNullOrWhiteSpace(returnUrl) + || !returnUrl.StartsWith("/", StringComparison.Ordinal) + || returnUrl.StartsWith("//", StringComparison.Ordinal)) + { + return "/admin"; + } + + return returnUrl; + } + + private static string? GetUserName(ClaimsPrincipal user) => + user.FindFirst("preferred_username")?.Value + ?? user.FindFirst(ClaimTypes.Name)?.Value + ?? user.Identity?.Name; +} + +internal sealed record DisableClientRequest( + string? Mode, + int? DurationMinutes, + string? Reason); + +internal sealed record ModelActionRequest( + string ClientId, + string Model, + string Action); + +internal sealed record CreateApiKeyRequest(string? Name); + +internal sealed record ClientSummary( + string Id, + bool Connected, + int PendingRequests, + IReadOnlyList Models, + IReadOnlyList ActiveModels, + DateTimeOffset? ModelsUpdatedAt, + bool Disabled, + DateTimeOffset? DisabledUntilUtc, + bool DisabledManually, + string? DisabledReason, + ClientRequestStats RequestStats); + +internal sealed record ModelSummary( + string Name, + IReadOnlyList ListedClients, + IReadOnlyList ActiveClients, + ModelUsageStats Metrics); diff --git a/src/ReverseLlama.Server/EmbeddingCache.cs b/src/ReverseLlama.Server/EmbeddingCache.cs new file mode 100644 index 0000000..ef997db --- /dev/null +++ b/src/ReverseLlama.Server/EmbeddingCache.cs @@ -0,0 +1,566 @@ +using System.Collections.Concurrent; +using System.Text.Json; +using Microsoft.AspNetCore.Http.Features; +using Microsoft.Data.Sqlite; +using ReverseLlama.Protocol; + +namespace ReverseLlama.Server; + +internal sealed class EmbeddingCache +{ + private const string JsonContentType = "application/json; charset=utf-8"; + + private readonly ConcurrentDictionary _entries = new(); + private string _connectionString = ""; + private string _databasePath = ""; + private bool _isAvailable; + private string? _lastError; + private readonly ILogger _logger; + private readonly SemaphoreSlim _storeLock = new(1, 1); + + public EmbeddingCache(ServerSettings settings, ILogger logger) + { + _logger = logger; + + try + { + _databasePath = ResolveDatabasePath(settings.EmbeddingCachePath); + _connectionString = new SqliteConnectionStringBuilder + { + DataSource = _databasePath, + Mode = SqliteOpenMode.ReadWriteCreate, + Pooling = true + }.ToString(); + + Initialize(); + _isAvailable = true; + } + catch (Exception exception) + { + _lastError = exception.Message; + _logger.LogError( + exception, + "Embedding cache is disabled because SQLite could not be initialized at {DatabasePath}.", + string.IsNullOrWhiteSpace(_databasePath) ? settings.EmbeddingCachePath : _databasePath); + } + } + + public int Count => _entries.Count; + + public string DatabasePath => _databasePath; + + public bool IsAvailable => _isAvailable; + + public string? LastError => _lastError; + + public async Task TryReadRequestAsync(HttpRequest request, PathString path) + { + if (!HttpMethods.IsPost(request.Method) + || !TryGetEndpointKind(path, out var kind) + || !CanHaveBody(request)) + { + return null; + } + + request.EnableBuffering(); + + try + { + using var document = await JsonDocument.ParseAsync( + request.Body, + cancellationToken: request.HttpContext.RequestAborted); + + var root = document.RootElement; + if (root.ValueKind != JsonValueKind.Object + || !TryReadRequiredString(root, "model", out var model) + || !TryReadInputTexts(root, kind, out var texts)) + { + return null; + } + + return new EmbeddingCacheRequest(kind, model, texts); + } + catch (JsonException) + { + return null; + } + finally + { + if (request.Body.CanSeek) + { + request.Body.Position = 0; + } + } + } + + public async Task TryWriteCachedResponseAsync(HttpContext context, EmbeddingCacheRequest request) + { + var embeddings = new List(request.Texts.Count); + + foreach (var text in request.Texts) + { + if (!_entries.TryGetValue(new EmbeddingCacheKey(request.Model, text), out var embedding)) + { + return false; + } + + embeddings.Add(embedding); + } + + byte[] body; + try + { + body = BuildResponseBody(request, embeddings); + } + catch (JsonException exception) + { + _logger.LogWarning(exception, "Ignoring invalid cached embedding JSON for model {Model}.", request.Model); + return false; + } + + context.Response.StatusCode = StatusCodes.Status200OK; + context.Response.ContentType = JsonContentType; + context.Response.ContentLength = body.Length; + context.Response.Headers["X-Reverse-Llama-Embedding-Cache"] = "hit"; + + await context.Response.Body.WriteAsync(body, context.RequestAborted); + return true; + } + + public async Task StoreResponseAsync( + EmbeddingCacheRequest request, + TunnelMessage responseHeaders, + byte[] body, + CancellationToken cancellationToken) + { + if (!_isAvailable + || responseHeaders.StatusCode is not >= 200 or >= 300 + || HasContentEncoding(responseHeaders) + || body.Length == 0) + { + return; + } + + List embeddings; + try + { + embeddings = ExtractEmbeddings(request, body); + } + catch (JsonException exception) + { + _logger.LogDebug(exception, "Embedding response for model {Model} was not cacheable JSON.", request.Model); + return; + } + + if (embeddings.Count != request.Texts.Count) + { + _logger.LogDebug( + "Embedding response for model {Model} returned {EmbeddingCount} vector(s) for {TextCount} text(s); skipping cache store.", + request.Model, + embeddings.Count, + request.Texts.Count); + return; + } + + await _storeLock.WaitAsync(cancellationToken); + try + { + using var connection = OpenConnection(); + using var transaction = connection.BeginTransaction(); + + using var command = connection.CreateCommand(); + command.Transaction = transaction; + command.CommandText = """ + INSERT INTO embedding_cache (model, text, embedding_json, created_at_utc, updated_at_utc) + VALUES ($model, $text, $embedding_json, $now, $now) + ON CONFLICT(model, text) DO UPDATE SET + embedding_json = excluded.embedding_json, + updated_at_utc = excluded.updated_at_utc + """; + + var modelParameter = command.Parameters.Add("$model", SqliteType.Text); + var textParameter = command.Parameters.Add("$text", SqliteType.Text); + var embeddingParameter = command.Parameters.Add("$embedding_json", SqliteType.Text); + var nowParameter = command.Parameters.Add("$now", SqliteType.Text); + + var now = DateTimeOffset.UtcNow.ToString("O"); + var stored = new List<(EmbeddingCacheKey Key, CachedEmbedding Embedding)>(embeddings.Count); + + for (var index = 0; index < request.Texts.Count; index++) + { + var key = new EmbeddingCacheKey(request.Model, request.Texts[index]); + var embedding = embeddings[index] with { UpdatedAtUtc = now }; + + modelParameter.Value = key.Model; + textParameter.Value = key.Text; + embeddingParameter.Value = embedding.EmbeddingJson; + nowParameter.Value = now; + + command.ExecuteNonQuery(); + stored.Add((key, embedding)); + } + + transaction.Commit(); + + foreach (var (key, embedding) in stored) + { + _entries[key] = embedding; + } + } + catch (Exception exception) when (exception is SqliteException or IOException or UnauthorizedAccessException) + { + _logger.LogWarning(exception, "Failed to persist embedding cache entries to {DatabasePath}.", _databasePath); + } + finally + { + _storeLock.Release(); + } + } + + private void Initialize() + { + var directory = Path.GetDirectoryName(_databasePath); + if (!string.IsNullOrWhiteSpace(directory)) + { + Directory.CreateDirectory(directory); + } + + using var connection = OpenConnection(); + + using (var pragma = connection.CreateCommand()) + { + pragma.CommandText = "PRAGMA journal_mode=WAL"; + pragma.ExecuteNonQuery(); + } + + using (var command = connection.CreateCommand()) + { + command.CommandText = """ + CREATE TABLE IF NOT EXISTS embedding_cache ( + model TEXT NOT NULL, + text TEXT NOT NULL, + embedding_json TEXT NOT NULL, + created_at_utc TEXT NOT NULL, + updated_at_utc TEXT NOT NULL, + PRIMARY KEY (model, text) + ) + """; + command.ExecuteNonQuery(); + } + + using (var command = connection.CreateCommand()) + { + command.CommandText = "SELECT model, text, embedding_json, updated_at_utc FROM embedding_cache"; + + using var reader = command.ExecuteReader(); + while (reader.Read()) + { + var key = new EmbeddingCacheKey(reader.GetString(0), reader.GetString(1)); + var embedding = new CachedEmbedding(reader.GetString(2), reader.GetString(3)); + _entries[key] = embedding; + } + } + + _logger.LogInformation( + "Loaded {EmbeddingCacheCount} embedding cache entries from {DatabasePath}.", + _entries.Count, + _databasePath); + } + + private SqliteConnection OpenConnection() + { + var connection = new SqliteConnection(_connectionString); + connection.Open(); + return connection; + } + + private static bool TryGetEndpointKind(PathString path, out EmbeddingEndpointKind kind) + { + var value = (path.Value ?? "").TrimEnd('/'); + if (value.Length == 0) + { + value = "/"; + } + + if (value.Equals("/api/embed", StringComparison.OrdinalIgnoreCase)) + { + kind = EmbeddingEndpointKind.OllamaEmbed; + return true; + } + + if (value.Equals("/api/embeddings", StringComparison.OrdinalIgnoreCase)) + { + kind = EmbeddingEndpointKind.OllamaEmbeddings; + return true; + } + + if (value.Equals("/v1/embeddings", StringComparison.OrdinalIgnoreCase)) + { + kind = EmbeddingEndpointKind.OpenAi; + return true; + } + + kind = default; + return false; + } + + 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 bool TryReadRequiredString(JsonElement root, string propertyName, out string value) + { + if (!TryReadString(root, propertyName, out value)) + { + return false; + } + + return !string.IsNullOrWhiteSpace(value); + } + + private static bool TryReadString(JsonElement root, string propertyName, out string value) + { + value = ""; + + if (!root.TryGetProperty(propertyName, out var element) + || element.ValueKind != JsonValueKind.String) + { + return false; + } + + value = element.GetString() ?? ""; + return true; + } + + private static bool TryReadInputTexts(JsonElement root, EmbeddingEndpointKind kind, out IReadOnlyList texts) + { + texts = []; + + if (kind == EmbeddingEndpointKind.OllamaEmbeddings) + { + if (!TryReadString(root, "prompt", out var prompt)) + { + return false; + } + + texts = [prompt]; + return true; + } + + if (!root.TryGetProperty("input", out var input)) + { + return false; + } + + if (input.ValueKind == JsonValueKind.String) + { + texts = [input.GetString() ?? ""]; + return true; + } + + if (input.ValueKind != JsonValueKind.Array) + { + return false; + } + + var values = new List(); + foreach (var item in input.EnumerateArray()) + { + if (item.ValueKind != JsonValueKind.String) + { + return false; + } + + values.Add(item.GetString() ?? ""); + } + + texts = values; + return values.Count > 0; + } + + private static bool HasContentEncoding(TunnelMessage responseHeaders) => + responseHeaders.Headers.Any(header => header.Name.Equals("Content-Encoding", StringComparison.OrdinalIgnoreCase)); + + private static List ExtractEmbeddings(EmbeddingCacheRequest request, byte[] body) + { + using var document = JsonDocument.Parse(body); + var root = document.RootElement; + + return request.Kind switch + { + EmbeddingEndpointKind.OllamaEmbeddings => ExtractOllamaEmbeddings(root), + EmbeddingEndpointKind.OllamaEmbed => ExtractOllamaEmbed(root), + EmbeddingEndpointKind.OpenAi => ExtractOpenAiEmbeddings(root), + _ => [] + }; + } + + private static List ExtractOllamaEmbeddings(JsonElement root) + { + if (root.ValueKind == JsonValueKind.Object + && root.TryGetProperty("embedding", out var embedding) + && embedding.ValueKind == JsonValueKind.Array) + { + return [new CachedEmbedding(embedding.GetRawText(), "")]; + } + + return []; + } + + private static List ExtractOllamaEmbed(JsonElement root) + { + if (root.ValueKind != JsonValueKind.Object + || !root.TryGetProperty("embeddings", out var embeddings) + || embeddings.ValueKind != JsonValueKind.Array) + { + return []; + } + + var values = new List(); + foreach (var embedding in embeddings.EnumerateArray()) + { + if (embedding.ValueKind != JsonValueKind.Array) + { + return []; + } + + values.Add(new CachedEmbedding(embedding.GetRawText(), "")); + } + + return values; + } + + private static List ExtractOpenAiEmbeddings(JsonElement root) + { + if (root.ValueKind != JsonValueKind.Object + || !root.TryGetProperty("data", out var data) + || data.ValueKind != JsonValueKind.Array) + { + return []; + } + + var values = new List<(int Index, int Position, CachedEmbedding Embedding)>(); + var position = 0; + + foreach (var item in data.EnumerateArray()) + { + if (item.ValueKind != JsonValueKind.Object + || !item.TryGetProperty("embedding", out var embedding) + || embedding.ValueKind != JsonValueKind.Array) + { + return []; + } + + var index = item.TryGetProperty("index", out var indexElement) + && indexElement.ValueKind == JsonValueKind.Number + && indexElement.TryGetInt32(out var parsedIndex) + ? parsedIndex + : position; + + values.Add((index, position, new CachedEmbedding(embedding.GetRawText(), ""))); + position++; + } + + return values + .OrderBy(value => value.Index) + .ThenBy(value => value.Position) + .Select(value => value.Embedding) + .ToList(); + } + + private static byte[] BuildResponseBody(EmbeddingCacheRequest request, IReadOnlyList embeddings) + { + using var memory = new MemoryStream(); + using var writer = new Utf8JsonWriter(memory); + + writer.WriteStartObject(); + + switch (request.Kind) + { + case EmbeddingEndpointKind.OllamaEmbeddings: + writer.WritePropertyName("embedding"); + writer.WriteRawValue(embeddings[0].EmbeddingJson); + break; + + case EmbeddingEndpointKind.OllamaEmbed: + writer.WriteString("model", request.Model); + writer.WritePropertyName("embeddings"); + WriteEmbeddingArray(writer, embeddings); + break; + + case EmbeddingEndpointKind.OpenAi: + writer.WriteString("object", "list"); + writer.WritePropertyName("data"); + writer.WriteStartArray(); + for (var index = 0; index < embeddings.Count; index++) + { + writer.WriteStartObject(); + writer.WriteString("object", "embedding"); + writer.WritePropertyName("embedding"); + writer.WriteRawValue(embeddings[index].EmbeddingJson); + writer.WriteNumber("index", index); + writer.WriteEndObject(); + } + + writer.WriteEndArray(); + writer.WriteString("model", request.Model); + writer.WriteStartObject("usage"); + writer.WriteNumber("prompt_tokens", 0); + writer.WriteNumber("total_tokens", 0); + writer.WriteEndObject(); + break; + } + + writer.WriteEndObject(); + writer.Flush(); + + return memory.ToArray(); + } + + private static void WriteEmbeddingArray(Utf8JsonWriter writer, IEnumerable embeddings) + { + writer.WriteStartArray(); + foreach (var embedding in embeddings) + { + writer.WriteRawValue(embedding.EmbeddingJson); + } + + writer.WriteEndArray(); + } + + private static string ResolveDatabasePath(string? configuredPath) + { + if (!string.IsNullOrWhiteSpace(configuredPath)) + { + var expanded = Environment.ExpandEnvironmentVariables(configuredPath); + return Path.IsPathRooted(expanded) + ? expanded + : Path.GetFullPath(Path.Combine(AppContext.BaseDirectory, expanded)); + } + + return Path.Combine(AppContext.BaseDirectory, "App_Data", "embedding-cache.sqlite"); + } +} + +internal sealed record EmbeddingCacheRequest( + EmbeddingEndpointKind Kind, + string Model, + IReadOnlyList Texts); + +internal enum EmbeddingEndpointKind +{ + OllamaEmbeddings, + OllamaEmbed, + OpenAi +} + +internal readonly record struct EmbeddingCacheKey(string Model, string Text); + +internal readonly record struct CachedEmbedding(string EmbeddingJson, string UpdatedAtUtc); diff --git a/src/ReverseLlama.Server/ManagementStore.cs b/src/ReverseLlama.Server/ManagementStore.cs new file mode 100644 index 0000000..67ec18f --- /dev/null +++ b/src/ReverseLlama.Server/ManagementStore.cs @@ -0,0 +1,700 @@ +using System.Security.Cryptography; +using System.Text; +using Microsoft.Data.Sqlite; + +namespace ReverseLlama.Server; + +internal sealed class ManagementStore +{ + private static readonly TimeSpan ApiKeyLastUsedWriteInterval = TimeSpan.FromMinutes(1); + + private readonly Dictionary _apiKeysByHash = new(StringComparer.Ordinal); + private readonly string _connectionString = ""; + private readonly string _databasePath = ""; + private readonly object _lock = new(); + private readonly ILogger _logger; + private bool _isAvailable; + private string? _lastError; + + public ManagementStore(ServerSettings settings, ILogger logger) + { + _logger = logger; + + try + { + _databasePath = ResolveDatabasePath(settings.ManagementDatabasePath); + _connectionString = new SqliteConnectionStringBuilder + { + DataSource = _databasePath, + Mode = SqliteOpenMode.ReadWriteCreate, + Pooling = true + }.ToString(); + + Initialize(); + _isAvailable = true; + } + catch (Exception exception) + { + _lastError = exception.Message; + _logger.LogError( + exception, + "Management database is disabled because SQLite could not be initialized at {DatabasePath}.", + string.IsNullOrWhiteSpace(_databasePath) ? settings.ManagementDatabasePath : _databasePath); + } + } + + public string DatabasePath => _databasePath; + + public bool IsAvailable => _isAvailable; + + public string? LastError => _lastError; + + public bool HasApiKeys + { + get + { + if (!_isAvailable) + { + return false; + } + + lock (_lock) + { + return _apiKeysByHash.Count > 0; + } + } + } + + public bool IsApiKeyValid(string apiKey, bool updateLastUsed) + { + if (!_isAvailable || string.IsNullOrWhiteSpace(apiKey)) + { + return false; + } + + var hash = HashApiKey(apiKey); + var now = DateTimeOffset.UtcNow; + + lock (_lock) + { + if (!_apiKeysByHash.TryGetValue(hash, out var key)) + { + return false; + } + + if (!updateLastUsed + || key.LastUsedUtc is not null + && now - key.LastUsedUtc.Value < ApiKeyLastUsedWriteInterval) + { + return true; + } + + key.LastUsedUtc = now; + + try + { + using var connection = OpenConnection(); + using var command = connection.CreateCommand(); + command.CommandText = "UPDATE api_keys SET last_used_at_utc = $last_used_at_utc WHERE id = $id"; + command.Parameters.AddWithValue("$last_used_at_utc", now.ToString("O")); + command.Parameters.AddWithValue("$id", key.Id); + command.ExecuteNonQuery(); + } + catch (Exception exception) when (exception is SqliteException or IOException or UnauthorizedAccessException) + { + _logger.LogWarning(exception, "Failed to update API key last-used timestamp."); + } + + return true; + } + } + + public IReadOnlyList ListApiKeys() + { + if (!_isAvailable) + { + return []; + } + + lock (_lock) + { + return _apiKeysByHash.Values + .OrderBy(key => key.Name, StringComparer.OrdinalIgnoreCase) + .ThenBy(key => key.CreatedAtUtc) + .Select(key => new ApiKeyInfo( + key.Id, + key.Name, + key.KeyPrefix, + key.CreatedAtUtc, + key.LastUsedUtc)) + .ToList(); + } + } + + public CreatedApiKey CreateApiKey(string? name) + { + EnsureAvailable(); + + var apiKey = GenerateApiKey(); + var now = DateTimeOffset.UtcNow; + var state = new ApiKeyState + { + Id = Guid.NewGuid().ToString("n"), + Name = string.IsNullOrWhiteSpace(name) ? "API key" : name.Trim(), + KeyHash = HashApiKey(apiKey), + KeyPrefix = GetKeyPrefix(apiKey), + CreatedAtUtc = now + }; + + lock (_lock) + { + using var connection = OpenConnection(); + using var command = connection.CreateCommand(); + command.CommandText = """ + INSERT INTO api_keys (id, name, key_hash, key_prefix, created_at_utc) + VALUES ($id, $name, $key_hash, $key_prefix, $created_at_utc) + """; + command.Parameters.AddWithValue("$id", state.Id); + command.Parameters.AddWithValue("$name", state.Name); + command.Parameters.AddWithValue("$key_hash", state.KeyHash); + command.Parameters.AddWithValue("$key_prefix", state.KeyPrefix); + command.Parameters.AddWithValue("$created_at_utc", state.CreatedAtUtc.ToString("O")); + command.ExecuteNonQuery(); + + _apiKeysByHash[state.KeyHash] = state; + } + + return new CreatedApiKey( + state.Id, + state.Name, + state.KeyPrefix, + state.CreatedAtUtc, + apiKey); + } + + public bool DeleteApiKey(string id) + { + if (!_isAvailable || string.IsNullOrWhiteSpace(id)) + { + return false; + } + + lock (_lock) + { + using var connection = OpenConnection(); + using var command = connection.CreateCommand(); + command.CommandText = "DELETE FROM api_keys WHERE id = $id"; + command.Parameters.AddWithValue("$id", id); + var deleted = command.ExecuteNonQuery() > 0; + + if (deleted) + { + foreach (var pair in _apiKeysByHash.Where(pair => pair.Value.Id == id).ToArray()) + { + _apiKeysByHash.Remove(pair.Key); + } + } + + return deleted; + } + } + + public ClientAccess GetClientAccess(string clientId) + { + if (!_isAvailable || string.IsNullOrWhiteSpace(clientId)) + { + return ClientAccess.Enabled; + } + + lock (_lock) + { + using var connection = OpenConnection(); + using var command = connection.CreateCommand(); + command.CommandText = """ + SELECT disabled_until_utc, disabled_manually, disabled_reason + FROM client_controls + WHERE client_id = $client_id + """; + command.Parameters.AddWithValue("$client_id", clientId); + + using var reader = command.ExecuteReader(); + if (!reader.Read()) + { + return ClientAccess.Enabled; + } + + var disabledUntil = ReadNullableDateTimeOffset(reader, 0); + var disabledManually = reader.GetInt32(1) != 0; + var reason = reader.IsDBNull(2) ? null : reader.GetString(2); + + if (disabledManually) + { + return new ClientAccess(true, null, true, reason); + } + + if (disabledUntil is { } until && until > DateTimeOffset.UtcNow) + { + return new ClientAccess(true, until, false, reason); + } + + return ClientAccess.Enabled; + } + } + + public IReadOnlyDictionary ListClientControls() + { + if (!_isAvailable) + { + return new Dictionary(StringComparer.OrdinalIgnoreCase); + } + + var result = new Dictionary(StringComparer.OrdinalIgnoreCase); + + lock (_lock) + { + using var connection = OpenConnection(); + using var command = connection.CreateCommand(); + command.CommandText = "SELECT client_id, disabled_until_utc, disabled_manually, disabled_reason FROM client_controls"; + + using var reader = command.ExecuteReader(); + while (reader.Read()) + { + var clientId = reader.GetString(0); + var disabledUntil = ReadNullableDateTimeOffset(reader, 1); + var disabledManually = reader.GetInt32(2) != 0; + var reason = reader.IsDBNull(3) ? null : reader.GetString(3); + + result[clientId] = disabledManually + ? new ClientAccess(true, null, true, reason) + : disabledUntil is { } until && until > DateTimeOffset.UtcNow + ? new ClientAccess(true, until, false, reason) + : ClientAccess.Enabled; + } + } + + return result; + } + + public void DisableClient(string clientId, TimeSpan? duration, bool manually, string? reason) + { + EnsureAvailable(); + + if (string.IsNullOrWhiteSpace(clientId)) + { + throw new ArgumentException("Client id is required.", nameof(clientId)); + } + + var now = DateTimeOffset.UtcNow; + var disabledUntil = manually ? null : now.Add(duration ?? TimeSpan.FromHours(1)).ToString("O"); + + lock (_lock) + { + using var connection = OpenConnection(); + using var command = connection.CreateCommand(); + command.CommandText = """ + INSERT INTO client_controls ( + client_id, + disabled_until_utc, + disabled_manually, + disabled_reason, + updated_at_utc) + VALUES ( + $client_id, + $disabled_until_utc, + $disabled_manually, + $disabled_reason, + $updated_at_utc) + ON CONFLICT(client_id) DO UPDATE SET + disabled_until_utc = excluded.disabled_until_utc, + disabled_manually = excluded.disabled_manually, + disabled_reason = excluded.disabled_reason, + updated_at_utc = excluded.updated_at_utc + """; + command.Parameters.AddWithValue("$client_id", clientId); + command.Parameters.AddWithValue("$disabled_until_utc", (object?)disabledUntil ?? DBNull.Value); + command.Parameters.AddWithValue("$disabled_manually", manually ? 1 : 0); + command.Parameters.AddWithValue("$disabled_reason", string.IsNullOrWhiteSpace(reason) ? DBNull.Value : reason.Trim()); + command.Parameters.AddWithValue("$updated_at_utc", now.ToString("O")); + command.ExecuteNonQuery(); + } + } + + public void EnableClient(string clientId) + { + EnsureAvailable(); + + if (string.IsNullOrWhiteSpace(clientId)) + { + throw new ArgumentException("Client id is required.", nameof(clientId)); + } + + lock (_lock) + { + using var connection = OpenConnection(); + using var command = connection.CreateCommand(); + command.CommandText = """ + INSERT INTO client_controls ( + client_id, + disabled_until_utc, + disabled_manually, + disabled_reason, + updated_at_utc) + VALUES ( + $client_id, + NULL, + 0, + NULL, + $updated_at_utc) + ON CONFLICT(client_id) DO UPDATE SET + disabled_until_utc = NULL, + disabled_manually = 0, + disabled_reason = NULL, + updated_at_utc = excluded.updated_at_utc + """; + command.Parameters.AddWithValue("$client_id", clientId); + command.Parameters.AddWithValue("$updated_at_utc", DateTimeOffset.UtcNow.ToString("O")); + command.ExecuteNonQuery(); + } + } + + public void RecordRequest(RequestMetric metric) + { + if (!_isAvailable) + { + return; + } + + try + { + lock (_lock) + { + using var connection = OpenConnection(); + using var command = connection.CreateCommand(); + command.CommandText = """ + INSERT INTO request_metrics ( + client_id, + model, + method, + path, + status_code, + token_count, + started_at_utc, + completed_at_utc, + duration_ms) + VALUES ( + $client_id, + $model, + $method, + $path, + $status_code, + $token_count, + $started_at_utc, + $completed_at_utc, + $duration_ms) + """; + command.Parameters.AddWithValue("$client_id", metric.ClientId); + command.Parameters.AddWithValue("$model", string.IsNullOrWhiteSpace(metric.Model) ? DBNull.Value : metric.Model); + command.Parameters.AddWithValue("$method", metric.Method); + command.Parameters.AddWithValue("$path", metric.Path); + command.Parameters.AddWithValue("$status_code", metric.StatusCode is null ? DBNull.Value : metric.StatusCode.Value); + command.Parameters.AddWithValue("$token_count", metric.TokenCount); + command.Parameters.AddWithValue("$started_at_utc", metric.StartedAtUtc.ToString("O")); + command.Parameters.AddWithValue("$completed_at_utc", metric.CompletedAtUtc.ToString("O")); + command.Parameters.AddWithValue("$duration_ms", metric.Duration.TotalMilliseconds); + command.ExecuteNonQuery(); + } + } + catch (Exception exception) when (exception is SqliteException or IOException or UnauthorizedAccessException) + { + _logger.LogWarning(exception, "Failed to record request metric for client {ClientId}.", metric.ClientId); + } + } + + public IReadOnlyDictionary GetClientRequestStats() + { + if (!_isAvailable) + { + return new Dictionary(StringComparer.OrdinalIgnoreCase); + } + + var now = DateTimeOffset.UtcNow; + var since10 = now.AddMinutes(-10).ToString("O"); + var sinceHour = now.AddHours(-1).ToString("O"); + var result = new Dictionary(StringComparer.OrdinalIgnoreCase); + + lock (_lock) + { + using var connection = OpenConnection(); + using var command = connection.CreateCommand(); + command.CommandText = """ + SELECT + client_id, + COUNT(*), + SUM(CASE WHEN started_at_utc >= $since10 THEN 1 ELSE 0 END), + SUM(CASE WHEN started_at_utc >= $sinceHour THEN 1 ELSE 0 END) + FROM request_metrics + GROUP BY client_id + """; + command.Parameters.AddWithValue("$since10", since10); + command.Parameters.AddWithValue("$sinceHour", sinceHour); + + using var reader = command.ExecuteReader(); + while (reader.Read()) + { + result[reader.GetString(0)] = new ClientRequestStats( + reader.GetInt64(1), + reader.GetInt64(2), + reader.GetInt64(3)); + } + } + + return result; + } + + public IReadOnlyDictionary GetModelUsageStats() + { + if (!_isAvailable) + { + return new Dictionary(StringComparer.OrdinalIgnoreCase); + } + + var now = DateTimeOffset.UtcNow; + var since10 = now.AddMinutes(-10).ToString("O"); + var sinceHour = now.AddHours(-1).ToString("O"); + var result = new Dictionary(StringComparer.OrdinalIgnoreCase); + + lock (_lock) + { + using var connection = OpenConnection(); + using var command = connection.CreateCommand(); + command.CommandText = """ + SELECT + model, + COUNT(*), + SUM(CASE WHEN started_at_utc >= $since10 THEN 1 ELSE 0 END), + SUM(CASE WHEN started_at_utc >= $sinceHour THEN 1 ELSE 0 END), + SUM(CASE WHEN started_at_utc >= $since10 THEN token_count ELSE 0 END), + SUM(CASE WHEN started_at_utc >= $sinceHour THEN token_count ELSE 0 END) + FROM request_metrics + WHERE model IS NOT NULL AND model <> '' + GROUP BY model + """; + command.Parameters.AddWithValue("$since10", since10); + command.Parameters.AddWithValue("$sinceHour", sinceHour); + + using var reader = command.ExecuteReader(); + while (reader.Read()) + { + result[reader.GetString(0)] = new ModelUsageStats( + reader.GetInt64(1), + reader.GetInt64(2), + reader.GetInt64(3), + reader.GetInt64(4), + reader.GetInt64(5)); + } + } + + return result; + } + + private void Initialize() + { + var directory = Path.GetDirectoryName(_databasePath); + if (!string.IsNullOrWhiteSpace(directory)) + { + Directory.CreateDirectory(directory); + } + + using var connection = OpenConnection(); + + using (var pragma = connection.CreateCommand()) + { + pragma.CommandText = "PRAGMA journal_mode=WAL"; + pragma.ExecuteNonQuery(); + } + + using (var command = connection.CreateCommand()) + { + command.CommandText = """ + CREATE TABLE IF NOT EXISTS client_controls ( + client_id TEXT NOT NULL PRIMARY KEY, + disabled_until_utc TEXT NULL, + disabled_manually INTEGER NOT NULL DEFAULT 0, + disabled_reason TEXT NULL, + updated_at_utc TEXT NOT NULL + ); + + CREATE TABLE IF NOT EXISTS api_keys ( + id TEXT NOT NULL PRIMARY KEY, + name TEXT NOT NULL, + key_hash TEXT NOT NULL UNIQUE, + key_prefix TEXT NOT NULL, + created_at_utc TEXT NOT NULL, + last_used_at_utc TEXT NULL + ); + + CREATE TABLE IF NOT EXISTS request_metrics ( + id INTEGER NOT NULL PRIMARY KEY AUTOINCREMENT, + client_id TEXT NOT NULL, + model TEXT NULL, + method TEXT NOT NULL, + path TEXT NOT NULL, + status_code INTEGER NULL, + token_count INTEGER NOT NULL DEFAULT 0, + started_at_utc TEXT NOT NULL, + completed_at_utc TEXT NOT NULL, + duration_ms REAL NOT NULL + ); + + CREATE INDEX IF NOT EXISTS idx_request_metrics_client_started + ON request_metrics (client_id, started_at_utc); + + CREATE INDEX IF NOT EXISTS idx_request_metrics_model_started + ON request_metrics (model, started_at_utc); + """; + command.ExecuteNonQuery(); + } + + using (var command = connection.CreateCommand()) + { + command.CommandText = """ + SELECT id, name, key_hash, key_prefix, created_at_utc, last_used_at_utc + FROM api_keys + """; + + using var reader = command.ExecuteReader(); + while (reader.Read()) + { + var state = new ApiKeyState + { + Id = reader.GetString(0), + Name = reader.GetString(1), + KeyHash = reader.GetString(2), + KeyPrefix = reader.GetString(3), + CreatedAtUtc = ReadDateTimeOffset(reader.GetString(4)), + LastUsedUtc = ReadNullableDateTimeOffset(reader, 5) + }; + + _apiKeysByHash[state.KeyHash] = state; + } + } + + _logger.LogInformation( + "Loaded {ApiKeyCount} API key(s) from {DatabasePath}.", + _apiKeysByHash.Count, + _databasePath); + } + + private SqliteConnection OpenConnection() + { + var connection = new SqliteConnection(_connectionString); + connection.Open(); + return connection; + } + + private void EnsureAvailable() + { + if (!_isAvailable) + { + throw new InvalidOperationException(_lastError ?? "The management database is not available."); + } + } + + private static DateTimeOffset? ReadNullableDateTimeOffset(SqliteDataReader reader, int ordinal) => + reader.IsDBNull(ordinal) ? null : ReadDateTimeOffset(reader.GetString(ordinal)); + + private static DateTimeOffset ReadDateTimeOffset(string value) => + DateTimeOffset.TryParse(value, out var parsed) ? parsed : DateTimeOffset.MinValue; + + private static string GenerateApiKey() + { + var bytes = RandomNumberGenerator.GetBytes(32); + return $"rl_{Base64UrlEncode(bytes)}"; + } + + private static string HashApiKey(string apiKey) => + Convert.ToHexString(SHA256.HashData(Encoding.UTF8.GetBytes(apiKey))).ToLowerInvariant(); + + private static string GetKeyPrefix(string apiKey) => + apiKey.Length <= 12 ? apiKey : apiKey[..12]; + + private static string Base64UrlEncode(byte[] bytes) => + Convert.ToBase64String(bytes) + .TrimEnd('=') + .Replace('+', '-') + .Replace('/', '_'); + + private static string ResolveDatabasePath(string? configuredPath) + { + if (!string.IsNullOrWhiteSpace(configuredPath)) + { + var expanded = Environment.ExpandEnvironmentVariables(configuredPath); + return Path.IsPathRooted(expanded) + ? expanded + : Path.GetFullPath(Path.Combine(AppContext.BaseDirectory, expanded)); + } + + return Path.Combine(AppContext.BaseDirectory, "App_Data", "management.sqlite"); + } + + private sealed class ApiKeyState + { + public string Id { get; init; } = ""; + + public string Name { get; init; } = ""; + + public string KeyHash { get; init; } = ""; + + public string KeyPrefix { get; init; } = ""; + + public DateTimeOffset CreatedAtUtc { get; init; } + + public DateTimeOffset? LastUsedUtc { get; set; } + } +} + +internal sealed record ClientAccess( + bool IsDisabled, + DateTimeOffset? DisabledUntilUtc, + bool DisabledManually, + string? DisabledReason) +{ + public static ClientAccess Enabled { get; } = new(false, null, false, null); +} + +internal sealed record ApiKeyInfo( + string Id, + string Name, + string KeyPrefix, + DateTimeOffset CreatedAtUtc, + DateTimeOffset? LastUsedUtc); + +internal sealed record CreatedApiKey( + string Id, + string Name, + string KeyPrefix, + DateTimeOffset CreatedAtUtc, + string Key); + +internal sealed record RequestMetric( + string ClientId, + string? Model, + string Method, + string Path, + int? StatusCode, + int TokenCount, + DateTimeOffset StartedAtUtc, + DateTimeOffset CompletedAtUtc, + TimeSpan Duration); + +internal sealed record ClientRequestStats( + long Total, + long Last10Minutes, + long LastHour); + +internal sealed record ModelUsageStats( + long TotalRequests, + long RequestsLast10Minutes, + long RequestsLastHour, + long TokensLast10Minutes, + long TokensLastHour); diff --git a/src/ReverseLlama.Server/PendingCommand.cs b/src/ReverseLlama.Server/PendingCommand.cs new file mode 100644 index 0000000..a2e11c4 --- /dev/null +++ b/src/ReverseLlama.Server/PendingCommand.cs @@ -0,0 +1,17 @@ +using ReverseLlama.Protocol; + +namespace ReverseLlama.Server; + +internal sealed class PendingCommand +{ + private readonly TaskCompletionSource _completion = new(TaskCreationOptions.RunContinuationsAsynchronously); + + public Task WaitAsync(CancellationToken cancellationToken) => + _completion.Task.WaitAsync(cancellationToken); + + public void Complete(TunnelMessage message) => + _completion.TrySetResult(message); + + public void Fail(string error) => + _completion.TrySetException(new InvalidOperationException(error)); +} diff --git a/src/ReverseLlama.Server/PendingProxyRequest.cs b/src/ReverseLlama.Server/PendingProxyRequest.cs new file mode 100644 index 0000000..3bf39ae --- /dev/null +++ b/src/ReverseLlama.Server/PendingProxyRequest.cs @@ -0,0 +1,52 @@ +using System.Threading.Channels; +using ReverseLlama.Protocol; + +namespace ReverseLlama.Server; + +internal sealed class PendingProxyRequest +{ + private readonly Channel _body = Channel.CreateUnbounded( + new UnboundedChannelOptions + { + SingleReader = true, + SingleWriter = false + }); + + private readonly TaskCompletionSource _responseHeaders = + new(TaskCreationOptions.RunContinuationsAsynchronously); + + public ChannelReader Body => _body.Reader; + + public Task WaitForHeadersAsync(CancellationToken cancellationToken) => + _responseHeaders.Task.WaitAsync(cancellationToken); + + public void SetResponseHeaders(TunnelMessage message) => + _responseHeaders.TrySetResult(message); + + public void AddBody(byte[] body) + { + if (body.Length > 0) + { + _body.Writer.TryWrite(body); + } + } + + public void Complete() + { + if (!_responseHeaders.Task.IsCompleted) + { + _responseHeaders.TrySetException(new InvalidOperationException("The client completed a response before sending response headers.")); + } + + _body.Writer.TryComplete(); + } + + public void Fail(string message) => + Fail(new InvalidOperationException(message)); + + public void Fail(Exception exception) + { + _responseHeaders.TrySetException(exception); + _body.Writer.TryComplete(exception); + } +} diff --git a/src/ReverseLlama.Server/Program.cs b/src/ReverseLlama.Server/Program.cs new file mode 100644 index 0000000..1473b75 --- /dev/null +++ b/src/ReverseLlama.Server/Program.cs @@ -0,0 +1,171 @@ +using System.Data.SqlClient; +using System.Net.WebSockets; +using ElmahCore; +using ElmahCore.Mvc; +using Microsoft.AspNetCore.Authentication.Cookies; +using Microsoft.AspNetCore.Authentication.OpenIdConnect; +using Microsoft.IdentityModel.Protocols.OpenIdConnect; +using ReverseLlama.Protocol; +using ReverseLlama.Server; + +var builder = WebApplication.CreateBuilder(args); +var settings = ServerSettings.FromConfiguration(builder.Configuration); + +builder.Services.AddSingleton(settings); +builder.Services.AddSingleton(); +builder.Services.AddSingleton(); +builder.Services.AddSingleton(); +builder.Services.AddElmah().Configure( + options => options.ConnectionString = builder.Configuration.GetConnectionString("ElmahConnection")); + +if (settings.Keycloak.IsConfigured) +{ + builder.Services + .AddAuthentication(options => + { + options.DefaultScheme = CookieAuthenticationDefaults.AuthenticationScheme; + options.DefaultChallengeScheme = OpenIdConnectDefaults.AuthenticationScheme; + }) + .AddCookie(options => + { + options.Cookie.Name = "ReverseLlama.Admin"; + options.Cookie.SameSite = SameSiteMode.Lax; + options.Cookie.SecurePolicy = CookieSecurePolicy.SameAsRequest; + options.LoginPath = "/admin/login"; + options.LogoutPath = "/admin/logout"; + }) + .AddOpenIdConnect(options => + { + options.Authority = settings.Keycloak.Authority; + options.ClientId = settings.Keycloak.ClientId; + options.ClientSecret = settings.Keycloak.ClientSecret; + options.RequireHttpsMetadata = settings.Keycloak.RequireHttpsMetadata; + options.ResponseType = OpenIdConnectResponseType.Code; + options.ResponseMode = OpenIdConnectResponseMode.Query; + options.SaveTokens = true; + options.GetClaimsFromUserInfoEndpoint = true; + options.CorrelationCookie.SameSite = SameSiteMode.Lax; + options.CorrelationCookie.SecurePolicy = CookieSecurePolicy.SameAsRequest; + options.NonceCookie.SameSite = SameSiteMode.Lax; + options.NonceCookie.SecurePolicy = CookieSecurePolicy.SameAsRequest; + options.Scope.Clear(); + options.Scope.Add("openid"); + options.Scope.Add("profile"); + options.Scope.Add("email"); + options.Events = new OpenIdConnectEvents + { + OnRemoteFailure = context => + { + var errorLog = context.HttpContext.RequestServices.GetService(); + if (context.Failure is not null) + { + errorLog?.Log(new Error(context.Failure)); + } + + context.HandleResponse(); + context.Response.Redirect("/admin/auth-error"); + return Task.CompletedTask; + } + }; + }); +} + +builder.Services.AddAuthorization(); + +var app = builder.Build(); + +if (settings.Keycloak.IsConfigured) +{ + app.UseAuthentication(); + app.UseAuthorization(); +} + +app.UseElmah(); + +app.UseWebSockets(new WebSocketOptions +{ + KeepAliveInterval = TimeSpan.FromSeconds(30) +}); + +app.MapAdminEndpoints(settings); + +app.MapGet("/", (TunnelHub hub) => + Results.Json(new + { + status = "ok", + connected = hub.HasClient, + pendingRequests = hub.PendingRequestCount, + clients = hub.ClientsSnapshot.Count + })); + +app.MapGet(settings.StatusPath, (HttpContext context, TunnelHub hub, ServerSettings serverSettings, EmbeddingCache embeddingCache, ManagementStore managementStore) => +{ + // Query token allowed so the status page can be checked in a browser. + if (!TokenAuthentication.IsAuthorized(context.Request, serverSettings, managementStore, allowQueryToken: true)) + { + return Results.Unauthorized(); + } + + return Results.Json(new + { + connected = hub.HasClient, + pendingRequests = hub.PendingRequestCount, + tunnelPath = serverSettings.TunnelPath, + embeddingCache = new + { + available = embeddingCache.IsAvailable, + count = embeddingCache.Count, + databasePath = embeddingCache.DatabasePath, + lastError = embeddingCache.LastError + }, + management = new + { + available = managementStore.IsAvailable, + databasePath = managementStore.DatabasePath, + lastError = managementStore.LastError + }, + clients = hub.ClientsSnapshot + }); +}); + +app.Map(settings.TunnelPath, async (HttpContext context, TunnelHub hub, ServerSettings serverSettings, ManagementStore managementStore) => +{ + if (!TokenAuthentication.IsAuthorized(context.Request, serverSettings, managementStore, allowQueryToken: true)) + { + context.Response.StatusCode = StatusCodes.Status401Unauthorized; + await context.Response.WriteAsync($"Missing or invalid {ProtocolConstants.TokenHeader}.", context.RequestAborted); + return; + } + + if (!context.WebSockets.IsWebSocketRequest) + { + context.Response.StatusCode = StatusCodes.Status400BadRequest; + await context.Response.WriteAsync("This endpoint only accepts WebSocket tunnel connections.", context.RequestAborted); + return; + } + + var clientId = context.Request.Headers[ProtocolConstants.ClientIdHeader].FirstOrDefault(); + if (string.IsNullOrWhiteSpace(clientId)) + { + clientId = $"anonymous-{Guid.NewGuid():n}"; + } + + using var socket = await context.WebSockets.AcceptWebSocketAsync(); + await hub.AcceptAsync(clientId, socket, context.RequestAborted); +}); + +app.Map("/clients/{clientId}/{**path}", ReverseProxyEndpoint.HandleClientAsync); + +app.Map("/{**path}", ReverseProxyEndpoint.HandleRootAsync) + .WithOrder(1000); + +var elmahService = app.Services.GetRequiredService(); +try +{ + app.Run(); +} +catch (Exception exception) +{ + elmahService.Log(new Error(exception)); + throw; +} diff --git a/src/ReverseLlama.Server/Properties/launchSettings.json b/src/ReverseLlama.Server/Properties/launchSettings.json new file mode 100644 index 0000000..f87c59f --- /dev/null +++ b/src/ReverseLlama.Server/Properties/launchSettings.json @@ -0,0 +1,38 @@ +{ + "$schema": "http://json.schemastore.org/launchsettings.json", + "iisSettings": { + "windowsAuthentication": false, + "anonymousAuthentication": true, + "iisExpress": { + "applicationUrl": "http://localhost:4407", + "sslPort": 44305 + } + }, + "profiles": { + "http": { + "commandName": "Project", + "dotnetRunMessages": true, + "launchBrowser": true, + "applicationUrl": "http://localhost:5001", + "environmentVariables": { + "ASPNETCORE_ENVIRONMENT": "Development" + } + }, + "https": { + "commandName": "Project", + "dotnetRunMessages": true, + "launchBrowser": true, + "applicationUrl": "https://localhost:7183;http://localhost:5174", + "environmentVariables": { + "ASPNETCORE_ENVIRONMENT": "Development" + } + }, + "IIS Express": { + "commandName": "IISExpress", + "launchBrowser": true, + "environmentVariables": { + "ASPNETCORE_ENVIRONMENT": "Development" + } + } + } +} diff --git a/src/ReverseLlama.Server/ResponseTokenCounter.cs b/src/ReverseLlama.Server/ResponseTokenCounter.cs new file mode 100644 index 0000000..b2978bf --- /dev/null +++ b/src/ReverseLlama.Server/ResponseTokenCounter.cs @@ -0,0 +1,157 @@ +using System.Text; +using System.Text.Json; + +namespace ReverseLlama.Server; + +internal sealed class ResponseTokenCounter +{ + private const int MaxBufferedBytes = 4 * 1024 * 1024; + + private readonly MemoryStream _buffer = new(); + + public void Add(ReadOnlySpan chunk) + { + if (chunk.Length == 0 || _buffer.Length >= MaxBufferedBytes) + { + return; + } + + var available = MaxBufferedBytes - (int)_buffer.Length; + var length = Math.Min(chunk.Length, available); + _buffer.Write(chunk[..length]); + } + + public int CountTokens() + { + if (_buffer.Length == 0) + { + return 0; + } + + var payload = Encoding.UTF8.GetString(_buffer.ToArray()); + var total = 0; + var parsedLines = false; + + foreach (var rawLine in payload.Split('\n')) + { + var line = rawLine.Trim(); + if (line.Length == 0) + { + continue; + } + + if (line.StartsWith("data:", StringComparison.OrdinalIgnoreCase)) + { + line = line["data:".Length..].Trim(); + } + + if (line.Equals("[DONE]", StringComparison.OrdinalIgnoreCase)) + { + continue; + } + + if (TryExtractFromJson(line, out var lineTokens)) + { + parsedLines = true; + total += lineTokens; + } + } + + if (parsedLines) + { + return total; + } + + return TryExtractFromJson(payload, out var tokens) ? tokens : 0; + } + + private static bool TryExtractFromJson(string json, out int tokens) + { + tokens = 0; + + try + { + using var document = JsonDocument.Parse(json); + tokens = ExtractTokens(document.RootElement); + return tokens > 0; + } + catch (JsonException) + { + return false; + } + } + + private static int ExtractTokens(JsonElement element) + { + if (element.ValueKind == JsonValueKind.Array) + { + var total = 0; + foreach (var item in element.EnumerateArray()) + { + total += ExtractTokens(item); + } + + return total; + } + + if (element.ValueKind != JsonValueKind.Object) + { + return 0; + } + + if (element.TryGetProperty("usage", out var usage) && usage.ValueKind == JsonValueKind.Object) + { + if (TryGetInt(usage, "total_tokens", out var totalTokens)) + { + return totalTokens; + } + + var usageTotal = 0; + if (TryGetInt(usage, "prompt_tokens", out var promptTokens)) + { + usageTotal += promptTokens; + } + + if (TryGetInt(usage, "completion_tokens", out var completionTokens)) + { + usageTotal += completionTokens; + } + + if (TryGetInt(usage, "input_tokens", out var inputTokens)) + { + usageTotal += inputTokens; + } + + if (TryGetInt(usage, "output_tokens", out var outputTokens)) + { + usageTotal += outputTokens; + } + + if (usageTotal > 0) + { + return usageTotal; + } + } + + var ollamaTotal = 0; + if (TryGetInt(element, "prompt_eval_count", out var promptEvalCount)) + { + ollamaTotal += promptEvalCount; + } + + if (TryGetInt(element, "eval_count", out var evalCount)) + { + ollamaTotal += evalCount; + } + + return ollamaTotal; + } + + private static bool TryGetInt(JsonElement element, string propertyName, out int value) + { + value = 0; + return element.TryGetProperty(propertyName, out var property) + && property.ValueKind == JsonValueKind.Number + && property.TryGetInt32(out value); + } +} diff --git a/src/ReverseLlama.Server/ReverseLlama.Server.csproj b/src/ReverseLlama.Server/ReverseLlama.Server.csproj new file mode 100644 index 0000000..cac2e31 --- /dev/null +++ b/src/ReverseLlama.Server/ReverseLlama.Server.csproj @@ -0,0 +1,20 @@ + + + + + + + + + + + + + + + net10.0 + enable + enable + + + diff --git a/src/ReverseLlama.Server/ReverseProxyEndpoint.cs b/src/ReverseLlama.Server/ReverseProxyEndpoint.cs new file mode 100644 index 0000000..abf16de --- /dev/null +++ b/src/ReverseLlama.Server/ReverseProxyEndpoint.cs @@ -0,0 +1,639 @@ +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); +} diff --git a/src/ReverseLlama.Server/ServerSettings.cs b/src/ReverseLlama.Server/ServerSettings.cs new file mode 100644 index 0000000..2c0ca3d --- /dev/null +++ b/src/ReverseLlama.Server/ServerSettings.cs @@ -0,0 +1,98 @@ +using Microsoft.Extensions.Configuration; +using ReverseLlama.Protocol; + +namespace ReverseLlama.Server; + +internal sealed class ServerSettings +{ + public string StatusPath { get; init; } = ProtocolConstants.DefaultStatusPath; + + public string TunnelPath { get; init; } = ProtocolConstants.DefaultTunnelPath; + + public string? Token { get; init; } + + public int ChunkSize { get; init; } = 64 * 1024; + + public string? EmbeddingCachePath { get; init; } + + public string? ManagementDatabasePath { get; init; } + + public KeycloakSettings Keycloak { get; init; } = new(); + + public static ServerSettings FromConfiguration(IConfiguration configuration) + { + return new ServerSettings + { + StatusPath = NormalizePath(Read(configuration, "ReverseLlama:StatusPath", "status-path") ?? ProtocolConstants.DefaultStatusPath), + TunnelPath = NormalizePath(Read(configuration, "ReverseLlama:TunnelPath", "tunnel-path") ?? ProtocolConstants.DefaultTunnelPath), + Token = Read(configuration, "ReverseLlama:Token", "token") ?? Environment.GetEnvironmentVariable("REVERSE_LLAMA_TOKEN"), + ChunkSize = ReadInt(configuration, 64 * 1024, "ReverseLlama:ChunkSize", "chunk-size", "REVERSE_LLAMA_CHUNK_SIZE"), + EmbeddingCachePath = Read( + configuration, + "ReverseLlama:EmbeddingCachePath", + "embedding-cache-path", + "REVERSE_LLAMA_EMBEDDING_CACHE_PATH"), + ManagementDatabasePath = Read( + configuration, + "ReverseLlama:ManagementDatabasePath", + "management-database-path", + "REVERSE_LLAMA_MANAGEMENT_DATABASE_PATH"), + Keycloak = new KeycloakSettings + { + Authority = Read(configuration, "Authentication:Keycloak:Authority", "REVERSE_LLAMA_KEYCLOAK_AUTHORITY"), + ClientId = Read(configuration, "Authentication:Keycloak:ClientId", "REVERSE_LLAMA_KEYCLOAK_CLIENT_ID"), + ClientSecret = Read(configuration, "Authentication:Keycloak:ClientSecret", "REVERSE_LLAMA_KEYCLOAK_CLIENT_SECRET"), + RequireHttpsMetadata = ReadBool( + configuration, + true, + "Authentication:Keycloak:RequireHttpsMetadata", + "REVERSE_LLAMA_KEYCLOAK_REQUIRE_HTTPS_METADATA") + } + }; + } + + private static string? Read(IConfiguration configuration, params string[] keys) + { + foreach (var key in keys) + { + var value = configuration[key] ?? Environment.GetEnvironmentVariable(key); + if (!string.IsNullOrWhiteSpace(value)) + { + return value; + } + } + + return null; + } + + private static int ReadInt(IConfiguration configuration, int fallback, params string[] keys) + { + var value = Read(configuration, keys); + return int.TryParse(value, out var parsed) && parsed > 0 ? parsed : fallback; + } + + private static bool ReadBool(IConfiguration configuration, bool fallback, params string[] keys) + { + var value = Read(configuration, keys); + return bool.TryParse(value, out var parsed) ? parsed : fallback; + } + + private static string NormalizePath(string path) => + path.StartsWith('/') ? path : $"/{path}"; +} + +internal sealed class KeycloakSettings +{ + public string? Authority { get; init; } + + public string? ClientId { get; init; } + + public string? ClientSecret { get; init; } + + public bool RequireHttpsMetadata { get; init; } = true; + + public bool IsConfigured => + !string.IsNullOrWhiteSpace(Authority) + && !string.IsNullOrWhiteSpace(ClientId) + && !string.IsNullOrWhiteSpace(ClientSecret); +} diff --git a/src/ReverseLlama.Server/TokenAuthentication.cs b/src/ReverseLlama.Server/TokenAuthentication.cs new file mode 100644 index 0000000..93901b3 --- /dev/null +++ b/src/ReverseLlama.Server/TokenAuthentication.cs @@ -0,0 +1,133 @@ +using ReverseLlama.Protocol; + +namespace ReverseLlama.Server; + +internal static class TokenAuthentication +{ + private static readonly PathString PathTokenPrefix = new("/token"); + + public static bool IsAuthorized( + HttpRequest request, + ServerSettings settings, + ManagementStore managementStore, + bool allowQueryToken, + bool allowPathToken = false) + { + if (string.IsNullOrWhiteSpace(settings.Token) && !managementStore.HasApiKeys) + { + return true; + } + + if (request.Headers.TryGetValue(ProtocolConstants.TokenHeader, out var headerValues) + && headerValues.Any(value => IsTokenAuthorized(value, settings, managementStore, updateApiKeyLastUsed: true))) + { + return true; + } + + // Bearer form for OpenAI-compatible clients (e.g. n8n's OpenAI nodes pointed + // at /clients/{id}/v1) that can send an API key but no custom headers. + if (request.Headers.TryGetValue("Authorization", out var authorizationValues) + && authorizationValues.Any(value => TryGetBearerToken(value, out var bearerToken) + && IsTokenAuthorized(bearerToken, settings, managementStore, updateApiKeyLastUsed: true))) + { + return true; + } + + if (allowPathToken + && TryGetPathToken(request.Path, out var pathToken, out _) + && IsTokenAuthorized(pathToken, settings, managementStore, updateApiKeyLastUsed: true)) + { + return true; + } + + return allowQueryToken + && request.Query.TryGetValue("token", out var queryValues) + && queryValues.Any(value => IsTokenAuthorized(value, settings, managementStore, updateApiKeyLastUsed: true)); + } + + public static bool TryRemovePathToken( + PathString path, + ServerSettings settings, + ManagementStore managementStore, + out PathString remainingPath) + { + remainingPath = path; + + if (!TryGetPathToken(path, out var pathToken, out var tokenRemainingPath) + || !IsTokenAuthorized(pathToken, settings, managementStore, updateApiKeyLastUsed: false)) + { + return false; + } + + remainingPath = string.IsNullOrEmpty(tokenRemainingPath.Value) + ? new PathString("/") + : tokenRemainingPath; + return true; + } + + public static bool IsOwnBearerValue(string? value, ServerSettings settings, ManagementStore managementStore) => + TryGetBearerToken(value, out var token) + && IsTokenAuthorized(token, settings, managementStore, updateApiKeyLastUsed: false); + + public static bool IsTokenAuthorized( + string? token, + ServerSettings settings, + ManagementStore managementStore, + bool updateApiKeyLastUsed) + { + if (string.IsNullOrWhiteSpace(token)) + { + return false; + } + + return !string.IsNullOrWhiteSpace(settings.Token) + && string.Equals(token, settings.Token, StringComparison.Ordinal) + || managementStore.IsApiKeyValid(token, updateApiKeyLastUsed); + } + + private static bool TryGetBearerToken(string? authorization, out string token) + { + token = ""; + + if (string.IsNullOrWhiteSpace(authorization) + || !authorization.StartsWith("Bearer ", StringComparison.OrdinalIgnoreCase)) + { + return false; + } + + token = authorization["Bearer ".Length..].Trim(); + return token.Length > 0; + } + + private static bool TryGetPathToken(PathString path, out string pathToken, out PathString remainingPath) + { + pathToken = ""; + remainingPath = PathString.Empty; + + if (!path.StartsWithSegments(PathTokenPrefix, out var pathAfterPrefix)) + { + return false; + } + + var value = pathAfterPrefix.Value ?? ""; + if (value.Length <= 1 || value[0] != '/') + { + return false; + } + + var nextSlash = value.IndexOf('/', 1); + pathToken = nextSlash < 0 + ? value[1..] + : value[1..nextSlash]; + + if (string.IsNullOrEmpty(pathToken)) + { + return false; + } + + remainingPath = nextSlash < 0 + ? PathString.Empty + : new PathString(value[nextSlash..]); + return true; + } +} diff --git a/src/ReverseLlama.Server/TunnelConnection.cs b/src/ReverseLlama.Server/TunnelConnection.cs new file mode 100644 index 0000000..489ed8c --- /dev/null +++ b/src/ReverseLlama.Server/TunnelConnection.cs @@ -0,0 +1,333 @@ +using System.Collections.Concurrent; +using System.Net.WebSockets; +using ReverseLlama.Protocol; + +namespace ReverseLlama.Server; + +internal sealed class TunnelConnection +{ + private readonly ConcurrentDictionary _commands = new(); + private readonly object _modelsLock = new(); + private readonly ConcurrentDictionary _pending = new(); + private readonly SemaphoreSlim _sendLock = new(1, 1); + private readonly WebSocket _socket; + private readonly ILogger _logger; + private string[] _activeModels = []; + private string[] _models = []; + private DateTimeOffset? _modelsUpdatedAt; + + public TunnelConnection(string clientId, WebSocket socket, ILogger logger) + { + ClientId = clientId; + _socket = socket; + _logger = logger; + } + + public string ClientId { get; } + + public string Id { get; } = Guid.NewGuid().ToString("n"); + + public bool IsOpen => _socket.State == WebSocketState.Open; + + public int PendingRequestCount => _pending.Count; + + public IReadOnlyList Models + { + get + { + lock (_modelsLock) + { + return _models; + } + } + } + + public IReadOnlyList ActiveModels + { + get + { + lock (_modelsLock) + { + return _activeModels; + } + } + } + + public DateTimeOffset? ModelsUpdatedAt + { + get + { + lock (_modelsLock) + { + return _modelsUpdatedAt; + } + } + } + + public PendingProxyRequest RegisterPending(string requestId) + { + var pending = new PendingProxyRequest(); + + if (!_pending.TryAdd(requestId, pending)) + { + throw new InvalidOperationException($"Request id {requestId} is already registered."); + } + + return pending; + } + + public void RemovePending(string requestId) + { + _pending.TryRemove(requestId, out _); + } + + public bool HasModel(string model) + { + if (string.IsNullOrWhiteSpace(model)) + { + return false; + } + + var requested = model.Trim(); + + foreach (var available in Models.Concat(ActiveModels)) + { + if (ModelNamesMatch(requested, available)) + { + return true; + } + } + + return false; + } + + public void UpdateModels(IEnumerable models, IEnumerable activeModels) + { + var snapshot = models + .Where(model => !string.IsNullOrWhiteSpace(model)) + .Select(model => model.Trim()) + .Distinct(StringComparer.OrdinalIgnoreCase) + .OrderBy(model => model, StringComparer.OrdinalIgnoreCase) + .ToArray(); + var activeSnapshot = activeModels + .Where(model => !string.IsNullOrWhiteSpace(model)) + .Select(model => model.Trim()) + .Distinct(StringComparer.OrdinalIgnoreCase) + .OrderBy(model => model, StringComparer.OrdinalIgnoreCase) + .ToArray(); + + lock (_modelsLock) + { + _models = snapshot; + _activeModels = activeSnapshot; + _modelsUpdatedAt = DateTimeOffset.UtcNow; + } + } + + public async Task SendModelCommandAsync( + string command, + string model, + string? payloadJson, + TimeSpan timeout, + CancellationToken cancellationToken) + { + if (string.IsNullOrWhiteSpace(command)) + { + throw new ArgumentException("Command is required.", nameof(command)); + } + + if (string.IsNullOrWhiteSpace(model)) + { + throw new ArgumentException("Model is required.", nameof(model)); + } + + var requestId = Guid.NewGuid().ToString("n"); + var pending = new PendingCommand(); + + if (!_commands.TryAdd(requestId, pending)) + { + throw new InvalidOperationException($"Command id {requestId} is already registered."); + } + + using var timeoutCts = CancellationTokenSource.CreateLinkedTokenSource(cancellationToken); + timeoutCts.CancelAfter(timeout); + + try + { + await SendAsync( + new TunnelMessage + { + Type = TunnelMessageTypes.ModelCommand, + RequestId = requestId, + Command = command, + Model = model, + PayloadJson = payloadJson + }, + timeoutCts.Token); + + return await pending.WaitAsync(timeoutCts.Token); + } + finally + { + _commands.TryRemove(requestId, out _); + } + } + + public async Task SendAsync(TunnelMessage message, CancellationToken cancellationToken) + { + await _sendLock.WaitAsync(cancellationToken); + try + { + if (_socket.State != WebSocketState.Open) + { + throw new InvalidOperationException("The tunnel client is not connected."); + } + + await WebSocketMessageTransport.SendAsync(_socket, message, cancellationToken); + } + finally + { + _sendLock.Release(); + } + } + + public async Task RunReceiveLoopAsync(CancellationToken cancellationToken) + { + try + { + while (_socket.State == WebSocketState.Open && !cancellationToken.IsCancellationRequested) + { + var message = await WebSocketMessageTransport.ReceiveAsync(_socket, cancellationToken); + if (message is null) + { + break; + } + + Dispatch(message); + } + } + catch (OperationCanceledException) when (cancellationToken.IsCancellationRequested) + { + } + catch (Exception exception) + { + _logger.LogWarning(exception, "Tunnel receive loop failed."); + } + finally + { + FailAll("Tunnel client disconnected."); + } + } + + public async Task CloseAsync(string reason, string? closeDescription = null) + { + try + { + if (_socket.State is WebSocketState.Open or WebSocketState.CloseReceived) + { + await _socket.CloseAsync(WebSocketCloseStatus.NormalClosure, closeDescription ?? reason, CancellationToken.None); + } + } + catch + { + } + finally + { + FailAll(reason); + } + } + + private void Dispatch(TunnelMessage message) + { + if (message.Type == TunnelMessageTypes.ModelSnapshot) + { + UpdateModels(message.Models, message.ActiveModels); + _logger.LogInformation( + "Tunnel client {ClientId} reported {ModelCount} listed and {ActiveModelCount} active model(s).", + ClientId, + Models.Count, + ActiveModels.Count); + return; + } + + if (message.Type == TunnelMessageTypes.ModelCommandResult) + { + if (_commands.TryRemove(message.RequestId, out var pendingCommand)) + { + if (!string.IsNullOrWhiteSpace(message.Error)) + { + pendingCommand.Fail(message.Error); + } + else + { + pendingCommand.Complete(message); + } + } + + return; + } + + if (string.IsNullOrWhiteSpace(message.RequestId)) + { + _logger.LogDebug("Ignoring tunnel message without a request id: {MessageType}", message.Type); + return; + } + + if (!_pending.TryGetValue(message.RequestId, out var pending)) + { + _logger.LogDebug("Ignoring tunnel message for unknown request {RequestId}: {MessageType}", message.RequestId, message.Type); + return; + } + + switch (message.Type) + { + case TunnelMessageTypes.HttpResponseHeaders: + pending.SetResponseHeaders(message); + break; + + case TunnelMessageTypes.HttpResponseBody: + pending.AddBody(message.Body ?? []); + break; + + case TunnelMessageTypes.HttpResponseComplete: + pending.Complete(); + break; + + case TunnelMessageTypes.Error: + pending.Fail(message.Error ?? "The tunnel client reported an error."); + break; + + default: + _logger.LogDebug("Ignoring unsupported tunnel message type from client: {MessageType}", message.Type); + break; + } + } + + private static bool ModelNamesMatch(string requested, string available) => + string.Equals(requested, available, StringComparison.OrdinalIgnoreCase) + || string.Equals(StripLatestTag(requested), StripLatestTag(available), StringComparison.OrdinalIgnoreCase); + + private static string StripLatestTag(string model) => + model.EndsWith(":latest", StringComparison.OrdinalIgnoreCase) + ? model[..^":latest".Length] + : model; + + private void FailAll(string reason) + { + foreach (var pair in _pending.ToArray()) + { + if (_pending.TryRemove(pair.Key, out var pending)) + { + pending.Fail(reason); + } + } + + foreach (var pair in _commands.ToArray()) + { + if (_commands.TryRemove(pair.Key, out var pending)) + { + pending.Fail(reason); + } + } + } +} diff --git a/src/ReverseLlama.Server/TunnelHub.cs b/src/ReverseLlama.Server/TunnelHub.cs new file mode 100644 index 0000000..dd5e909 --- /dev/null +++ b/src/ReverseLlama.Server/TunnelHub.cs @@ -0,0 +1,148 @@ +using System.Collections.Concurrent; +using System.Net.WebSockets; +using ReverseLlama.Protocol; + +namespace ReverseLlama.Server; + +internal sealed class TunnelHub +{ + private readonly ConcurrentDictionary _connections = new(StringComparer.OrdinalIgnoreCase); + private readonly ILogger _logger; + private readonly ILoggerFactory _loggerFactory; + + public TunnelHub(ILogger logger, ILoggerFactory loggerFactory) + { + _logger = logger; + _loggerFactory = loggerFactory; + } + + public bool HasClient => _connections.Values.Any(connection => connection.IsOpen); + + public int PendingRequestCount => _connections.Values.Sum(connection => connection.PendingRequestCount); + + public TunnelConnection? Get(string clientId) => + _connections.TryGetValue(clientId, out var connection) && connection.IsOpen ? connection : null; + + public TunnelConnection? SelectBest(string? model, Func? isAvailable = null) + { + var allOpen = _connections.Values + .Where(connection => connection.IsOpen) + .Where(connection => isAvailable?.Invoke(connection.ClientId) ?? true) + .ToList(); + + if (allOpen.Count == 0) + { + return null; + } + + if (!string.IsNullOrWhiteSpace(model)) + { + var withModel = allOpen.Where(connection => connection.HasModel(model)).ToList(); + if (withModel.Count > 0) + { + return withModel + .OrderBy(connection => connection.PendingRequestCount) + .ThenBy(connection => connection.ClientId, StringComparer.OrdinalIgnoreCase) + .First(); + } + } + + return allOpen + .OrderBy(connection => connection.PendingRequestCount) + .ThenBy(connection => connection.ClientId, StringComparer.OrdinalIgnoreCase) + .First(); + } + + /// The only open connection, or null when zero or more than one client is connected. + public TunnelConnection? Single + { + get + { + TunnelConnection? single = null; + + foreach (var connection in _connections.Values) + { + if (!connection.IsOpen) + { + continue; + } + + if (single is not null) + { + return null; + } + + single = connection; + } + + return single; + } + } + + public IReadOnlyList ClientSnapshots => + _connections.Values + .Where(connection => connection.IsOpen) + .OrderBy(connection => connection.ClientId, StringComparer.OrdinalIgnoreCase) + .Select(connection => new TunnelClientSnapshot( + connection.ClientId, + connection.PendingRequestCount, + connection.Models, + connection.ActiveModels, + connection.ModelsUpdatedAt)) + .ToList(); + + public IReadOnlyList ClientsSnapshot => + ClientSnapshots + .Select(client => (object)new + { + id = client.Id, + pendingRequests = client.PendingRequests, + models = client.Models, + activeModels = client.ActiveModels, + modelsUpdatedAt = client.ModelsUpdatedAt + }) + .ToList(); + + public async Task AcceptAsync(string clientId, WebSocket socket, CancellationToken cancellationToken) + { + var connection = new TunnelConnection(clientId, socket, _loggerFactory.CreateLogger()); + + TunnelConnection? previous = null; + _connections.AddOrUpdate( + clientId, + connection, + (_, existing) => + { + previous = existing; + return connection; + }); + + if (previous is not null) + { + _logger.LogInformation( + "Replacing existing tunnel client {ClientId} ({ConnectionId}) with {NewConnectionId}.", + clientId, previous.Id, connection.Id); + await previous.CloseAsync("A newer tunnel client connected.", ProtocolConstants.ReplacedCloseDescription); + } + + _logger.LogInformation("Tunnel client {ClientId} ({ConnectionId}) connected.", clientId, connection.Id); + + try + { + await connection.RunReceiveLoopAsync(cancellationToken); + } + finally + { + _connections.TryRemove(new KeyValuePair(clientId, connection)); + await connection.CloseAsync("Tunnel closed."); + _logger.LogInformation("Tunnel client {ClientId} ({ConnectionId}) disconnected.", clientId, connection.Id); + } + } +} + +internal sealed record TunnelClientSnapshot( + string Id, + int PendingRequests, + IReadOnlyList Models, + IReadOnlyList ActiveModels, + DateTimeOffset? ModelsUpdatedAt); diff --git a/src/ReverseLlama.Server/appsettings.Example.json b/src/ReverseLlama.Server/appsettings.Example.json new file mode 100644 index 0000000..d09ec67 --- /dev/null +++ b/src/ReverseLlama.Server/appsettings.Example.json @@ -0,0 +1,19 @@ +{ + "Logging": { + "LogLevel": { + "Default": "Information", + "Microsoft.AspNetCore": "Warning" + } + }, + "ConnectionStrings": { + "ElmahConnection": "Data Source=localhost;Initial Catalog=elmah;Persist Security Info=True;User ID=elmah;Password=elmah;" + }, + "Authentication": { + "Keycloak": { + "Authority": "http://your-keycloak-server/realms/master", + "ClientId": "ReverseLlama", + "ClientSecret": "YOUR-Client-SECRET-GOES-HERE-AND-YES-ITS-VERY-LONG", + "RequireHttpsMetadata": false + } + } +} diff --git a/src/ReverseLlama.Server/appsettings.json b/src/ReverseLlama.Server/appsettings.json new file mode 100644 index 0000000..10f68b8 --- /dev/null +++ b/src/ReverseLlama.Server/appsettings.json @@ -0,0 +1,9 @@ +{ + "Logging": { + "LogLevel": { + "Default": "Information", + "Microsoft.AspNetCore": "Warning" + } + }, + "AllowedHosts": "*" +} diff --git a/src/ReverseLlama.Server/wwwroot/admin/app.css b/src/ReverseLlama.Server/wwwroot/admin/app.css new file mode 100644 index 0000000..9e03e20 --- /dev/null +++ b/src/ReverseLlama.Server/wwwroot/admin/app.css @@ -0,0 +1,496 @@ +:root { + color-scheme: light; + --bg: #f5f7f9; + --panel: #ffffff; + --panel-alt: #eef2f6; + --text: #18212f; + --muted: #667085; + --line: #d9e0e7; + --line-strong: #b8c3cf; + --blue: #2f6fed; + --blue-dark: #1f4fb7; + --green: #117a55; + --red: #b42318; + --amber: #9a6500; + --shadow: 0 12px 30px rgba(24, 33, 47, 0.08); +} + +* { + box-sizing: border-box; +} + +body { + margin: 0; + min-height: 100vh; + background: var(--bg); + color: var(--text); + font-family: Inter, "Segoe UI", system-ui, -apple-system, BlinkMacSystemFont, sans-serif; + font-size: 14px; +} + +button, +input, +select, +textarea { + font: inherit; +} + +.app-shell { + display: grid; + grid-template-columns: 248px minmax(0, 1fr); + min-height: 100vh; +} + +.sidebar { + display: flex; + flex-direction: column; + gap: 20px; + padding: 20px 16px; + background: #151b24; + color: #f8fafc; +} + +.brand { + display: flex; + align-items: center; + gap: 10px; + min-width: 0; +} + +.brand-mark { + display: grid; + place-items: center; + width: 34px; + height: 34px; + border: 1px solid #4c5d74; + border-radius: 6px; + background: #223044; + color: #c8f2df; + font-weight: 800; + letter-spacing: 0; +} + +.brand strong, +.brand small { + display: block; + overflow-wrap: anywhere; +} + +.brand small { + color: #aeb8c6; + margin-top: 2px; +} + +.nav { + display: grid; + gap: 6px; +} + +.nav a { + color: #d8dee8; + text-decoration: none; + padding: 9px 10px; + border-radius: 6px; +} + +.nav a.active, +.nav a:hover { + background: #263449; + color: #ffffff; +} + +.sidebar-meta { + margin-top: auto; + padding-top: 16px; + border-top: 1px solid #344154; + color: #b8c3cf; + font-size: 12px; + line-height: 1.5; + overflow-wrap: anywhere; +} + +.main { + min-width: 0; + padding: 22px; +} + +.topbar { + display: flex; + align-items: flex-start; + justify-content: space-between; + gap: 16px; + margin-bottom: 16px; +} + +.topbar h1 { + margin: 0; + font-size: 24px; + line-height: 1.2; + letter-spacing: 0; +} + +.topbar p { + margin: 6px 0 0; + color: var(--muted); +} + +.topbar-actions { + display: flex; + align-items: center; + gap: 8px; + flex-wrap: wrap; +} + +.content { + display: grid; + gap: 14px; +} + +.toolbar { + display: flex; + align-items: end; + justify-content: space-between; + gap: 12px; + flex-wrap: wrap; +} + +.form-row { + display: flex; + align-items: end; + gap: 8px; + flex-wrap: wrap; +} + +.field { + display: grid; + gap: 5px; +} + +.field label { + color: var(--muted); + font-size: 12px; + font-weight: 600; +} + +.input, +.select, +.textarea { + width: 100%; + min-height: 36px; + border: 1px solid var(--line-strong); + border-radius: 6px; + background: #ffffff; + color: var(--text); + padding: 8px 10px; +} + +.textarea { + min-height: 72px; + resize: vertical; +} + +.input:focus, +.select:focus, +.textarea:focus { + outline: 2px solid rgba(47, 111, 237, 0.25); + border-color: var(--blue); +} + +.button { + min-height: 36px; + border: 1px solid var(--blue); + border-radius: 6px; + background: var(--blue); + color: #ffffff; + padding: 8px 12px; + font-weight: 700; + cursor: pointer; + white-space: nowrap; +} + +.button:hover { + background: var(--blue-dark); + border-color: var(--blue-dark); +} + +.button.secondary { + background: #ffffff; + border-color: var(--line-strong); + color: var(--text); +} + +.button.secondary:hover { + background: var(--panel-alt); +} + +.button.danger { + background: #ffffff; + border-color: #e5aaa4; + color: var(--red); +} + +.button.danger:hover { + background: #fff1f0; +} + +.button.warning { + background: #ffffff; + border-color: #e4c37c; + color: var(--amber); +} + +.button.warning:hover { + background: #fff8e5; +} + +.button:disabled { + cursor: not-allowed; + opacity: 0.55; +} + +.panel { + background: var(--panel); + border: 1px solid var(--line); + border-radius: 8px; + box-shadow: var(--shadow); +} + +.panel-header { + display: flex; + align-items: center; + justify-content: space-between; + gap: 10px; + padding: 14px 16px; + border-bottom: 1px solid var(--line); +} + +.panel-header h2 { + margin: 0; + font-size: 15px; + letter-spacing: 0; +} + +.panel-body { + padding: 14px 16px; +} + +.table-wrap { + overflow-x: auto; +} + +table { + width: 100%; + border-collapse: collapse; + min-width: 760px; +} + +th, +td { + padding: 10px 12px; + border-bottom: 1px solid var(--line); + text-align: left; + vertical-align: top; +} + +th { + color: var(--muted); + font-size: 12px; + font-weight: 800; + text-transform: uppercase; +} + +tr:last-child td { + border-bottom: 0; +} + +.cell-main { + font-weight: 800; + overflow-wrap: anywhere; +} + +.cell-sub { + margin-top: 3px; + color: var(--muted); + font-size: 12px; + overflow-wrap: anywhere; +} + +.actions { + display: flex; + gap: 6px; + flex-wrap: wrap; +} + +.badge-row { + display: flex; + gap: 5px; + flex-wrap: wrap; +} + +.badge { + display: inline-flex; + align-items: center; + min-height: 24px; + border: 1px solid var(--line); + border-radius: 999px; + padding: 3px 8px; + background: var(--panel-alt); + color: var(--text); + font-size: 12px; + font-weight: 700; + overflow-wrap: anywhere; +} + +.badge.good { + border-color: #9bd7bd; + background: #eaf8f1; + color: var(--green); +} + +.badge.bad { + border-color: #efb2ad; + background: #fff1f0; + color: var(--red); +} + +.badge.warn { + border-color: #ead09a; + background: #fff8e5; + color: var(--amber); +} + +.metric-grid { + display: grid; + grid-template-columns: repeat(4, minmax(120px, 1fr)); + gap: 10px; +} + +.metric { + padding: 12px; + border: 1px solid var(--line); + border-radius: 8px; + background: #ffffff; +} + +.metric strong { + display: block; + font-size: 22px; + letter-spacing: 0; +} + +.metric span { + display: block; + margin-top: 4px; + color: var(--muted); + font-size: 12px; +} + +.notice { + margin-bottom: 12px; + padding: 10px 12px; + border: 1px solid #bad2ff; + border-radius: 8px; + background: #edf4ff; + color: #163d85; +} + +.notice.error { + border-color: #efb2ad; + background: #fff1f0; + color: var(--red); +} + +.empty { + padding: 28px 16px; + color: var(--muted); + text-align: center; +} + +.pre { + max-height: 460px; + overflow: auto; + padding: 12px; + border: 1px solid var(--line); + border-radius: 8px; + background: #101820; + color: #e9eef5; + font-family: "Cascadia Mono", Consolas, monospace; + font-size: 12px; + line-height: 1.45; + white-space: pre-wrap; + overflow-wrap: anywhere; +} + +.new-key { + display: grid; + gap: 8px; + padding: 12px; + border: 1px solid #9bd7bd; + border-radius: 8px; + background: #eaf8f1; +} + +.new-key code { + display: block; + padding: 10px; + border: 1px solid #9bd7bd; + border-radius: 6px; + background: #ffffff; + color: var(--text); + overflow-wrap: anywhere; +} + +@media (max-width: 860px) { + .app-shell { + grid-template-columns: 1fr; + } + + .sidebar { + position: sticky; + top: 0; + z-index: 2; + display: grid; + grid-template-columns: 1fr; + gap: 12px; + } + + .nav { + grid-template-columns: repeat(3, 1fr); + } + + .nav a { + text-align: center; + } + + .sidebar-meta { + display: none; + } + + .main { + padding: 16px; + } + + .topbar { + display: grid; + } + + .metric-grid { + grid-template-columns: repeat(2, minmax(0, 1fr)); + } +} + +@media (max-width: 520px) { + .nav { + grid-template-columns: 1fr; + } + + .metric-grid { + grid-template-columns: 1fr; + } + + .button, + .input, + .select { + width: 100%; + } +} diff --git a/src/ReverseLlama.Server/wwwroot/admin/app.js b/src/ReverseLlama.Server/wwwroot/admin/app.js new file mode 100644 index 0000000..3db8bc6 --- /dev/null +++ b/src/ReverseLlama.Server/wwwroot/admin/app.js @@ -0,0 +1,766 @@ +const state = { + summary: null, + detail: null, + newKey: null, + loading: false +}; + +const content = document.getElementById("content"); +const notice = document.getElementById("notice"); +const pageTitle = document.getElementById("pageTitle"); +const pageSubtitle = document.getElementById("pageSubtitle"); +const sidebarMeta = document.getElementById("sidebarMeta"); + +document.getElementById("refreshButton").addEventListener("click", () => refresh(true)); +window.addEventListener("hashchange", () => renderRoute()); + +content.addEventListener("click", async (event) => { + const button = event.target.closest("button[data-action]"); + if (!button) { + return; + } + + const { action, clientId, model, keyId } = button.dataset; + + try { + setBusy(button, true); + + if (action === "disable-hour") { + await api(`/clients/${encodeURIComponent(clientId)}/disable`, { + method: "POST", + body: { mode: "duration", durationMinutes: 60 } + }); + setNotice(`Disabled ${clientId} for one hour.`); + await refresh(); + } + + if (action === "disable-manual") { + await api(`/clients/${encodeURIComponent(clientId)}/disable`, { + method: "POST", + body: { mode: "manual" } + }); + setNotice(`Disabled ${clientId}.`); + await refresh(); + } + + if (action === "enable-client") { + await api(`/clients/${encodeURIComponent(clientId)}/enable`, { method: "POST" }); + setNotice(`Enabled ${clientId}.`); + await refresh(); + } + + if (action === "model-command") { + await runModelCommand(clientId, model, button.dataset.modelAction); + } + + if (action === "delete-key") { + if (!confirm("Delete this API key?")) { + return; + } + + await api(`/api-keys/${encodeURIComponent(keyId)}`, { method: "DELETE" }); + setNotice("API key deleted."); + await refresh(); + } + + if (action === "copy-key") { + await navigator.clipboard.writeText(button.dataset.key); + setNotice("API key copied."); + } + } catch (error) { + setNotice(error.message, true); + } finally { + setBusy(button, false); + } +}); + +content.addEventListener("submit", async (event) => { + const form = event.target; + if (!(form instanceof HTMLFormElement)) { + return; + } + + event.preventDefault(); + const data = Object.fromEntries(new FormData(form).entries()); + const submit = form.querySelector("button[type=submit]"); + + try { + setBusy(submit, true); + + if (form.dataset.form === "model-action") { + await runModelCommand(data.clientId, data.model, data.action); + form.reset(); + } + + if (form.dataset.form === "model-detail") { + await loadModelDetail(data.model, data.clientId); + } + + if (form.dataset.form === "api-key") { + state.newKey = await api("/api-keys", { + method: "POST", + body: { name: data.name } + }); + setNotice("API key created."); + await refresh(); + } + } catch (error) { + setNotice(error.message, true); + } finally { + setBusy(submit, false); + } +}); + +async function boot() { + await refresh(); + setInterval(() => refresh(), 15000); +} + +async function refresh(showNotice = false, render = true) { + state.loading = true; + + try { + state.summary = await api("/summary"); + updateShell(); + if (render) { + await renderRoute(); + } + + if (showNotice) { + setNotice("Refreshed."); + } + } catch (error) { + setNotice(error.message, true); + content.innerHTML = `
${escapeHtml(error.message)}
`; + } finally { + state.loading = false; + } +} + +async function api(path, options = {}) { + const headers = options.headers ? { ...options.headers } : {}; + const init = { + method: options.method || "GET", + credentials: "same-origin", + headers + }; + + if (options.body !== undefined) { + headers["Content-Type"] = "application/json"; + init.body = JSON.stringify(options.body); + } + + const response = await fetch(`/api/admin${path}`, init); + const contentType = response.headers.get("content-type") || ""; + const body = contentType.includes("application/json") + ? await response.json() + : await response.text(); + + if (!response.ok) { + const message = body?.detail || body?.error || body?.title || response.statusText; + throw new Error(message); + } + + return body; +} + +async function renderRoute() { + const hash = (window.location.hash || "#clients").slice(1); + const [view, encodedModel] = hash.split("/"); + + document.querySelectorAll("[data-nav]").forEach((link) => { + link.classList.toggle("active", link.dataset.nav === view); + }); + + if (view === "models" && encodedModel) { + const model = decodeURIComponent(encodedModel); + await loadModelDetail(model); + return; + } + + state.detail = null; + + if (view === "models") { + renderModels(); + return; + } + + if (view === "api-keys") { + renderApiKeys(); + return; + } + + renderClients(); +} + +function updateShell() { + const summary = state.summary; + if (!summary) { + return; + } + + sidebarMeta.innerHTML = ` +
${escapeHtml(summary.user?.name || "Signed in")}
+
${summary.clients.length} clients
+
${summary.models.length} models
+
${formatDate(summary.generatedAtUtc)}
+ `; +} + +function renderClients() { + const clients = state.summary?.clients || []; + pageTitle.textContent = "Clients"; + pageSubtitle.textContent = "Connected tunnel clients, request counts, and forwarding controls."; + + content.innerHTML = ` +
+
+

Clients

+ ${clients.length} total +
+
+ ${clients.length ? clientsTable(clients) : emptyState("No clients have connected yet.")} +
+
+ `; +} + +function clientsTable(clients) { + const rows = clients.map((client) => ` + + +
${escapeHtml(client.id)}
+
${client.connected ? "Connected" : "Offline"}${client.modelsUpdatedAt ? `, models ${formatDate(client.modelsUpdatedAt)}` : ""}
+ + +
+ ${client.connected ? badge("Connected", "good") : badge("Offline", "")} + ${client.disabled ? badge(client.disabledManually ? "Disabled manual" : "Disabled timed", "bad") : badge("Enabled", "good")} +
+ ${client.disabled ? `
${escapeHtml(disabledText(client))}
` : ""} + + ${number(client.pendingRequests)} + +
${number(client.requestStats.total)}
+
${number(client.requestStats.last10Minutes)} in 10m, ${number(client.requestStats.lastHour)} in 1h
+ + ${modelBadges(client.models)} + ${modelBadges(client.activeModels)} + +
+ + + +
+ + + `).join(""); + + return ` + + + + + + + + + + + + + ${rows} +
ClientStatusPendingRequestsListed modelsActive modelsActions
+ `; +} + +function renderModels() { + const models = state.summary?.models || []; + const clients = connectedClients(); + pageTitle.textContent = "Models"; + pageSubtitle.textContent = "Listed and active models, recent request volume, and model operations."; + + content.innerHTML = ` +
+
+

Run model action

+
+
+ ${modelActionForm(clients)} +
+
+
+
+

Models

+ ${models.length} total +
+
+ ${models.length ? modelsTable(models) : emptyState("No models have been reported yet.")} +
+
+ `; +} + +function modelActionForm(clients, selectedModel = "", selectedClient = "") { + return ` +
+
+ + +
+
+ + +
+
+ + +
+ +
+ `; +} + +function modelsTable(models) { + const rows = models.map((model) => ` + + +
${escapeHtml(model.name)}
+
${number(model.metrics.totalRequests)} total requests
+ + ${modelBadges(model.listedClients)} + ${modelBadges(model.activeClients)} + +
${number(model.metrics.requestsLast10Minutes)}
+
${number(model.metrics.requestsLastHour)} in last hour
+ + +
${number(model.metrics.tokensLast10Minutes)}
+
${number(model.metrics.tokensLastHour)} in last hour
+ + + Details + + + `).join(""); + + return ` + + + + + + + + + + + + ${rows} +
ModelListed onActive onRequests 10mTokens 10m
+ `; +} + +async function loadModelDetail(model, clientId) { + pageTitle.textContent = "Model Detail"; + pageSubtitle.textContent = model; + content.innerHTML = `
Loading model detail...
`; + + const query = new URLSearchParams({ model }); + if (clientId) { + query.set("clientId", clientId); + } + + state.detail = await api(`/models/detail?${query.toString()}`); + renderModelDetail(); +} + +function renderModelDetail() { + const detail = state.detail; + const clients = connectedClients(); + const model = detail.model; + const selectedClient = detail.selectedClientId || clients[0]?.id || ""; + const metrics = detail.metrics; + const showBody = detail.show?.body === undefined ? null : detail.show.body; + + pageTitle.textContent = "Model Detail"; + pageSubtitle.textContent = model; + + content.innerHTML = ` + +
+ ${metric(number(metrics.requestsLast10Minutes), "Requests in 10 minutes")} + ${metric(number(metrics.requestsLastHour), "Requests in 1 hour")} + ${metric(number(metrics.tokensLast10Minutes), "Tokens in 10 minutes")} + ${metric(number(metrics.tokensLastHour), "Tokens in 1 hour")} +
+
+
+

Placement

+
+
+
+ ${badge("Listed", detail.listedClients.length ? "good" : "")} + ${modelBadges(detail.listedClients)} +
+
+ ${badge("Active", detail.activeClients.length ? "good" : "warn")} + ${modelBadges(detail.activeClients)} +
+
+
+
+
+

Client action

+
+
+
+ +
+ + +
+ +
+
+ + + +
+
+
+
+
+

Ollama show response

+ ${detail.show ? badge(detail.show.ok ? "OK" : `HTTP ${detail.show.statusCode}`, detail.show.ok ? "good" : "bad") : ""} +
+
+ ${detail.show ? `
${escapeHtml(formatJson(showBody))}
` : emptyState("No connected client was available for details.")} +
+
+ `; +} + +async function runModelCommand(clientId, model, action) { + if (!clientId || !model || !action) { + throw new Error("Client, model, and action are required."); + } + + if (action === "remove" && !confirm(`Remove ${model} from ${clientId}?`)) { + return; + } + + const result = await api("/models/actions", { + method: "POST", + body: { clientId, model, action } + }); + + if (!result.ok) { + throw new Error(modelActionError(result)); + } + + const detail = modelActionResultDetail(result); + const completed = `${capitalize(action)} completed for ${model} on ${clientId}${detail ? ` (${detail})` : ""}.`; + + if (action === "load" || action === "unload") { + const shouldBeActive = action === "load"; + setNotice(`${completed} Waiting for active model snapshot.`); + + if (await waitForModelActiveState(clientId, model, shouldBeActive)) { + setNotice(completed); + return; + } + + setNotice(`${completed} The active model snapshot has not reflected the change yet.`); + return; + } + + setNotice(completed); + await refreshAfterModelCommand(model, clientId); +} + +function renderApiKeys() { + const keys = state.summary?.apiKeys || []; + pageTitle.textContent = "API Keys"; + pageSubtitle.textContent = "Keys accepted by the proxy token header, bearer auth, query token, and token path."; + + content.innerHTML = ` + ${state.newKey ? newKeyPanel(state.newKey) : ""} +
+
+

Create API key

+
+
+
+
+ + +
+ +
+
+
+
+
+

API keys

+ ${keys.length} total +
+
+ ${keys.length ? apiKeysTable(keys) : emptyState("No API keys have been created.")} +
+
+ `; +} + +function newKeyPanel(key) { + return ` +
+ New API key + ${escapeHtml(key.key)} +
+ +
+
+ `; +} + +function apiKeysTable(keys) { + const rows = keys.map((key) => ` + + +
${escapeHtml(key.name)}
+
${escapeHtml(key.keyPrefix)}...
+ + ${formatDate(key.createdAtUtc)} + ${key.lastUsedUtc ? formatDate(key.lastUsedUtc) : "Never"} + + + + + `).join(""); + + return ` + + + + + + + + + + ${rows} +
NameCreatedLast used
+ `; +} + +function connectedClients() { + return (state.summary?.clients || []).filter((client) => client.connected); +} + +async function waitForModelActiveState(clientId, model, shouldBeActive) { + const deadline = Date.now() + 15000; + + while (Date.now() <= deadline) { + await refreshAfterModelCommand(model, clientId); + + const client = (state.summary?.clients || []) + .find((item) => sameText(item.id, clientId)); + + if (client && modelListContains(client.activeModels, model) === shouldBeActive) { + return true; + } + + await delay(1000); + } + + return false; +} + +async function refreshAfterModelCommand(model, clientId) { + await refresh(false, false); + + const hash = window.location.hash || ""; + if (hash.startsWith(`#models/${encodeURIComponent(model)}`)) { + await loadModelDetail(model, clientId); + return; + } + + await renderRoute(); +} + +function modelActionResultDetail(result) { + const body = result?.body; + + if (!body || typeof body === "string") { + return body || ""; + } + + if (body.status) { + return body.status; + } + + if (body.done_reason) { + return `done: ${body.done_reason}`; + } + + if (body.done) { + return "done"; + } + + return ""; +} + +function modelActionError(result) { + const body = result?.body; + + if (typeof body === "string" && body) { + return body; + } + + if (body?.error) { + return body.error; + } + + if (body?.message) { + return body.message; + } + + return `Model action failed with HTTP ${result.statusCode}.`; +} + +function modelListContains(models, model) { + return (models || []).some((item) => sameModelName(item, model)); +} + +function sameModelName(left, right) { + return stripLatestTag(left).toLowerCase() === stripLatestTag(right).toLowerCase(); +} + +function sameText(left, right) { + return String(left || "").trim().toLowerCase() === String(right || "").trim().toLowerCase(); +} + +function stripLatestTag(model) { + const value = String(model || "").trim(); + return value.toLowerCase().endsWith(":latest") ? value.slice(0, -":latest".length) : value; +} + +function delay(milliseconds) { + return new Promise((resolve) => setTimeout(resolve, milliseconds)); +} + +function clientOptions(clients, selectedClient) { + if (!clients.length) { + return ``; + } + + return clients.map((client) => ` + + `).join(""); +} + +function modelBadges(items) { + if (!items || !items.length) { + return `None`; + } + + return `
${items.map((item) => badge(item, "")).join("")}
`; +} + +function badge(text, kind) { + return `${escapeHtml(text)}`; +} + +function metric(value, label) { + return ` +
+ ${escapeHtml(value)} + ${escapeHtml(label)} +
+ `; +} + +function emptyState(text) { + return `
${escapeHtml(text)}
`; +} + +function disabledText(client) { + if (client.disabledManually) { + return "Until enabled manually"; + } + + return client.disabledUntilUtc ? `Until ${formatDate(client.disabledUntilUtc)}` : "Disabled"; +} + +function formatDate(value) { + if (!value) { + return ""; + } + + return new Intl.DateTimeFormat(undefined, { + dateStyle: "short", + timeStyle: "medium" + }).format(new Date(value)); +} + +function formatJson(value) { + if (typeof value === "string") { + return value; + } + + return JSON.stringify(value, null, 2); +} + +function number(value) { + return new Intl.NumberFormat().format(value || 0); +} + +function capitalize(value) { + return value ? value[0].toUpperCase() + value.slice(1) : value; +} + +function setNotice(message, isError = false) { + notice.hidden = false; + notice.textContent = message; + notice.classList.toggle("error", isError); + + clearTimeout(setNotice.timer); + setNotice.timer = setTimeout(() => { + notice.hidden = true; + }, isError ? 7000 : 3500); +} + +function setBusy(element, busy) { + if (!element) { + return; + } + + element.disabled = busy; +} + +function escapeHtml(value) { + return String(value ?? "") + .replaceAll("&", "&") + .replaceAll("<", "<") + .replaceAll(">", ">") + .replaceAll('"', """) + .replaceAll("'", "'"); +} + +function escapeAttr(value) { + return escapeHtml(value); +} + +boot(); diff --git a/src/ReverseLlama.Server/wwwroot/admin/index.html b/src/ReverseLlama.Server/wwwroot/admin/index.html new file mode 100644 index 0000000..1273314 --- /dev/null +++ b/src/ReverseLlama.Server/wwwroot/admin/index.html @@ -0,0 +1,48 @@ + + + + + + ReverseLlama Admin + + + +
+ + +
+
+
+

Clients

+

Connected tunnel clients and forwarding controls.

+
+
+ +
+ +
+
+
+ + +
+
+
+ + + + diff --git a/tests/ReverseLlama.Client.Tests/PackageAuditTests.cs b/tests/ReverseLlama.Client.Tests/PackageAuditTests.cs new file mode 100644 index 0000000..29bba8a --- /dev/null +++ b/tests/ReverseLlama.Client.Tests/PackageAuditTests.cs @@ -0,0 +1,153 @@ +using System.Diagnostics; +using System.Text.Json; +using Xunit; + +namespace ReverseLlama.Client.Tests; + +public sealed class PackageAuditTests +{ + [Fact] + public async Task Solution_HasNoKnownVulnerablePackages() + { + var solutionPath = FindSolutionPath(); + var result = await RunDotnetPackageAuditAsync(solutionPath); + + Assert.True( + result.ExitCode == 0, + $"Package audit command failed with exit code {result.ExitCode}.{Environment.NewLine}{result.Output}{result.Error}"); + + using var document = JsonDocument.Parse(result.Output); + var findings = new List(); + CollectVulnerablePackages(document.RootElement, findings); + + Assert.True( + findings.Count == 0, + "Known vulnerable packages were found:" + + Environment.NewLine + + string.Join(Environment.NewLine, findings)); + } + + private static string FindSolutionPath() + { + var directory = new DirectoryInfo(AppContext.BaseDirectory); + while (directory is not null) + { + var solutionPath = Path.Combine(directory.FullName, "ReverseLlama.sln"); + if (File.Exists(solutionPath)) + { + return solutionPath; + } + + directory = directory.Parent; + } + + throw new InvalidOperationException("Could not find ReverseLlama.sln from the test output directory."); + } + + private static async Task RunDotnetPackageAuditAsync(string solutionPath) + { + var startInfo = new ProcessStartInfo + { + FileName = "dotnet", + WorkingDirectory = Path.GetDirectoryName(solutionPath)!, + RedirectStandardOutput = true, + RedirectStandardError = true, + UseShellExecute = false + }; + + startInfo.ArgumentList.Add("list"); + startInfo.ArgumentList.Add(solutionPath); + startInfo.ArgumentList.Add("package"); + startInfo.ArgumentList.Add("--vulnerable"); + startInfo.ArgumentList.Add("--include-transitive"); + startInfo.ArgumentList.Add("--format"); + startInfo.ArgumentList.Add("json"); + startInfo.Environment["DOTNET_CLI_TELEMETRY_OPTOUT"] = "1"; + + using var process = Process.Start(startInfo) + ?? throw new InvalidOperationException("Could not start dotnet package audit."); + + var outputTask = process.StandardOutput.ReadToEndAsync(); + var errorTask = process.StandardError.ReadToEndAsync(); + + using var timeout = new CancellationTokenSource(TimeSpan.FromMinutes(2)); + try + { + await process.WaitForExitAsync(timeout.Token); + } + catch (OperationCanceledException) + { + try + { + process.Kill(entireProcessTree: true); + } + catch + { + // Best effort; the assertion below will fail with the timeout message. + } + + throw new TimeoutException("dotnet package audit did not finish within 2 minutes."); + } + + return new CommandResult( + process.ExitCode, + await outputTask, + await errorTask); + } + + private static void CollectVulnerablePackages(JsonElement element, List findings) + { + if (element.ValueKind == JsonValueKind.Object) + { + if (element.TryGetProperty("vulnerabilities", out var vulnerabilities) + && vulnerabilities.ValueKind == JsonValueKind.Array + && vulnerabilities.GetArrayLength() > 0) + { + findings.Add(DescribePackageFinding(element, vulnerabilities)); + } + + foreach (var property in element.EnumerateObject()) + { + CollectVulnerablePackages(property.Value, findings); + } + } + else if (element.ValueKind == JsonValueKind.Array) + { + foreach (var item in element.EnumerateArray()) + { + CollectVulnerablePackages(item, findings); + } + } + } + + private static string DescribePackageFinding(JsonElement package, JsonElement vulnerabilities) + { + var id = package.TryGetProperty("id", out var idElement) + ? idElement.GetString() + : ""; + var version = package.TryGetProperty("resolvedVersion", out var versionElement) + ? versionElement.GetString() + : ""; + + var advisories = vulnerabilities + .EnumerateArray() + .Select(vulnerability => DescribeVulnerability(vulnerability)) + .ToArray(); + + return $"- {id} {version}: {string.Join(", ", advisories)}"; + } + + private static string DescribeVulnerability(JsonElement vulnerability) + { + var severity = vulnerability.TryGetProperty("severity", out var severityElement) + ? severityElement.GetString() + : "unknown severity"; + var advisoryUrl = vulnerability.TryGetProperty("advisoryUrl", out var advisoryElement) + ? advisoryElement.GetString() + : "unknown advisory"; + + return $"{severity} {advisoryUrl}"; + } + + private sealed record CommandResult(int ExitCode, string Output, string Error); +} diff --git a/tests/ReverseLlama.Client.Tests/ReverseLlama.Client.Tests.csproj b/tests/ReverseLlama.Client.Tests/ReverseLlama.Client.Tests.csproj new file mode 100644 index 0000000..579d3cd --- /dev/null +++ b/tests/ReverseLlama.Client.Tests/ReverseLlama.Client.Tests.csproj @@ -0,0 +1,24 @@ + + + + net10.0 + enable + enable + false + true + + + + + + + runtime; build; native; contentfiles; analyzers; buildtransitive + all + + + + + + + + diff --git a/tests/ReverseLlama.Client.Tests/TunnelClientModelTests.cs b/tests/ReverseLlama.Client.Tests/TunnelClientModelTests.cs new file mode 100644 index 0000000..7d81f33 --- /dev/null +++ b/tests/ReverseLlama.Client.Tests/TunnelClientModelTests.cs @@ -0,0 +1,62 @@ +using System.Text.Json; +using ReverseLlama.Client; +using Xunit; + +namespace ReverseLlama.Client.Tests; + +public sealed class TunnelClientModelTests +{ + private static readonly Uri Upstream = new("http://localhost:11434"); + + [Fact] + public void ExtractModelNames_ReadsActiveOllamaPsModelNames() + { + using var document = JsonDocument.Parse( + """ + { + "models": [ + { "name": "qwen3.5:0.8b", "model": "qwen3.5:0.8b" }, + { "name": "llama3.2:latest", "model": "llama3.2:latest" }, + { "name": " QWEN3.5:0.8B " } + ] + } + """); + + var models = TunnelClient.ExtractModelNames(document.RootElement); + + Assert.Equal(["llama3.2:latest", "qwen3.5:0.8b"], models); + } + + [Fact] + public async Task BuildModelCommandRequest_LoadUsesOllamaPreloadRequest() + { + using var request = TunnelClient.BuildModelCommandRequest(Upstream, "load", " qwen3.5:0.8b "); + var body = await request.Content!.ReadAsStringAsync(); + using var document = JsonDocument.Parse(body); + + Assert.Equal(HttpMethod.Post, request.Method); + Assert.Equal("http://localhost:11434/api/generate", request.RequestUri!.AbsoluteUri); + Assert.Equal("qwen3.5:0.8b", document.RootElement.GetProperty("model").GetString()); + Assert.False(document.RootElement.TryGetProperty("prompt", out _)); + Assert.False(document.RootElement.GetProperty("stream").GetBoolean()); + Assert.Equal(-1, document.RootElement.GetProperty("keep_alive").GetInt32()); + } + + [Theory] + [InlineData("load", -1)] + [InlineData("unload", 0)] + public async Task BuildEmbeddingModelCommandRequest_UsesOllamaEmbedWarmupRequest( + string command, + int expectedKeepAlive) + { + using var request = TunnelClient.BuildEmbeddingModelCommandRequest(Upstream, command, " bge-m3:latest "); + var body = await request.Content!.ReadAsStringAsync(); + using var document = JsonDocument.Parse(body); + + Assert.Equal(HttpMethod.Post, request.Method); + Assert.Equal("http://localhost:11434/api/embed", request.RequestUri!.AbsoluteUri); + Assert.Equal("bge-m3:latest", document.RootElement.GetProperty("model").GetString()); + Assert.Equal("ReverseLlama warmup", document.RootElement.GetProperty("input").GetString()); + Assert.Equal(expectedKeepAlive, document.RootElement.GetProperty("keep_alive").GetInt32()); + } +} diff --git a/tests/ReverseLlama.Client.Tests/UpstreamRequestTests.cs b/tests/ReverseLlama.Client.Tests/UpstreamRequestTests.cs new file mode 100644 index 0000000..bd5b5a0 --- /dev/null +++ b/tests/ReverseLlama.Client.Tests/UpstreamRequestTests.cs @@ -0,0 +1,35 @@ +using ReverseLlama.Client; +using Xunit; + +namespace ReverseLlama.Client.Tests; + +public sealed class UpstreamRequestTests +{ + private static readonly Uri Upstream = new("http://localhost:11434"); + + [Theory] + [InlineData("/api/tags", "http://localhost:11434/api/tags")] + [InlineData("/api/tags?model=llama3.1", "http://localhost:11434/api/tags?model=llama3.1")] + [InlineData("/api//tags", "http://localhost:11434/api//tags")] + public void BuildUpstreamUri_AcceptsOriginFormPaths(string pathAndQuery, string expected) + { + var uri = UpstreamRequest.BuildUpstreamUri(Upstream, pathAndQuery); + + Assert.Equal(expected, uri.AbsoluteUri); + } + + [Theory] + [InlineData("//169.254.169.254/latest")] + [InlineData("http://169.254.169.254/latest")] + [InlineData("https://localhost:11434/api/tags")] + [InlineData(@"\\169.254.169.254\latest")] + [InlineData(@"/\169.254.169.254/latest")] + [InlineData("api/tags")] + public void BuildUpstreamUri_RejectsPathsThatCanEscapeTheUpstreamOrigin(string pathAndQuery) + { + var exception = Assert.Throws( + () => UpstreamRequest.BuildUpstreamUri(Upstream, pathAndQuery)); + + Assert.Contains("origin-form path", exception.Message); + } +}