From f4500761dee9ebe87d7aa32af8a66522b656207e Mon Sep 17 00:00:00 2001 From: Tommaso Bailetti Date: Tue, 6 Oct 2026 09:09:36 +0200 Subject: [PATCH] refactor: merging controller repo inside the module --- .github/workflows/clean-registry.yml | 2 +- .github/workflows/controller-tests.yml | 61 + .github/workflows/update-controller.yml | 38 - .gitignore | 4 + AGENTS.md | 45 +- README.md | 10 +- build-images.sh | 19 +- controller/README.md | 309 ++++ controller/api/Containerfile | 38 + controller/api/README.md | 1216 ++++++++++++++++ controller/api/account_test.go | 329 +++++ controller/api/configuration/configuration.go | 354 +++++ .../api/configuration/configuration_test.go | 178 +++ controller/api/entrypoint.sh | 31 + controller/api/go.mod | 68 + controller/api/go.sum | 210 +++ controller/api/logs/logs.go | 25 + controller/api/logs/logs_test.go | 24 + controller/api/main.go | 277 ++++ controller/api/main_test.go | 1240 ++++++++++++++++ controller/api/methods/account.go | 401 +++++ controller/api/methods/auth.go | 423 ++++++ controller/api/methods/defaults.go | 42 + controller/api/methods/report.go | 326 +++++ controller/api/methods/unit.go | 1002 +++++++++++++ controller/api/middleware/middleware.go | 467 ++++++ controller/api/middleware/middleware_test.go | 247 ++++ controller/api/middleware_test.go | 318 ++++ controller/api/models/account.go | 39 + controller/api/models/auth.go | 40 + controller/api/models/platform.go | 18 + controller/api/models/report.go | 111 ++ controller/api/models/ubus.go | 24 + controller/api/models/unit.go | 67 + controller/api/report_test.go | 142 ++ controller/api/response/response.go | 54 + controller/api/routines/routines.go | 47 + controller/api/socket/socket.go | 59 + controller/api/socket/socket_test.go | 40 + controller/api/storage/grafana_user.sql.tmpl | 13 + controller/api/storage/report_schema.sql.tmpl | 587 ++++++++ controller/api/storage/storage.go | 1290 +++++++++++++++++ controller/api/storage/storage_test.go | 213 +++ controller/api/storage/upgrade_schema.sql | 36 + controller/api/storage_test.go | 142 ++ controller/api/utils/geoip.go | 100 ++ controller/api/utils/utils.go | 215 +++ controller/api/utils/utils_test.go | 419 ++++++ controller/controller.te | 25 + controller/dev.sh | 109 ++ controller/proxy/Containerfile | 4 + controller/proxy/entrypoint.sh | 159 ++ controller/test/smoke.sh | 178 +++ controller/ui/Containerfile | 16 + controller/ui/entrypoint.sh | 30 + controller/vpn/Containerfile | 13 + controller/vpn/controller-auth | 7 + controller/vpn/entrypoint.sh | 87 ++ controller/vpn/handle-connection | 61 + controller/vpn/handle-disconnection | 7 + controller/vpn/ip | 3 + controller/vpn/renew-certs | 68 + renovate.json | 21 +- 63 files changed, 12084 insertions(+), 64 deletions(-) create mode 100644 .github/workflows/controller-tests.yml delete mode 100644 .github/workflows/update-controller.yml create mode 100644 controller/README.md create mode 100644 controller/api/Containerfile create mode 100644 controller/api/README.md create mode 100644 controller/api/account_test.go create mode 100644 controller/api/configuration/configuration.go create mode 100644 controller/api/configuration/configuration_test.go create mode 100755 controller/api/entrypoint.sh create mode 100644 controller/api/go.mod create mode 100644 controller/api/go.sum create mode 100644 controller/api/logs/logs.go create mode 100644 controller/api/logs/logs_test.go create mode 100644 controller/api/main.go create mode 100644 controller/api/main_test.go create mode 100644 controller/api/methods/account.go create mode 100644 controller/api/methods/auth.go create mode 100644 controller/api/methods/defaults.go create mode 100644 controller/api/methods/report.go create mode 100644 controller/api/methods/unit.go create mode 100644 controller/api/middleware/middleware.go create mode 100644 controller/api/middleware/middleware_test.go create mode 100644 controller/api/middleware_test.go create mode 100644 controller/api/models/account.go create mode 100644 controller/api/models/auth.go create mode 100644 controller/api/models/platform.go create mode 100644 controller/api/models/report.go create mode 100644 controller/api/models/ubus.go create mode 100644 controller/api/models/unit.go create mode 100644 controller/api/report_test.go create mode 100644 controller/api/response/response.go create mode 100644 controller/api/routines/routines.go create mode 100644 controller/api/socket/socket.go create mode 100644 controller/api/socket/socket_test.go create mode 100644 controller/api/storage/grafana_user.sql.tmpl create mode 100644 controller/api/storage/report_schema.sql.tmpl create mode 100644 controller/api/storage/storage.go create mode 100644 controller/api/storage/storage_test.go create mode 100644 controller/api/storage/upgrade_schema.sql create mode 100644 controller/api/storage_test.go create mode 100644 controller/api/utils/geoip.go create mode 100644 controller/api/utils/utils.go create mode 100644 controller/api/utils/utils_test.go create mode 100644 controller/controller.te create mode 100755 controller/dev.sh create mode 100644 controller/proxy/Containerfile create mode 100755 controller/proxy/entrypoint.sh create mode 100755 controller/test/smoke.sh create mode 100644 controller/ui/Containerfile create mode 100755 controller/ui/entrypoint.sh create mode 100644 controller/vpn/Containerfile create mode 100755 controller/vpn/controller-auth create mode 100755 controller/vpn/entrypoint.sh create mode 100755 controller/vpn/handle-connection create mode 100755 controller/vpn/handle-disconnection create mode 100755 controller/vpn/ip create mode 100755 controller/vpn/renew-certs diff --git a/.github/workflows/clean-registry.yml b/.github/workflows/clean-registry.yml index 2e7102ea..fba93e99 100644 --- a/.github/workflows/clean-registry.yml +++ b/.github/workflows/clean-registry.yml @@ -12,5 +12,5 @@ jobs: steps: - uses: NethServer/ns8-github-actions/.github/actions/delete-image@v1 with: - images: "nethsecurity-controller webssh" + images: "nethsecurity-controller webssh nethsecurity-vpn nethsecurity-api nethsecurity-ui nethsecurity-proxy" delete_image_token: ${{ secrets.IMAGES_CLEANUP_TOKEN }} diff --git a/.github/workflows/controller-tests.yml b/.github/workflows/controller-tests.yml new file mode 100644 index 00000000..72b8112f --- /dev/null +++ b/.github/workflows/controller-tests.yml @@ -0,0 +1,61 @@ +name: Controller tests + +on: + push: + paths: + - "controller/**" + - ".github/workflows/controller-tests.yml" + pull_request: + paths: + - "controller/**" + - ".github/workflows/controller-tests.yml" + workflow_dispatch: + +defaults: + run: + working-directory: controller + +jobs: + api: + name: API Tests + runs-on: ubuntu-24.04 + steps: + - name: Checkout + uses: actions/checkout@v7 + - name: Install Podman + run: | + sudo apt-get update + sudo apt-get install -y podman oathtool + - name: Start TimescaleDB + run: | + # renovate: datasource=docker depName=docker.io/timescale/timescaledb + podman run --rm -d --name timescaledb -p 5432:5432 -e POSTGRES_PASSWORD=password -e POSTGRES_USER=report docker.io/timescale/timescaledb:2.23.1-pg16 + # Wait for DB to be ready + for i in {1..30}; do + podman exec timescaledb pg_isready -U report && break + sleep 1 + done + - name: Setup Go + uses: actions/setup-go@v7 + with: + go-version-file: controller/api/go.mod + - name: Test with the Go CLI + run: cd api && go test ./... -coverpkg=./... -coverprofile=coverage.out -v + + smoke: + name: Smoke Tests + runs-on: ubuntu-24.04 + steps: + - name: Checkout + uses: actions/checkout@v7 + - name: Install dependencies + run: | + sudo apt-get update + sudo apt-get install -y podman buildah jq + - name: Create tunsec device + run: | + sudo ip tuntap add dev tunsec mod tun + sudo ip addr add 172.21.0.1/16 dev tunsec + sudo ip link set dev tunsec up + - name: Run smoke test + run: ./test/smoke.sh diff --git a/.github/workflows/update-controller.yml b/.github/workflows/update-controller.yml deleted file mode 100644 index 60f1b23f..00000000 --- a/.github/workflows/update-controller.yml +++ /dev/null @@ -1,38 +0,0 @@ -name: Update controller package - -# **What it does**: Every nigth, at midnight checks if a new version of nethsecurity-controller is available. -# **Why we have it**: To avoid manually updating the package. -# **Who does it impact**: build-images.sh and the UI_COMMIT value - -on: - workflow_dispatch: - - schedule: - - cron: "0 0 * * *" - -jobs: - update-package: - name: Update nethsecurity-controller package - - runs-on: ubuntu-latest - - steps: - - name: Checkout - uses: actions/checkout@v7 - with: - fetch-depth: 0 - - name: Update apt - run: sudo apt update - - name: Install deps - run: sudo apt-get install -y curl jq git - - name: Check if new UI commit is different - run: | - NEW_TAG=$(curl https://api.github.com/repos/NethServer/nethsecurity-controller/tags | jq -r .[0].name) - sed -i "s/controller_version=.*/controller_version=\"$NEW_TAG\"/g" build-images.sh - - name: Commit and create PR - uses: peter-evans/create-pull-request@v8 - with: - title: 'build(deps): Update nethsecurity-controller package (automated)' - branch: 'build-update-nethsecurity-controller-package-automated' - commit-message: 'build(deps): nethsecurity-controller package: update nethsecurity-controller package (automated)' - base: main diff --git a/.gitignore b/.gitignore index bbb1494c..b05c3fa5 100644 --- a/.gitignore +++ b/.gitignore @@ -1,5 +1,6 @@ # Logs logs +!controller/api/logs/ *.log npm-debug.log* yarn-debug.log* @@ -105,3 +106,6 @@ dist # Tests output tests/outputs/ + +# Controller API binary +controller/api/api diff --git a/AGENTS.md b/AGENTS.md index f9d075e3..a8f355dd 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -4,28 +4,53 @@ This file provides guidance to AI agents when working with code in this reposito ## Repository overview -This is an **NS8 module** (`ns8-nethsecurity-controller`) that packages and configures an instance of [nethsecurity-controller](https://github.com/NethServer/nethsecurity-controller) on a NethServer 8 node. A single controller instance manages a fleet of NethSecurity firewall units over an OpenVPN tunnel, collecting their logs/metrics and exposing a Vue cluster-admin UI. +This is an **NS8 module** (`ns8-nethsecurity-controller`) that packages and configures an instance of nethsecurity-controller on a NethServer 8 node. A single controller instance manages a fleet of NethSecurity firewall units over an OpenVPN tunnel, collecting their logs/metrics and exposing a Vue cluster-admin UI. The module's containers run as one systemd-managed Podman pod: -- `nethsecurity-api` — Go REST API (source lives in the separate `NethServer/nethsecurity-controller` repo, **not** in this repo; only consumed as a prebuilt image) -- `nethsecurity-vpn` — OpenVPN server (external image) -- `nethsecurity-ui` — lighttpd static UI server (external image; distinct from this repo's own `ui/`) -- `nethsecurity-proxy` — Traefik reverse proxy (external image) +- `nethsecurity-api` — Go REST API (`controller/api/`) +- `nethsecurity-vpn` — OpenVPN server (`controller/vpn/`) +- `nethsecurity-ui` — lighttpd serving the NethSecurity UI built from `NethServer/nethsecurity-ui` (`controller/ui/`; distinct from this repo's own `ui/`) +- `nethsecurity-proxy` — Traefik reverse proxy (`controller/proxy/`) - `promtail` / `loki` — log shipping and storage - `prometheus` — metrics scraping - `timescale` — TimescaleDB for network-traffic/DPI/VPN time-series data - `grafana` — dashboards over Prometheus + Loki + Timescale - `webssh` — web SSH client, built locally from upstream `huashengdun/webssh` with a custom template -Only the module glue (`imageroot/`), the cluster-admin frontend (`ui/`), and the `webssh` UI override are implemented in this repo; the API server, VPN, UI-proxy, and other images are pulled prebuilt and pinned in `build-images.sh`. +The module glue (`imageroot/`), the cluster-admin frontend (`ui/`), the `webssh` UI override and the controller services (`controller/`) are implemented in this repo; the other images are pulled prebuilt and pinned in `build-images.sh`. ## Build `build-images.sh` builds images with **buildah** (no top-level Containerfile — built imperatively via `buildah from/add/config/commit`). It: 1. Builds `webssh` from `python:3.13.14-alpine`, unpacking the upstream `huashengdun/webssh` release and replacing its UI with `webssh/index.html`. -2. Builds the Vue `ui/` app in a `node:24.17.0-slim` builder container (`corepack enable && yarn install --frozen-lockfile && yarn build`), then assembles the main `nethsecurity-controller` image from `scratch` with `imageroot/` → `/imageroot` and `ui/dist` → `/ui`, plus NS8 image labels (`org.nethserver.authorizations`, `org.nethserver.min-core`, `org.nethserver.images`, etc.). +2. Builds `nethsecurity-{vpn,api,ui,proxy}` with `buildah build --target dist` from `controller//Containerfile`. The module references them with the same `IMAGETAG` as the module image. +3. Builds the Vue `ui/` app in a `node:24.17.0-slim` builder container (`corepack enable && yarn install --frozen-lockfile && yarn build`), then assembles the main `nethsecurity-controller` image from `scratch` with `imageroot/` → `/imageroot` and `ui/dist` → `/ui`, plus NS8 image labels (`org.nethserver.authorizations`, `org.nethserver.min-core`, `org.nethserver.images`, etc.). -Version pins for all external images (`controller_version`, `promtail_image`, `loki_image`, `prometheus_image`, `grafana_image`, `timescale_image`, `webssh_version`) live at the top of `build-images.sh`. `controller_version` is kept in sync automatically both by Renovate (custom regex manager in `renovate.json`) and by the nightly `update-controller.yml` workflow. +Version pins for all external images (`promtail_image`, `loki_image`, `prometheus_image`, `grafana_image`, `timescale_image`, `webssh_version`) live at the top of `build-images.sh`. The nethsecurity-ui version is the `UI_VERSION` ARG in `controller/ui/Containerfile`, bumped by Renovate. + +## Controller services (`controller/`) + +`api/`, `vpn/`, `proxy/`, `ui/` each hold a Containerfile; `controller.te` is the SELinux policy. `ui-new/` is a Vue 3/Vite UI, not built or tested in CI yet. + +Dev environment: `controller/dev.sh start|stop` runs a local pod with all services plus TimescaleDB and writes `api.env`. It needs a `tunsec` device, created once as root: + +```bash +sudo ip tuntap add dev tunsec mod tun +sudo ip addr add 172.21.0.1/16 dev tunsec +sudo ip link set dev tunsec up +``` + +`dev.sh` defaults to images tagged with the current branch (as published by CI); run `IMAGE_TAG=latest ./dev.sh start` for images built locally with `build-images.sh`. `controller/test/smoke.sh` runs `build-images.sh`, starts the pod and checks login, units and health. + +Go API tests need TimescaleDB running. Use the image pinned as `timescale_image` in `build-images.sh`: + +```bash +podman run --rm -d --name timescaledb -p 5432:5432 -e POSTGRES_PASSWORD=password -e POSTGRES_USER=report +cd controller/api && go test ./... +podman stop timescaledb +``` + +Add or update tests for any API change, and keep the README in each service directory up to date. Commit scope for controller changes is the service (`api`, `vpn`, `ui`, `proxy`), e.g. `fix(vpn): resolve authentication handshake failure`. ## UI development (`ui/`) @@ -67,8 +92,8 @@ Integration tests use **Robot Framework** driven over SSH against a live NS8 nod ## CI (`.github/workflows/`) -All workflows are thin wrappers around reusable workflows in `NethServer/ns8-github-actions`; there is no dedicated local lint job: +Most workflows are thin wrappers around reusable workflows in `NethServer/ns8-github-actions`; there is no dedicated local lint job: - `publish-images.yml` — on push, runs `build-images.sh` via `publish-branch.yml@v1`; also runs `module-info.yml@v1` and (on stable/latest releases) `scan-with-trivy.yml@v1`. - `test-module.yml` — runs the Robot Framework suite after images publish, or manually via `workflow_dispatch`. -- `update-controller.yml` — nightly cron that bumps `controller_version` in `build-images.sh` and opens a PR. +- `controller-tests.yml` — Go API tests and `controller/test/smoke.sh`, on push/PR touching `controller/`. - `build-apidoc.yml` / `clean-apidoc.yml` — build/clean API docs from `validate-input.json`/`validate-output.json` changes. diff --git a/README.md b/README.md index a3f40028..5e68569f 100644 --- a/README.md +++ b/README.md @@ -1,6 +1,6 @@ # ns8-nethsecurity-controller -Setup and start an instance of [nethsecurity-controller](https://github.com/NethServer/nethsecurity-controller). +Setup and start an instance of [nethsecurity-controller](controller). Each node can host multiple controller instances. @@ -81,19 +81,19 @@ The module is composed by the following systemd units: ### API Server -The [api server](https://github.com/NethServer/nethsecurity-controller/tree/master/api) gives NethSecurity the ability to register itself to NS8 (through [`ns-plug`](https://dev.nethsecurity.org/nethsecurity/packages/ns-plug/)) and gives access to the on-demand generated credentials for the VPN. +The [api server](controller/api) gives NethSecurity the ability to register itself to NS8 (through [`ns-plug`](https://dev.nethsecurity.org/nethsecurity/packages/ns-plug/)) and gives access to the on-demand generated credentials for the VPN. The API also registers the endpoints for the [Traefik Proxy](#proxy-and-ui) that allows direct interaction with the firewall, even if it's not in the same network. ### VPN -The [OpenVPN container](https://github.com/NethServer/nethsecurity-controller/tree/master/vpn) tunnels connection from the NethSecurity to the NS8 through a VPN tunnel, due to [firewall configuration](https://github.com/NethServer/ns8-nethsecurity-controller/blob/main/imageroot/actions/configure-module/20configure#L87) in NS8, no client can be reached from other clients and only client-server communication is allowed. +The [OpenVPN container](controller/vpn) tunnels connection from the NethSecurity to the NS8 through a VPN tunnel, due to [firewall configuration](https://github.com/NethServer/ns8-nethsecurity-controller/blob/main/imageroot/actions/configure-module/20configure#L87) in NS8, no client can be reached from other clients and only client-server communication is allowed. The module uses the NS8 [TUN feature](https://dev.nethsecurity.org/ns8-core/core/tun/) to create a new network interface and assign it to the VPN container. ### Proxy and UI -The [UI](https://github.com/NethServer/nethsecurity-controller/tree/master/ui) allows the browse of the interface directly off the NethSecurity installation, this is possible due to the [Traefik Proxy](https://github.com/NethServer/nethsecurity-controller/tree/master/proxy) server that redirects the urls to the correct IP inside the VPN. +The [UI](controller/ui) allows the browse of the interface directly off the NethSecurity installation, this is possible due to the [Traefik Proxy](controller/proxy) server that redirects the urls to the correct IP inside the VPN. ### Promtail @@ -243,7 +243,7 @@ journalctl _UID=$(id -u nethsecurity-controller1) --grep 'MIGRATION' ### Database maintenance The database is used to store configuration and metrics. -See [Database design](https://github.com/NethServer/nethsecurity-controller/tree/main/api#database-design) for more details. +See [Database design](controller/api/README.md#database-design) for more details. #### DPI stats cleanup diff --git a/build-images.sh b/build-images.sh index 2a9259fa..acea8454 100755 --- a/build-images.sh +++ b/build-images.sh @@ -9,7 +9,6 @@ images=() repobase="${REPOBASE:-ghcr.io/nethserver}" # Configure the image name reponame="nethsecurity-controller" -controller_version="2.4.1" promtail_image="docker.io/grafana/promtail:3.6.11" loki_image="docker.io/grafana/loki:2.9.17" prometheus_image="docker.io/prom/prometheus:v3.15.0" @@ -53,6 +52,16 @@ buildah commit "${webssh}" "${repobase}/webssh" # Append the image URL to the images array images+=("${repobase}/webssh") +# Build controller service images +for service in vpn api ui proxy; do + echo "Build nethsecurity-${service} container" + buildah build --layers --target dist \ + --file "controller/${service}/Containerfile" \ + --tag "${repobase}/nethsecurity-${service}" \ + "controller/${service}" + images+=("${repobase}/nethsecurity-${service}") +done + # Create a new empty container image container=$(buildah from scratch) @@ -78,10 +87,10 @@ buildah config --entrypoint=/ \ --label="org.nethserver.min-core=3.20.1" \ --label="org.nethserver.tcp-ports-demand=11" \ --label="org.nethserver.images=\ - ghcr.io/nethserver/nethsecurity-vpn:$controller_version \ - ghcr.io/nethserver/nethsecurity-api:$controller_version \ - ghcr.io/nethserver/nethsecurity-ui:$controller_version \ - ghcr.io/nethserver/nethsecurity-proxy:$controller_version \ + ghcr.io/nethserver/nethsecurity-vpn:${IMAGETAG:-latest} \ + ghcr.io/nethserver/nethsecurity-api:${IMAGETAG:-latest} \ + ghcr.io/nethserver/nethsecurity-ui:${IMAGETAG:-latest} \ + ghcr.io/nethserver/nethsecurity-proxy:${IMAGETAG:-latest} \ $promtail_image \ $loki_image \ $prometheus_image \ diff --git a/controller/README.md b/controller/README.md new file mode 100644 index 00000000..f5daed5c --- /dev/null +++ b/controller/README.md @@ -0,0 +1,309 @@ +# nethsecurity-controller + +The controller (server) is a set of containers that allow the admin to remotely manage multiple [NethSecurity](https://github.com/NethServer/nethsecurity) installations (firewalls). + +Firewalls can register to the server using [ns-plug](https://github.com/NethServer/nethsecurity/tree/master/packages/ns-plug) client. Upon registration the server will: +- create a VPN configuration which is sent back to the firewall +- create a route inside the proxy to access the firewall Luci RPC +- store credentials to access the remote firewall + +## Development environments + +Depending on your needs, you can setup a development environment to test the controller locally or you can install it on a NethServer 8 machine. + +### Local environment + +This environment integrates only the basic componentes: +- VPN server +- API server +- proxy server +- UI server + +It's suitable for development of the the basic controller features, but it does not include the full features like reporting and logging. You can use it also to develop the UI. + +If you need a local development environment, you can use the `dev.sh` script to start a podman pod with all the containers needed to run the controller. +First make sure to have [podman](https://podman.io/) installed on your server. +Containers should run under non-root users, but first you need to configure the tun device and the user. + +As root, execute: +``` +ip tuntap add dev tunsec mod tun +ip addr add 172.21.0.1/16 dev tunsec +ip link set dev tunsec up +``` + +If you're running the dev environment on a distro with SELinux enabled, you also may need to create a module to allow the controller to access the tun device. +Just execute: +``` +checkmodule -M -m -o controller.mod controller.te +semodule_package -o controller.pp -m controller.mod +semodule -i controller.pp +``` + +Then change to non-root user, clone this repository and execute: +``` +su - controller + +./dev.sh start +``` + +To stop the pod, execute: +``` +./dev.sh stop +``` + +By default, the script will use images tagged with the current branch name. +If you want to use a specific image tag, you can set the `IMAGE_TAG` environment + +To run a specific image tag, you can use: +``` +IMAGE_TAG= ./dev.sh start +``` + +The server will be available at `http://localhost:8080/`. +Default credentials are: `admin/admin`. + +### NethServer 8 environment + +You can install the controller on [NS8](https://github.com/NethServer/ns8-nethsecurity-controller#install). + +After the installation and first configuration of the controller, enable the debug mode to allow the UI to connect to the controller API: +``` +runagent -m nethsecurity-controller1 sed -i 's/GIN_MODE=release/GIN_MODE=debug/' api.env +runagent -m nethsecurity-controller1 systemctl --user restart controller +``` + +Above commands assume that the controller instance is named `nethsecurity-controller1`. + +### UI development + +If you need to the develop the UI, first clone the [nethsecurity-ui](https://github.com/nethserver/nethsecurity-ui). + +Then you can choose to connect to the local development environment or the NethServer 8 environment. + +#### Connect to local development environment + +To connect to the local development environment, you need to: +- start the controller using the `dev.sh` script from the `nethsecurity-controller` repository: + ``` + IMAGE_TAG=pr-123 ./dev.sh start + ``` +- move to the `nethsecurity-ui` directory: + ``` + git clone git@github.com:NethServer/nethsecurity-ui.git + cd nethsecurity-ui + ``` +- inside the UI directory, setup the `.env.development` file to connect to the controller API + ``` + cat < .env.development + VITE_API_SCHEME=http + VITE_CONTROLLER_API_HOST=localhost:8080 + VITE_UI_MODE=controller + EOF + ``` +- still inside the UI directory, start the UI in dev mode: + ``` + ./dev.sh + ``` +- access to the dev UI URL generated by vite, usually `http://localhost:5173/` + + +#### Connect to NethServer 8 environment + +First, make sure the controller is running in debug mode, as described above. + +To connect to the NethServer 8 environment, you need to: +- move to the `nethsecurity-ui` directory: + ``` + git clone git@github.com:NethServer/nethsecurity-ui.git + cd nethsecurity-ui + ``` +- inside the UI directory, setup the `.env.development` file to connect to the controller API + ``` + cat < .env.development + VITE_API_SCHEME=https + VITE_CONTROLLER_API_HOST=controller.example.com + VITE_UI_MODE=controller + EOF +- still inside the UI directory, start the UI in dev mode: + ``` + ./dev.sh + ``` +- access to the dev UI URL generated by vite, usually `http://localhost:5173/` + +### API server development + +It's possible to develop the API server without the need to run the full controller stack. +This environment is suitable to develop the API server and test it using the `curl` command, but +does not integrate nor the VPN server nor the proxy server. + +To start the API server in development mode: +- first, make sure that a Timescale DB is running and accessible: + + ``` + podman run --rm --name timescaledb -p 5432:5432 -e POSTGRES_PASSWORD=password -e POSTGRES_USER=report docker.io/timescale/timescaledb:2.23.1-pg16 + ``` +- then move to the `api` directory, create the `data` directory and build the API server: + ``` + cd api + mkdir -p data + go build + ``` +- start the API server by setting up all required environment variables: + ``` + LISTEN_ADDRESS=0.0.0.0:5000 ADMIN_USERNAME=admin ADMIN_PASSWORD=admin SECRET_JWT=secret PROMTAIL_ADDRESS=127.0.0.1 PROMTAIL_PORT=6565 PROMETHEUS_PATH="/prometheus" WEBSSH_PATH="/webssh" GRAFANA_PATH="/grafana" REGISTRATION_TOKEN=1234 REPORT_DB_URI=postgres://report:password@127.0.0.1:5432/report GRAFANA_POSTGRES_PASSWORD=password ISSUER_2FA=test ENCRYPTION_KEY=12345678901234567890123456789012 VALID_SUBSCRIPTION=true CREDENTIALS_DIR=data DATA_DIR=data ./api + ``` +- the API server will be available at `http://localhost:5000/` + +#### Testing the API server + +To run the testing suite: +- first, make sure that a Timescale DB is running and accessible: + + ``` + podman run --rm --name timescaledb -p 5432:5432 -e POSTGRES_PASSWORD=password -e POSTGRES_USER=report docker.io/timescale/timescaledb:2.23.1-pg16 + ``` +- then move to the `api` directory and run the tests: + ``` + cd api + go test + ``` + +## How it works + +General workflow: + +1. Access the controller and add a new machine using the `add` API below. This will generate a join code containing the FQDN of the controller, a registration token, and the unit UUID. +2. Connect the NethSecurity unit and register the machine using the join code. +3. Return to the controller and manage the unit. + - The UI retrieves a token for the NethSecurity unit: `curl http://localhost:8080/api/servers/login/clientX` + - THe UI Uses the token to invoke Luci APIs: `curl http://localhost:8080/clientX/cgi-bin/luci/rpc/...` + + +### Serving unit UIs + +A unit reporting a supported `ui_version` is opened at `https:////`, where +traefik proxies the unit's own nginx and the unit serves its own UI at its own version. Older units +render the copy of the standalone UI bundled into the controller. + +This runs unit-supplied JavaScript on the controller's origin — a path prefix is not an origin +boundary, so everything under `//` shares one `localStorage` with the controller UI. + +### Services + +The controller is composed by 4 services: +- nethsec-vpn: OpenVPN server, it authenticates the machines and create routes for the proxy, it listens on port 1194 +- nethsec-proxy: traefik forwards requests to the connected machines using the machine name as path prefix, it listens on port 8181 +- nethsec-api: REST API python server to manage nethsec-vpn clients, it listens on port 5000 +- nethsec-ui: lighttpd instance serving static UI files, it listens on port 3000 + +### Certificate renewal + +The OpenVPN PKI lasts 10 years. At every restart the VPN container renews +whatever has less than 6 months left: the CA, the server certificate, the unit +certificates and the revocation list. + +Renewal keeps the existing key. Units get the new certificate the next time +they register. + +Renewing the CA re-issues every certificate, so connected units drop and +reconnect. + +## Environment configuration + +The following environment variables can be used to configure the containers: + +- `FQDN`: default is the container/pod hostname +- `OVPN_NETWORK`: OpenVPN network, default is `172.21.0.0` +- `OVPN_NETMASK`: OpenVPN netmask, default is `255.255.0.0` +- `OVPN_CN`: OpenVPN certificate CN, default is `nethsec` +- `OVPN_UDP_PORT`: OpenVPN UDP port, default is `1194` +- `OVPN_TUN`: OpenVPN tun device name, default is `tunsec` +- `OVPN_TUN_MTU`: OpenVPN tun device MTU, default is `1500` +- `OVPN_MSSFIX`: OpenVPN mssfix value, default is `1450` +- `UI_PORT`: UI listening port, default is `3000` +- `UI_BIND_IP`: UI binding IP, default is `0.0.0.0` +- `API_PORT`: API server listening port, default is `5000` +- `API_BIND_IP`: API server listening IP, default is `127.0.0.1` +- `API_USER`: controller admin user, default is `admin` +- `API_PASSWORD`: controller admin password, it must be passed as SHA56SUM, default is `admin` +- `API_SECRET`: JWT secret token +- `API_DEBUG`: enable debug logging and CORS if set to `1`, default is `0` +- `API_SESSION_DURATION`: JWT session duration in seconds, default is 7 days +- `GLOBAL_RATE_LIMIT_AVERAGE`: max sustained requests per second per client IP across all API routes, default is `25`; set to `0` to disable rate limiting +- `GLOBAL_RATE_LIMIT_BURST`: burst allowance above the average before requests are rejected with HTTP 429, default is `100` +- `PROXY_PORT`: proxy listening port, default is `8080` +- `PROXY_BIND_IP`: proxy binding IP, default is `0.0.0.0` +- `REPORT_DB_URI`: Timescale DB URI, like `postgresql://user:password@host:port/dbname` +- `ALLOWED_IPS`: comma-separated list of allowed IPs, if empty, all IPs are allowed, default is empty +- `PUBLIC_ENDPOINTS`: comma-separated list of public endpoints, that can be accessed even if `ALLOWED_IPS` is set, default is empty + If ALLOWED_IPS is set, the public endpoints should allow registration and ingestions from units, a good value should be: `/api/ingest,/api/units/register` + +## REST API + +Manage server registrations using the REST API server. +Request should be sent to the proxy server. + +Almost all APIs are authenticated using [JWT](https://flask-jwt-extended.readthedocs.io/en/stable/). + +Authentication work-flow: + +1. send user name and password to `/login` API +2. retrieve authorization tokens: + - `access_token`: it's the token used to executed all APIs, it expires after an hour + - `refresh_token`: this token can be used only to call the `/refresh` API and request a new `access_token`, it expires after `API_SESSION_DURATION` seconds (default to 7 days) +3. invoke other APIs by setting the header `Authorization: Bearer "` + +Unauthenticated APIs: + +- `/login`: execute the login and retrieve the tokens +- `/register`: invoked by firewalls to register themselves, this API should be always invoked using a valid HTTPS endpoint to + ensure the identity of the server + +See the [API documentation](api/README.md) for more details. + +## Build + +Each container is built using a Containerfile, which is both compatible with `docker build` command and `podman build`. + +All images, including these, are built by `build-images.sh` in the repository root, using buildah: +```bash +cd .. +./build-images.sh +``` + +Images are tagged `latest`. + +To build the images using podman, you can use the following: + +```bash +podman build --target dist --layers --force-rm --jobs 0 +``` + +Where `` is the path to any of the directory to build. + +Optionally, you can add the `--tag ` to tag the image with a specific name. + +## Smoke Testing + +A comprehensive smoke test is provided to verify the entire stack is working correctly. The test: +1. Builds all containers with `build-images.sh` +2. Starts the development environment +3. Verifies all services are running and responsive + +To run the smoke test: +```bash +./test/smoke.sh +``` + +The script will: +- Check prerequisites (podman, buildah, curl, jq, tunsec device) +- Build all containers using `../build-images.sh` +- Start the stack with the `latest` images +- Verify all 5 containers (vpn, db, api, ui, proxy) are running +- Test login and JWT token generation +- Test unit creation and retrieval +- Verify health endpoints +- Check VPN PKI directory and database connectivity + +The test automatically cleans up and stops the stack on exit. diff --git a/controller/api/Containerfile b/controller/api/Containerfile new file mode 100644 index 00000000..908b02f4 --- /dev/null +++ b/controller/api/Containerfile @@ -0,0 +1,38 @@ +FROM docker.io/golang:1.26.8 AS build +WORKDIR /build +COPY go.mod . +COPY go.sum . +RUN go mod download +COPY configuration configuration +COPY logs logs +COPY methods methods +COPY middleware middleware +COPY models models +COPY response response +COPY routines routines +COPY socket socket +COPY storage storage +COPY utils utils +COPY main.go . +COPY main_test.go . +ENV GOOS=linux +ENV GOARCH=amd64 +ENV CGO_ENABLED=1 +RUN go build -ldflags='-extldflags=-static' -tags sqlite_omit_load_extension + +FROM build AS test +RUN go test + +FROM docker.io/alpine:3.22.6 AS dist +RUN apk add --no-cache \ + curl \ + easy-rsa \ + oath-toolkit-oathtool \ + openssh \ + sqlite +WORKDIR /nethsecurity-api +COPY --from=build /build/api /nethsecurity-api/api +COPY entrypoint.sh /entrypoint.sh +ENTRYPOINT ["/entrypoint.sh"] +CMD ["./api"] + diff --git a/controller/api/README.md b/controller/api/README.md new file mode 100644 index 00000000..6c4501b0 --- /dev/null +++ b/controller/api/README.md @@ -0,0 +1,1216 @@ +# nethsecurity-controller + +## Build + +```bash +CGO_ENABLED=0 go build +``` + +# Environment variables + +**Mandatory** + +- `ADMIN_USERNAME`: admin username to login +- `ADMIN_PASSWORD`: admin password to login +- `SECRET_JWT`: secret to sing JWT tokens +- `REGISTRATION_TOKEN`: secret token used to register units + +- `CREDENTIALS_DIR`: directory to save credentials of connected units + +- `PROMTAIL_ADDRESS`: promtail address +- `PROMTAIL_PORT`: promtail port + +- `PROMETHEUS_PATH`: prometheus web path +- `WEBSSH_PATH`: webssh web path +- `GRAFANA_PATH`: grafana web path +- `GRAFANA_POSTGRES_PASSWORD`: password to access grafana postgres database +- `REPORT_DB_URI`: Timescale database URI for reports + +**Optional** + +- `LISTEN_ADDRESS`: a comma-separated list of listen addresses for the server. Each entry is in the form `
:` - _default_: `127.0.0.1:5000` + Example: `127.0.0.1:5000,192.168.100.1:5000` + +- `OVPN_DIR`: openvpn configuration directory - _default_: `/etc/openvpn` +- `OVPN_NETWORK`: openvpn network address - _default_: `172.21.0.0` +- `OVPN_NETMASK`: openvpn netmask - _default_: `255.255.0.0` +- `OVPN_UDP_PORT`: openvpn udp port - _default_: `1194` + +- `OVPN_C_DIR`: openvpn path of ccd directory - _default_: OVPN_DIR + `/ccd` +- `OVPN_P_DIR`: openvpn path of proxy directory - _default_: OVPN_DIR + `/proxy` +- `OVPN_K_DIR`: openvpn path of pki directory - _default_: OVPN_DIR + `/pki` +- `OVPN_M_SOCK`: opevpn management socket path - _default_: OVPN_DIR + `/run/mgmt.sock` + +- `EASYRSA_PATH`: easyrsa command path - _default_: `/usr/share/easy-rsa/easyrsa` + +- `PROXY_PROTOCOL`: traefik protocol - _default_: `http://` +- `PROXY_HOST`: traefik host - _default_: `localhost` +- `PROXY_PORT`: traefik port - _default_: `8080` +- `LOGIN_ENDPOINT`: unit login endpoint, on stand-alone api server - _default_: `/api/login` + +- `FQDN`: fully qualified domain name of the machine - _default_: `hostname -f` + +- `CACHE_TTL`: cache time to live for unit information in seconds - _default_: `7200` (2 hours) + Unit information are fetched from the connected units. The cache is refreshed every hour. + +- `RETENTION_DAYS`: configure how many days the metrics should be kept - _default_: `60` + +- `MAXMIND_LICENSE`: license key for maxmind geolite2 database - _default_: `` + If the license key is not set, the geolite2 database will not be downloaded. +- `GEOIP_DB_DIR`: directory to save geolite2 database - _default_: current directory + +- `SENSITIVE_LIST`: list of sensitive information to be redacted in logs +- `VALID_SUBSCRIPTION`: valid subscription status - _default_: `false` + +- `ENCRYPTION_KEY`: key to encrypt/decrypt sensitive data, it must be 32 bytes long + +- `PLATFORM_INFO`: a JSON string with platform information, used to store the controller version and other information. It can be left empty. + Example: `{"vpn_port":"1194","vpn_network":"192.168.100.0/24", "controller_version":"1.0.0", "metrics_retention_days":30, "logs_retention_days":90}` + +- `PROMETHEUS_AUTH_PASSWORD` and `PROMETHEUS_AUTH_USERNAME`: credentials to access the `/prometheus/targets` endpoint for listing connected units + in Prometheus target format. - _default_: `prometheus:prometheus` + +## User and units authorizations + +A unit is a NethSecurity firewall that is connected to the controller. +Units are identified by a unique ID and are stored in the database. VPN info of the unit are stored in a file in the `OVPN_DIR` directory. + +A group of units is a collection of units that can be managed together. +Units can be added to a group via the API. The group is identified by a unique ID and is stored in the database. + +A user is an account that can access the API and the UI. +User accounts are stored in the database and can be managed via the API. +By default, a user can't access any unit. +A user can be promoted to an admin user, which allows the user to manage other users and units. + +Admin users have the `admin` flag set to `true` and can: + +- create, modify and delete user accounts +- create, modify and delete units +- create, modify and delete units groups +- assign a user to one or more groups of units +- see all units, despite the groups assigned to the user + +The following rules apply: + +- a user can be assigned to one or more groups of units, if the user is not assigned to any group, the user can't see any unit +- a non existing-unit cannot be added to a group +- a unit group that is associated to a user account can't be deleted +- a non-existing unit group cannot be assigned to a user account +- when a unit is deleted from the database, it is removed from all groups + +## Database design + +The database is designed to efficiently manage and report on a large number of NethSecurity firewall units connected to a central controller. Each unit is uniquely identified and related data is stored in several reporting tables (such as `dpi_stats`, `ovpnrw_connections`, `ts_attacks`, etc.), supporting analytics and monitoring via TimescaleDB continuous aggregates. + +Key design choices include: + +- **No Foreign Keys:** + Foreign key constraints (especially with `ON DELETE CASCADE`) were removed from all tables. With high-volume tables like `dpi_stats` (which can grow by millions of rows per day on a controller with 70 units), cascading deletes caused severe performance issues—deleting a unit could take hours due to the volume of related data. + +- **Orphaned Data Cleanup:** + Instead of relying on foreign keys, a stored procedure (`cleanup_orphaned_unit_data`) is scheduled to run daily. This procedure deletes all records from reporting tables where the `uuid` is not present in the `units` table, ensuring data consistency and preventing orphaned records. + +- **Indexes on UUID:** + All reporting tables have an index on the `uuid` field to speed up deletion and lookup operations. For the `dpi_stats` table, the index is only created on new (empty) installations, as creating it on an existing, large table would be prohibitively slow. + +- **TimescaleDB Features:** + The schema leverages TimescaleDB hypertables and continuous aggregates for efficient time-series data storage and reporting, with retention policies to automatically drop old data. + Queries can be perfomed on aggregated tables even if the original data is deleted, as the continuous aggregates are updated periodically. + +## APIs + +### Auth + +- `GET /health` + + REQ + + ```json + Content-Type: application/json + ``` + + RES + + ```json + HTTP/1.1 200 OK + Content-Type: application/json; charset=utf-8 + + { + "status": "ok" + } + ``` + +- `POST /login` + + REQ + + ```json + Content-Type: application/json + + { + "username": "root", + "password": "Nethesis,1234" + } + ``` + + RES + + ```json + HTTP/1.1 200 OK + Content-Type: application/json; charset=utf-8 + + { + "code": 200, + "expire": "2023-05-25T14:04:03.734920987Z", + "token": "eyJh...E-f0" + } + ``` + +- `POST /logout` + + REQ + + ```json + Content-Type: application/json + Authorization: Bearer + ``` + + RES + + ```json + HTTP/1.1 200 OK + Content-Type: application/json; charset=utf-8 + + { + "code": 200 + } + ``` + +- `GET /refresh` + + REQ + + ```json + Content-Type: application/json + Authorization: Bearer + ``` + + RES + + ```json + HTTP/1.1 200 OK + Content-Type: application/json; charset=utf-8 + + { + "code": 200, + "expire": "2023-05-25T14:04:03.734920987Z", + "token": "eyJh...E-f0" + } + ``` + +### Units + +- `GET /units` + + REQ + + ```json + Content-Type: application/json + Authorization: Bearer + ``` + + RES + + ```json + HTTP/1.1 200 OK + Content-Type: application/json; charset=utf-8 + + { + "code": 200, + "data": [ + { + "ipaddress": "172.23.21.3", + "id": "", + "netmask": "255.255.255.0", + "groups": ["group1", "group2"], + "vpn": { + "bytes_rcvd": "21830", + "bytes_sent": "5641", + "connected_since": "1686312722", + "real_address": "192.168.122.220:41445", + "virtual_address": "172.23.21.3" + }, + "info": { + "unit_name": "myfw1", + "version": "8-23.05.2-ns.0.0.2-beta2-37-g6e74afc", + "subscription_type": "enterprise", + "system_id": "XXXXXXXX-XXXX", + "ssh_port": 22, + "fqdn": "fw.local", + "api_version": "1.0.0" + } + }, + ... + { + "ipaddress": "", + "id": "", + "netmask": "", + "vpn": {}, + "groups": [], + "info": { + "unit_name": "", + "version": "", + "subscription_type": "", + "system_id": "", + "ssh_port": 0, + "fqdn": "", + "api_version": "1.0.0" + } + } + ], + "message": "units listed successfully" + } + ``` + + The API takes a query parameter `cache`. If `cache` is set to `true`, the API will return the cached data, if data are fresh enough. + If `cache` is set to `false`, the API will always fetch the data from the connected units. + +- `GET /units/` + + REQ + + ```json + Content-Type: application/json + Authorization: Bearer + ``` + + RES + + ```json + HTTP/1.1 200 OK + Content-Type: application/json; charset=utf-8 + + { + "code": 200, + "data": { + "ipaddress": "172.23.21.3", + "id": "", + "netmask": "255.255.255.0", + "registered": true, + "vpn": { + "bytes_rcvd": "22030", + "bytes_sent": "5841", + "connected_since": "1686312722", + "real_address": "192.168.122.220:41445", + "virtual_address": "172.23.21.3" + }, + "info": { + "unit_name": "myfw1", + "version": "8-23.05.2-ns.0.0.2-beta2-37-g6e74afc", + "subscription_type": "enterprise", + "system_id": "XXXXXXXX-XXXX", + "ssh_port": 22, + "fqdn": "fw.local", + "api_version": "1.0.0" + }, + "join_code": "eyJmcWRuIjoiY29udHJvbGxlci5ncy5uZXRoc2VydmVyLm5ldCIsInRva2VuIjoiMTIzNCIsInVuaXRfaWQiOiI5Njk0Y2Y4ZC03ZmE5LTRmN2EtYjFjNC1iY2Y0MGUzMjhjMDIifQ==" + }, + "message": "unit listed successfully" + } + ``` + + The API takes a query parameter `cache`. If `cache` is set to `true`, the API will return the cached data, if data are fresh enough. + If `cache` is set to `false`, the API will always fetch the data from the connected units. + +- `GET /units//info` + + REQ + + ```json + Content-Type: application/json + Authorization: Bearer + ``` + + RES + + ```json + HTTP/1.1 200 OK + Content-Type: application/json; charset=utf-8 + + { + "code": 200, + "data": { + "fqdn": "NethSec", + "ssh_port": 22, + "subscription_type": "enterprise", + "system_id": "XXXXXXXX-XXXX", + "unit_name": "NethSec", + "version": "NethSecurity 8 23.05.3-ns.1.0.1", + "api_version": "1.0.0" + }, + "message": "unit info retrieved successfully" + } + ``` + + The backend stores data inside the database. This is useful for retrieving new information of the unit without waiting for cron to store it. + +- `GET /units//token` + + REQ + + ```json + Content-Type: application/json + Authorization: Bearer + ``` + + RES + + ```json + HTTP/1.1 200 OK + Content-Type: application/json; charset=utf-8 + + { + "code": 200, + "data": { + "expire": "2023-06-10T12:23:39.46160793Z", + "token": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9...rRikiEG83smBWPdHWzzhKOnfgzOkRXQntxdKGdaIhk8" + }, + "message": "unit token retrieved successfully" + } + ``` + +- `POST /units` + + REQ + + ```json + Content-Type: application/json + Authorization: Bearer + + { + "unit_id": "" + } + ``` + + RES + + ```json + HTTP/1.1 200 OK + Content-Type: application/json; charset=utf-8 + + { + "code": 200, + "data": { + "join_code": "eyJmcWRuIjoiY29udHJvbGxlci5ncy5uZXRoc2VydmVyLm5ldCIsInRva2VuIjoiMTIzNCIsInVuaXRfaWQiOiI2OThhMDQzZC02MGRiLTQyNmMtODRjZi1lODZhMTZmM2QxMzMifQ==" + }, + "message": "unit added successfully" + } + ``` + +- `POST /units/register` + + REQ + + ```json + Content-Type: application/json + Authorization: Bearer + + { + "unit_name": "fw.nethsecurity.local", + "unit_id": "d330b2db-cdfe-4c56-b9b6-f97e5b838748", + "username": "test", + "password": "Nethesis,1234", + "version": "8-23.05.2-ns.0.0.2-beta2-37-g6e74afc", + "subscription_type": "enterprise", + "system_id": "XXXXXXXX-XXXX" + } + ``` + + RES | unit previously added + + ```json + HTTP/1.1 200 OK + Content-Type: application/json; charset=utf-8 + + { + "code": 200, + "data": { + "ca": "-----BEGIN CERTIFICATE-----\n\n-----END CERTIFICATE-----", + "cert": "Certificate:\n\n-----END CERTIFICATE-----", + "host": "ns8.local", + "key": "-----BEGIN PRIVATE KEY-----\n\n-----END PRIVATE KEY-----", + "port": "1194", + "promtail_address": "172.21.0.1", + "promtail_port": "5151", + "api_port": "20001", + "vpn_address": "192.168.0.1" + }, + "message": "unit registered successfully" + } + ``` + + RES | unit not added in waiting list + + ```json + HTTP/1.1 200 OK + Content-Type: application/json; charset=utf-8 + + { + "code": 403, + "data": "", + "message": "unit added to waiting list" + } + ``` + +- `DELETE /units/` + + REQ + + ```json + Content-Type: application/json + Authorization: Bearer + ``` + + RES + + ```json + HTTP/1.1 200 OK + Content-Type: application/json; charset=utf-8 + + { + "code": 200, + "data": "", + "message": "unit deleted successfully" + } + ``` + +### Unit Groups + +- `GET /unit_groups` + + REQ + + ```json + Content-Type: application/json + Authorization: Bearer + ``` + + RES + + ```json + HTTP/1.1 200 OK + Content-Type: application/json; charset=utf-8 + + { + "code": 200, + "data": { + "unit_groups": [ + { + "id": 1, + "name": "Group 1", + "description": "This is a test group", + "units": ["unit_id_1", "unit_id_2"], + "created_at": "2024-03-14T09:37:28+01:00", + "updated_at": "2024-03-14T10:00:00+01:00", + "used_by": ["account_id_1", "account_id_2"] + } + ... + ] + }, + "message": "unit groups listed successfully" + } + ``` + +- `GET /unit_groups/` + + REQ + + ```json + Content-Type: application/json + Authorization: Bearer + ``` + + RES + + ```json + HTTP/1.1 200 OK + Content-Type: application/json; charset=utf-8 + + { + "code": 200, + "data": { + "unit_group": { + "id": 1, + "name": "Group 1", + "description": "This is a test group", + "units": ["unit_id_1", "unit_id_2"], + "created_at": "2024-03-14T09:37:28+01:00", + "updated_at": "2024-03-14T10:00:00+01:00" + } + }, + "message": "success" + } + ``` + +- `POST /unit_groups` + + REQ + + ```json + Content-Type: application/json + Authorization: Bearer + + { + "name": "Group 1", + "descrption": "This is a test group", + "units": ["unit_id_1", "unit_id_2"] + } + ``` + + RES + + ```json + HTTP/1.1 201 OK + Content-Type: application/json; charset=utf-8 + + { + "code": 201, + "data": {"id": 1}, + "message": "success" + } + ``` + +- `PUT /unit_groups/` + + REQ + + ```json + Content-Type: application/json + Authorization: Bearer + + { + "name": "Group 1 updated", + "description": "This is an updated test group", + "units": ["unit_id_1", "unit_id_3"] + } + ``` + + RES + + ```json + HTTP/1.1 200 OK + Content-Type: application/json; charset=utf-8 + + { + "code": 200, + "data": null, + "message": "success" + } + ``` + +- `DELETE /unit_groups/` + + REQ + + ```json + Content-Type: application/json + Authorization: Bearer + ``` + + RES + + ```json + HTTP/1.1 200 OK + Content-Type: application/json; charset=utf-8 + + { + "code": 200, + "data": "", + "message": "success" + } + ``` + +### Accounts + +- `GET /accounts` + + REQ + + ```json + Content-Type: application/json + Authorization: Bearer + ``` + + RES + + ```json + HTTP/1.1 200 OK + Content-Type: application/json; charset=utf-8 + + { + "code": 200, + "data": { + "accounts": [ + { + "id": 2, + "username": "test1", + "password": "", + "admin": true, + "display_name": "Test 1", + "created": "2024-03-14T09:37:28+01:00" + }, + ... + { + "id": 6, + "username": "test2", + "password": "", + "admin": false, + "display_name": "Test 2", + "created": "2024-03-14T11:43:33+01:00" + } + ], + "total": 5 + }, + "message": "success" + } + + ``` + +- `GET /accounts/` + + REQ + + ```json + Content-Type: application/json + Authorization: Bearer + ``` + + RES + + ```json + HTTP/1.1 200 OK + Content-Type: application/json; charset=utf-8 + + { + "code": 200, + "data": { + "account": { + "id": 2, + "username": "test3", + "password": "", + "admin": false, + "display_name": "Test 3", + "created": "2024-03-14T09:37:28+01:00" + } + }, + "message": "success" + } + ``` + +- `POST /accounts` + + REQ + + ```json + Content-Type: application/json + Authorization: Bearer + + { + "username": "test1", + "password": "Nethesis,1234", + "display_name": "Test 1", + "unit_groups": [1, 2], + "admin": false + } + ``` + + RES + + ```json + HTTP/1.1 201 OK + Content-Type: application/json; charset=utf-8 + + { + "code": 201, + "data": {"id": 5}, + "message": "success" + } + ``` + +- `PUT /accounts/` + + REQ + + ```json + Content-Type: application/json + Authorization: Bearer + + { + "password": "Nethesis,4321", + "display_name": "Test 5", + "unit_groups": [1, 2], + "admin": false + } + ``` + + RES + + ```json + HTTP/1.1 200 OK + Content-Type: application/json; charset=utf-8 + + { + "code": 200, + "data": null, + "message": "success" + } + ``` + +- `PUT /accounts/password` + + REQ + + ```json + Content-Type: application/json + Authorization: Bearer + + { + "old_password": "Nethesis,1234", + "new_password": "Nethesis,4321" + } + ``` + + RES + + ```json + HTTP/1.1 200 OK + Content-Type: application/json; charset=utf-8 + + { + "code": 200, + "data": null, + "message": "success" + } + ``` + +- `DELETE /accounts/` + + REQ + + ```json + Content-Type: application/json + Authorization: Bearer + ``` + + RES + + ```json + HTTP/1.1 200 OK + Content-Type: application/json; charset=utf-8 + + { + "code": 200, + "data": "", + "message": "success" + } + ``` + +- `GET /accounts/ssh-keys` + + REQ + + ```json + Content-Type: application/json + Authorization: Bearer + ``` + + RES + + ```json + HTTP/1.1 200 OK + Content-Type: application/json; charset=utf-8 + + { + "code": 200, + "data": { + "key": "-----BEGIN OPENSSH PRIVATE KEY-----\nb3BlbnNza...m3XHi7DiRCmyqbwdp86eV\n-----END OPENSSH PRIVATE KEY-----", + "key_pub": "ssh-rsa AAAAB3NzaC1yc2EAAAADAQABAAABgQDVled...UxVF6O0Esc3gFe0XMUT9Y+GtqM1O2s= test@local.domain" + }, + "message": "success" + } + ``` + +- `POST /accounts/ssh-keys` + + REQ + + ```json + Content-Type: application/json + Authorization: Bearer + + { + "passphrase": "Nethesis,2222" + } + ``` + + RES + + ```json + HTTP/1.1 201 OK + Content-Type: application/json; charset=utf-8 + + { + "code": 200, + "data": { + "key_pub": "ssh-rsa AAAAB3NzaC1yc2EAAAADAQABAAABgQDVled...UxVF6O0Esc3gFe0XMUT9Y+GtqM1O2s= test@local.domain" + }, + "message": "success" + } + ``` + +- `DELETE /accounts/ssh-keys` + + REQ + + ```json + Content-Type: application/json + Authorization: Bearer + ``` + + RES + + ```json + HTTP/1.1 200 OK + Content-Type: application/json; charset=utf-8 + + { + "code": 200, + "data": "", + "message": "success" + } + ``` + +- `GET /platform` + + REQ + + ```json + Content-Type: application/json + Authorization: Bearer + ``` + + RES + + ```json + HTTP/1.1 200 OK + Content-Type: application/json; charset=utf-8 + + { + "code": 200, + "data": { + "vpn_port": "1194", + "vpn_network": "172.21.0.0/16", + "controller_version": "1.2.3", + "nethserver_version": "8-23.05.3-ns.1.0.1", + "nethserver_system_id": "XXXXXXXX-XXXX", + "metrics_retention_days": 60, + "logs_retention_days": 30 + }, + "message": "success" + } + ``` + +## Basic authentication API + +- GET `/auth` + + This endpoint is used to check if the user is authenticated. It returns a 200 status code if the user is authenticated, otherwise it returns a 401 status code. + It can be used by external applications to check if the user is authenticated without needing to handle JWT tokens. + + REQ + + ```json + Content-Type: application/json + Authorization: Bearer + ``` + + RES + + ```json + HTTP/1.1 200 OK + Content-Type: application/json; charset=utf-8 + X-Auth-User: + + { + "authentication": "ok" + } + ``` + +### Defaults + +- `GET /defaults` + + REQ + + ```json + Content-Type: application/json + Authorization: Bearer + ``` + + RES + + ```json + HTTP/1.1 200 OK + Content-Type: application/json; charset=utf-8 + + { + "code": 200, + "data": { + "fqdn": "controller.ns8.local", + "grafana_path": "/grafana", + "prometheus_path": "/prometheus", + "webssh_path": "/webssh", + "valid_subscription": false + }, + "message": "success" + } + ``` + +### Ingest + +This API is used to ingest metrics from connected units. It requires basic authentication and +takes `firewall_api` as a parameter. +The `firewall_api` paramater is the name of the firewall API that is sending the metrics. +The API accepts only POST requests abd requires the following headers: + +- `Authorization:`: basic authentication header, where the username is the unit uuid and the password is the registration token +- `Content-Type: application/json`: the content type must be JSON + +It responds with a 200 status code in case of success. Success example: + +```json +{ "code": 200, "data": null, "message": "success" } +``` + +Possible error status codes are: + +- 400 if the request is malformed +- 401 if the authentication headers are missing or invalid +- 500 if there is an internal server error + +Error example: + +```json +{ "code": 401, "data": null, "message": "invalid unit id" } +``` + +- `POST /ingest/dump-nsplug-config` + + Create the unit record where all metrics are connected, it also stores the unit name in the report database. + This endpoint is mandatory and must be called at least once before sending all other metrics. + + REQ + + ```json + { "name": "fw.test.local" } + ``` + +- `POST /ingest/dump-mwan-events` + + Store all multiwan events in the report database. + + REQ + + ```json + { + "data": [ + { + "timestamp": 1726819981, + "wan": "wan", + "interface": "eth1", + "event": "online" + }, + { + "timestamp": 1726820241, + "wan": "wan2", + "interface": "eth2", + "event": "offline" + } + ] + } + ``` + +- `POST /ingest/dump-ts-attacks` + + Store all threat shield brute force blocks (fail2ban-like) in the report database. + + REQ + + ```json + { + "data": [ + { + "timestamp": 1726812650, + "ip": "200.91.234.36" + } + ] + } + ``` + +- `POST /ingest/dump-ts-malware` + + Store all threat shield blocks based on category in the report database. + + REQ + + ```json + { + "data": [ + { + "timestamp": 1726811160, + "src": "5.6.32.54", + "dst": "1.2.3.4", + "category": "nethesislvl3v4", + "chain": "inp-wan" + } + ] + } + ``` + +- `POST /ingest/dump-ovpn-connections` + + Store all openvpn connections in the report database. + + REQ + + ```json + { + "data": [ + { + "timestamp": 1726812276, + "instance": "ns_roadwarrior1", + "common_name": "user1", + "virtual_ip_addr": "10.9.10.41", + "remote_ip_addr": "1.2.3.4", + "start_time": 1726819476, + "duration": 4, + "bytes_received": 16343, + "bytes_sent": 7666 + } + ] + } + ``` + +- `POST /ingest/dump-dpi-stats` + + Store all network traffic stats in the report database. + + REQ + + ```json + { + "data": [ + { + "timestamp": 1726819203, + "client_address": "fe80::10ac:f709:5fb8:8fc3", + "client_name": "host1.test.local", + "protocol": "mdns", + "bytes": 123 + } + ] + } + ``` + + ``` + + ``` + +- `POST /ingest/dump-ovpn-config` + + Store the openvpn configuration in the report database. + + REQ + + ```json + { + "data": [ + { + "instance": "ns_roadwarrior1", + "device": "tunrw1", + "type": "rw", + "name": "srv1" + } + ] + } + ``` + +- `POST /ingest/dump-wan-config` + + REQ + + ```json + { + "data": [ + { "interface": "wan1", "device": "eth0", "status": "online" }, + { "interface": "wan2", "device": "eth5", "status": "offline" } + ] + } + ``` + +- `POST /ingest/info` + + This endpoint is used to store general information about the unit in the report database. It requires basic authentication and accepts the following headers: + + - `Authorization:`: basic authentication header, where the username is the unit UUID and the password is the registration token. + - `Content-Type: application/json`: the content type must be JSON. + + REQ + + ```json + { + "unit_name": "NethSec", + "version": "NethSecurity 8 24.10.0-ns.1.6.0", + "subscription_type": "", + "system_id": "", + "ssh_port": 22, + "fqdn": "NethSec", + "description": "This is my description", + "api_version": "3.2.0-r1", + "scheduled_update": -1, + "version_update": "NethSecurity 8-24.10.0-ns.1.6.0" + } + ``` + + RES + + ```json + { + "code": 200, + "data": null, + "message": "success" + } + ``` + + Possible error status codes: + + - 400 if the request is malformed. + - 401 if the authentication headers are missing or invalid. + - 500 if there is an internal server error. + + Error example: + + ```json + { + "code": 401, + "data": null, + "message": "invalid unit id" + } + ``` + +## Testing + +Execute tests with coverage: +``` +go test ./... -coverpkg=./... -coverprofile=coverage.out -v +``` diff --git a/controller/api/account_test.go b/controller/api/account_test.go new file mode 100644 index 00000000..d615106c --- /dev/null +++ b/controller/api/account_test.go @@ -0,0 +1,329 @@ +/* + * Copyright (C) 2025 Nethesis S.r.l. + * http://www.nethesis.it - info@nethesis.it + * + * SPDX-License-Identifier: GPL-2.0-only + * + * author: Giacomo Sanchietti + */ + +package main + +import ( + "bytes" + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/assert" +) + +// TestGetSSHKeys tests retrieving SSH keys for a user. +func TestGetSSHKeys(t *testing.T) { + gin.SetMode(gin.TestMode) + router := setupRouter() + + // First, login to get a valid token + var loginResp map[string]interface{} + w := httptest.NewRecorder() + loginBody := `{"username": "admin", "password": "admin"}` + req, _ := http.NewRequest("POST", "/login", bytes.NewBuffer([]byte(loginBody))) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusOK, w.Code, "login should succeed") + json.NewDecoder(w.Body).Decode(&loginResp) + assert.NotNil(t, loginResp["token"], "token should be present in login response") + token := loginResp["token"].(string) + + // Call GET /accounts/ssh-keys + w = httptest.NewRecorder() + req, _ = http.NewRequest("GET", "/accounts/ssh-keys", nil) + req.Header.Set("Authorization", "Bearer "+token) + router.ServeHTTP(w, req) + + // Should return 200 OK (even if no keys exist yet) + assert.Equal(t, http.StatusOK, w.Code, "GetSSHKeys should return 200 OK") + + var resp map[string]interface{} + json.NewDecoder(w.Body).Decode(&resp) + assert.Equal(t, float64(200), resp["code"], "response code should be 200") + assert.Equal(t, "success", resp["message"], "response message should be 'success'") +} + +// TestAddSSHKeys tests generating a new SSH key pair. +func TestAddSSHKeys(t *testing.T) { + gin.SetMode(gin.TestMode) + router := setupRouter() + + // First, login to get a valid token + var loginResp map[string]interface{} + w := httptest.NewRecorder() + loginBody := `{"username": "admin", "password": "admin"}` + req, _ := http.NewRequest("POST", "/login", bytes.NewBuffer([]byte(loginBody))) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusOK, w.Code, "login should succeed") + json.NewDecoder(w.Body).Decode(&loginResp) + token := loginResp["token"].(string) + + // Call POST /accounts/ssh-keys with passphrase + w = httptest.NewRecorder() + sshGenBody := `{"passphrase": "test-passphrase"}` + req, _ = http.NewRequest("POST", "/accounts/ssh-keys", bytes.NewBuffer([]byte(sshGenBody))) + req.Header.Set("Authorization", "Bearer "+token) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(w, req) + + // SSH key generation might fail if ssh-keygen is not available in test environment + // Only assert on successful response + if w.Code == http.StatusOK { + var resp map[string]interface{} + json.NewDecoder(w.Body).Decode(&resp) + assert.Equal(t, float64(200), resp["code"], "response code should be 200") + assert.Equal(t, "success", resp["message"], "response message should be 'success'") + data := resp["data"].(map[string]interface{}) + assert.NotEmpty(t, data["key_pub"], "key_pub should not be empty") + } +} + +// TestDeleteSSHKeys tests deleting SSH keys for a user. +func TestDeleteSSHKeys(t *testing.T) { + gin.SetMode(gin.TestMode) + router := setupRouter() + + // First, login to get a valid token + var loginResp map[string]interface{} + w := httptest.NewRecorder() + loginBody := `{"username": "admin", "password": "admin"}` + req, _ := http.NewRequest("POST", "/login", bytes.NewBuffer([]byte(loginBody))) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusOK, w.Code, "login should succeed") + json.NewDecoder(w.Body).Decode(&loginResp) + token := loginResp["token"].(string) + + // Call DELETE /accounts/ssh-keys + w = httptest.NewRecorder() + req, _ = http.NewRequest("DELETE", "/accounts/ssh-keys", nil) + req.Header.Set("Authorization", "Bearer "+token) + router.ServeHTTP(w, req) + + // Should return 200 OK + assert.Equal(t, http.StatusOK, w.Code, "DeleteSSHKeys should return 200 OK") + + var resp map[string]interface{} + json.NewDecoder(w.Body).Decode(&resp) + assert.Equal(t, float64(200), resp["code"], "response code should be 200") + assert.Equal(t, "success", resp["message"], "response message should be 'success'") +} + +// TestSSHKeyValidation tests that SSH key operations handle missing keys gracefully. +func TestSSHKeyValidation(t *testing.T) { + gin.SetMode(gin.TestMode) + router := setupRouter() + + // First, login to get a valid token + var loginResp map[string]interface{} + w := httptest.NewRecorder() + loginBody := `{"username": "admin", "password": "admin"}` + req, _ := http.NewRequest("POST", "/login", bytes.NewBuffer([]byte(loginBody))) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusOK, w.Code, "login should succeed") + json.NewDecoder(w.Body).Decode(&loginResp) + token := loginResp["token"].(string) + + // Delete any existing keys first + w = httptest.NewRecorder() + req, _ = http.NewRequest("DELETE", "/accounts/ssh-keys", nil) + req.Header.Set("Authorization", "Bearer "+token) + router.ServeHTTP(w, req) + + // Now try to get keys (should return empty without error) + w = httptest.NewRecorder() + req, _ = http.NewRequest("GET", "/accounts/ssh-keys", nil) + req.Header.Set("Authorization", "Bearer "+token) + router.ServeHTTP(w, req) + + assert.Equal(t, http.StatusOK, w.Code, "GetSSHKeys should return 200 OK even with missing keys") +} + +// TestUpdatePassword tests updating user password with correct old password. +func TestUpdatePassword(t *testing.T) { + gin.SetMode(gin.TestMode) + router := setupRouter() + + // First, login to get a valid token + var loginResp map[string]interface{} + w := httptest.NewRecorder() + loginBody := `{"username": "admin", "password": "admin"}` + req, _ := http.NewRequest("POST", "/login", bytes.NewBuffer([]byte(loginBody))) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusOK, w.Code, "login should succeed") + json.NewDecoder(w.Body).Decode(&loginResp) + token := loginResp["token"].(string) + + // Call PUT /accounts/password with old and new password + w = httptest.NewRecorder() + passChangeBody := `{"old_password": "admin", "new_password": "newpassword123"}` + req, _ = http.NewRequest("PUT", "/accounts/password", bytes.NewBuffer([]byte(passChangeBody))) + req.Header.Set("Authorization", "Bearer "+token) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(w, req) + + // Should return 200 OK + assert.Equal(t, http.StatusOK, w.Code, "UpdatePassword should return 200 OK with correct old password") + + var resp map[string]interface{} + json.NewDecoder(w.Body).Decode(&resp) + assert.Equal(t, float64(200), resp["code"], "response code should be 200") + assert.Equal(t, "success", resp["message"], "response message should be 'success'") + + // Reset password back to "admin" for other tests + w = httptest.NewRecorder() + passChangeBody = `{"old_password": "newpassword123", "new_password": "admin"}` + req, _ = http.NewRequest("PUT", "/accounts/password", bytes.NewBuffer([]byte(passChangeBody))) + req.Header.Set("Authorization", "Bearer "+token) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(w, req) +} + +// TestPasswordMismatch tests that password update fails with incorrect old password. +func TestPasswordMismatch(t *testing.T) { + gin.SetMode(gin.TestMode) + router := setupRouter() + + // First, login to get a valid token + var loginResp map[string]interface{} + w := httptest.NewRecorder() + loginBody := `{"username": "admin", "password": "admin"}` + req, _ := http.NewRequest("POST", "/login", bytes.NewBuffer([]byte(loginBody))) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusOK, w.Code, "login should succeed") + json.NewDecoder(w.Body).Decode(&loginResp) + token := loginResp["token"].(string) + + // Call PUT /accounts/password with wrong old password + w = httptest.NewRecorder() + passChangeBody := `{"old_password": "wrongpassword", "new_password": "newpassword123"}` + req, _ = http.NewRequest("PUT", "/accounts/password", bytes.NewBuffer([]byte(passChangeBody))) + req.Header.Set("Authorization", "Bearer "+token) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(w, req) + + // Should return 400 Bad Request + assert.Equal(t, http.StatusBadRequest, w.Code, "UpdatePassword should return 400 with incorrect old password") + + var resp map[string]interface{} + json.NewDecoder(w.Body).Decode(&resp) + assert.Equal(t, float64(400), resp["code"], "response code should be 400") + assert.Contains(t, resp["message"].(string), "mismatch", "response should indicate password mismatch") +} + +// TestGetAccountsAuthorizationForbidden tests that non-admin users cannot access /accounts endpoint. +func TestGetAccountsAuthorizationForbidden(t *testing.T) { + gin.SetMode(gin.TestMode) + router := setupRouter() + + // Create and login with limited user + // First, login as admin to create a limited user account + var adminLoginResp map[string]interface{} + w := httptest.NewRecorder() + adminLoginBody := `{"username": "admin", "password": "admin"}` + req, _ := http.NewRequest("POST", "/login", bytes.NewBuffer([]byte(adminLoginBody))) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusOK, w.Code, "admin login should succeed") + json.NewDecoder(w.Body).Decode(&adminLoginResp) + adminToken := adminLoginResp["token"].(string) + + // Create a limited user account + w = httptest.NewRecorder() + addBody := `{"username": "limiteduser", "password": "limited", "display_name": "Limited User", "admin": false}` + req, _ = http.NewRequest("POST", "/accounts", bytes.NewBuffer([]byte(addBody))) + req.Header.Set("Authorization", "Bearer "+adminToken) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(w, req) + + // Now login as the limited user + var limitedLoginResp map[string]interface{} + w = httptest.NewRecorder() + limitedLoginBody := `{"username": "limiteduser", "password": "limited"}` + req, _ = http.NewRequest("POST", "/login", bytes.NewBuffer([]byte(limitedLoginBody))) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusOK, w.Code, "limited user login should succeed") + json.NewDecoder(w.Body).Decode(&limitedLoginResp) + limitedToken := limitedLoginResp["token"].(string) + + // Try to access /accounts endpoint as limited user + w = httptest.NewRecorder() + req, _ = http.NewRequest("GET", "/accounts", nil) + req.Header.Set("Authorization", "Bearer "+limitedToken) + router.ServeHTTP(w, req) + + // Should return 403 Forbidden + assert.Equal(t, http.StatusForbidden, w.Code, "non-admin user should not access /accounts endpoint") + + var resp map[string]interface{} + json.NewDecoder(w.Body).Decode(&resp) + assert.Equal(t, float64(403), resp["code"], "response code should be 403") +} + +// TestAddSSHKeysInvalidRequest tests that AddSSHKeys rejects malformed requests. +func TestAddSSHKeysInvalidRequest(t *testing.T) { + gin.SetMode(gin.TestMode) + router := setupRouter() + + // First, login to get a valid token + var loginResp map[string]interface{} + w := httptest.NewRecorder() + loginBody := `{"username": "admin", "password": "admin"}` + req, _ := http.NewRequest("POST", "/login", bytes.NewBuffer([]byte(loginBody))) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusOK, w.Code, "login should succeed") + json.NewDecoder(w.Body).Decode(&loginResp) + token := loginResp["token"].(string) + + // Call POST /accounts/ssh-keys with malformed JSON + w = httptest.NewRecorder() + req, _ = http.NewRequest("POST", "/accounts/ssh-keys", bytes.NewBuffer([]byte("invalid json"))) + req.Header.Set("Authorization", "Bearer "+token) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(w, req) + + // Should return 400 Bad Request + assert.Equal(t, http.StatusBadRequest, w.Code, "AddSSHKeys should reject malformed JSON") +} + +// TestUpdatePasswordInvalidRequest tests that UpdatePassword rejects malformed requests. +func TestUpdatePasswordInvalidRequest(t *testing.T) { + gin.SetMode(gin.TestMode) + router := setupRouter() + + // First, login to get a valid token + var loginResp map[string]interface{} + w := httptest.NewRecorder() + loginBody := `{"username": "admin", "password": "admin"}` + req, _ := http.NewRequest("POST", "/login", bytes.NewBuffer([]byte(loginBody))) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusOK, w.Code, "login should succeed") + json.NewDecoder(w.Body).Decode(&loginResp) + token := loginResp["token"].(string) + + // Call PUT /accounts/password with malformed JSON + w = httptest.NewRecorder() + req, _ = http.NewRequest("PUT", "/accounts/password", bytes.NewBuffer([]byte("invalid json"))) + req.Header.Set("Authorization", "Bearer "+token) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(w, req) + + // Should return 400 Bad Request + assert.Equal(t, http.StatusBadRequest, w.Code, "UpdatePassword should reject malformed JSON") +} diff --git a/controller/api/configuration/configuration.go b/controller/api/configuration/configuration.go new file mode 100644 index 00000000..93448b8b --- /dev/null +++ b/controller/api/configuration/configuration.go @@ -0,0 +1,354 @@ +/* + * Copyright (C) 2024 Nethesis S.r.l. + * http://www.nethesis.it - info@nethesis.it + * + * SPDX-License-Identifier: GPL-2.0-only + * + * author: Edoardo Spadoni + */ + +package configuration + +import ( + "encoding/json" + "os" + "strconv" + "strings" + + "github.com/NethServer/nethsecurity-controller/api/logs" + "github.com/NethServer/nethsecurity-controller/api/models" + "github.com/Showmax/go-fqdn" +) + +type Configuration struct { + OpenVPNDir string `json:"openvpn_dir"` + OpenVPNNetwork string `json:"openvpn_network"` + OpenVPNNetmask string `json:"openvpn_netmask"` + OpenVPNUDPPort string `json:"openvpn_udp_port"` + + OpenVPNStatusDir string `json:"openvpn_status_dir"` // Deprecated: it can be removed in the future + OpenVPNCCDDir string `json:"openvpn_ccd_dir"` // Deprecated: it can be removed in the future + OpenVPNProxyDir string `json:"openvpn_proxy_dir"` + OpenVPNPKIDir string `json:"openvpn_pki_dir"` + OpenVPNMGMTSock string `json:"openvpn_mgmt_sock"` + + ListenAddress []string `json:"listen_address"` + + AdminUsername string `json:"admin_username"` + AdminPassword string `json:"admin_password"` + SecretJWT string `json:"secret_jwt"` + SensitiveList []string `json:"sensitive_list"` + RegistrationToken string `json:"registration_token"` + + CredentialsDir string `json:"credentials_dir"` // Deprecated: it can be removed in the future + DataDir string `json:"data_dir"` + Issuer2FA string `json:"issuer_2fa"` + SecretsDir string `json:"secrets_dir"` // Deprecated: it can be removed in the future + + PromtailAddress string `json:"promtail_address"` + PromtailPort string `json:"promtail_port"` + PrometheusPath string `json:"prometheus_path"` + WebSSHPath string `json:"webssh_path"` + GrafanaPath string `json:"grafana_path"` + + EasyRSAPath string `json:"easy_rsa_path"` + + ProxyProtocol string `json:"proxy_protocol"` + ProxyHost string `json:"proxy_host"` + ProxyPort string `json:"proxy_port"` + LoginEndpoint string `json:"login_endpoint"` + + FQDN string `json:"fqdn"` + + CacheTTL string `json:"cache_ttl"` + + ValidSubscription bool `json:"valid_subscription"` + + ReportDbUri string `json:"report_db_uri"` + + GeoIPDbDir string `json:"geoip_db_dir"` + MaxmindLicense string `json:"maxmind_license"` + + GrafanaPostgresPassword string `json:"grafana_postgres_password"` + + RetentionDays string `json:"retention_days"` + + EncryptionKey string `json:"encryption_key"` + + PlatformInfo models.PlatformInfo `json:"platform_info"` + + // Prometheus basi authenticatin to access target list + PrometheusAuthUsername string `json:"prometheus_auth_username"` + PrometheusAuthPassword string `json:"prometheus_auth_password"` + + // Generous global per-IP rate limit applied to every API route as a coarse + // safety net; 0 disables it + GlobalRateLimitAverage int `json:"global_rate_limit_average"` + GlobalRateLimitBurst int `json:"global_rate_limit_burst"` +} + +var Config = Configuration{} + +func Init() { + // read configuration from ENV + if os.Getenv("LISTEN_ADDRESS") != "" { + Config.ListenAddress = strings.Split(os.Getenv("LISTEN_ADDRESS"), ",") + } else { + Config.ListenAddress = []string{"127.0.0.1:5000"} + } + + if os.Getenv("ADMIN_USERNAME") != "" { + Config.AdminUsername = os.Getenv("ADMIN_USERNAME") + } else { + logs.Logs.Println("[CRITICAL][ENV] ADMIN_USERNAME variable is empty") + os.Exit(1) + } + if os.Getenv("ADMIN_PASSWORD") != "" { + Config.AdminPassword = os.Getenv("ADMIN_PASSWORD") + } else { + logs.Logs.Println("[CRITICAL][ENV] ADMIN_PASSWORD variable is empty") + os.Exit(1) + } + if os.Getenv("SECRET_JWT") != "" { + Config.SecretJWT = os.Getenv("SECRET_JWT") + } else { + logs.Logs.Println("[CRITICAL][ENV] SECRET_JWT variable is empty") + os.Exit(1) + } + if os.Getenv("SENSITIVE_LIST") != "" { + Config.SensitiveList = strings.Split(os.Getenv("SENSITIVE_LIST"), ",") + } else { + Config.SensitiveList = []string{"password", "secret", "token", "passphrase", "private", "key"} + } + if os.Getenv("REGISTRATION_TOKEN") != "" { + Config.RegistrationToken = os.Getenv("REGISTRATION_TOKEN") + } else { + logs.Logs.Println("[CRITICAL][ENV] REGISTRATION_TOKEN variable is empty") + os.Exit(1) + } + + if os.Getenv("CREDENTIALS_DIR") != "" { + Config.CredentialsDir = os.Getenv("CREDENTIALS_DIR") + } else { + logs.Logs.Println("[CRITICAL][ENV] CREDENTIALS_DIR variable is empty") + os.Exit(1) + } + if os.Getenv("DATA_DIR") != "" { + Config.DataDir = os.Getenv("DATA_DIR") + } else { + logs.Logs.Println("[CRITICAL][ENV] DATA_DIR variable is empty") + os.Exit(1) + } + + if os.Getenv("ISSUER_2FA") != "" { + Config.Issuer2FA = os.Getenv("ISSUER_2FA") + } else { + logs.Logs.Println("[CRITICAL][ENV] ISSUER_2FA variable is empty") + os.Exit(1) + } + + if os.Getenv("SECRETS_DIR") != "" { + Config.SecretsDir = os.Getenv("SECRETS_DIR") + } else { + logs.Logs.Println("[INFO][ENV] SECRETS_DIR variable is empty") + } + + if os.Getenv("OVPN_DIR") != "" { + Config.OpenVPNDir = os.Getenv("OVPN_DIR") + } else { + Config.OpenVPNDir = "/etc/openvpn" + } + if os.Getenv("OVPN_NETWORK") != "" { + Config.OpenVPNNetwork = os.Getenv("OVPN_NETWORK") + } else { + Config.OpenVPNNetwork = "172.21.0.0" + } + if os.Getenv("OVPN_NETMASK") != "" { + Config.OpenVPNNetmask = os.Getenv("OVPN_NETMASK") + } else { + Config.OpenVPNNetmask = "255.255.0.0" + } + if os.Getenv("OVPN_UDP_PORT") != "" { + Config.OpenVPNUDPPort = os.Getenv("OVPN_UDP_PORT") + } else { + Config.OpenVPNUDPPort = "1194" + } + + Config.OpenVPNStatusDir = Config.OpenVPNDir + "/status" + if os.Getenv("OVPN_C_DIR") != "" { + Config.OpenVPNCCDDir = os.Getenv("OVPN_C_DIR") + } else { + Config.OpenVPNCCDDir = Config.OpenVPNDir + "/ccd" + } + if os.Getenv("OVPN_P_DIR") != "" { + Config.OpenVPNProxyDir = os.Getenv("OVPN_P_DIR") + } else { + Config.OpenVPNProxyDir = Config.OpenVPNDir + "/proxy" + } + if os.Getenv("OVPN_K_DIR") != "" { + Config.OpenVPNPKIDir = os.Getenv("OVPN_K_DIR") + } else { + Config.OpenVPNPKIDir = Config.OpenVPNDir + "/pki" + } + if os.Getenv("OVPN_M_SOCK") != "" { + Config.OpenVPNMGMTSock = os.Getenv("OVPN_M_SOCK") + } else { + Config.OpenVPNMGMTSock = Config.OpenVPNDir + "/run/mgmt.sock" + } + + if os.Getenv("PROMTAIL_ADDRESS") != "" { + Config.PromtailAddress = os.Getenv("PROMTAIL_ADDRESS") + } else { + logs.Logs.Println("[CRITICAL][ENV] PROMTAIL_ADDRESS variable is empty") + os.Exit(1) + } + if os.Getenv("PROMTAIL_PORT") != "" { + Config.PromtailPort = os.Getenv("PROMTAIL_PORT") + } else { + logs.Logs.Println("[CRITICAL][ENV] PROMTAIL_PORT variable is empty") + os.Exit(1) + } + if os.Getenv("PROMETHEUS_PATH") != "" { + Config.PrometheusPath = os.Getenv("PROMETHEUS_PATH") + } else { + logs.Logs.Println("[CRITICAL][ENV] PROMETHEUS_PATH variable is empty") + os.Exit(1) + } + if os.Getenv("WEBSSH_PATH") != "" { + Config.WebSSHPath = os.Getenv("WEBSSH_PATH") + } else { + logs.Logs.Println("[CRITICAL][ENV] WEBSSH_PATH variable is empty") + os.Exit(1) + } + if os.Getenv("GRAFANA_PATH") != "" { + Config.GrafanaPath = os.Getenv("GRAFANA_PATH") + } else { + logs.Logs.Println("[CRITICAL][ENV] GRAFANA_PATH variable is empty") + os.Exit(1) + } + + if os.Getenv("EASYRSA_PATH") != "" { + Config.EasyRSAPath = os.Getenv("EASYRSA_PATH") + } else { + Config.EasyRSAPath = "/usr/share/easy-rsa/easyrsa" + } + + if os.Getenv("PROXY_PROTOCOL") != "" { + Config.ProxyProtocol = os.Getenv("PROXY_PROTOCOL") + } else { + Config.ProxyProtocol = "http://" + } + if os.Getenv("PROXY_HOST") != "" { + Config.ProxyHost = os.Getenv("PROXY_HOST") + } else { + Config.ProxyHost = "localhost" + } + if os.Getenv("PROXY_PORT") != "" { + Config.ProxyPort = os.Getenv("PROXY_PORT") + } else { + Config.ProxyPort = "8080" + } + if os.Getenv("LOGIN_ENDPOINT") != "" { + Config.LoginEndpoint = os.Getenv("LOGIN_ENDPOINT") + } else { + Config.LoginEndpoint = "/api/login" + } + + if os.Getenv("FQDN") != "" { + Config.FQDN = os.Getenv("FQDN") + } else { + Config.FQDN, _ = fqdn.FqdnHostname() + } + + if os.Getenv("CACHE_TTL") != "" { + Config.CacheTTL = os.Getenv("CACHE_TTL") + } else { + Config.CacheTTL = "7200" + } + + if os.Getenv("VALID_SUBSCRIPTION") != "" { + Config.ValidSubscription = os.Getenv("VALID_SUBSCRIPTION") == "true" + } else { + Config.ValidSubscription = false + } + + if os.Getenv("REPORT_DB_URI") != "" { + Config.ReportDbUri = os.Getenv("REPORT_DB_URI") + } else { + logs.Logs.Println("[CRITICAL][ENV] REPORT_DB_URI variable is empty") + os.Exit(1) + } + + // Assuming the file is named GeoLite2-Country.mmdb + if os.Getenv("GEOIP_DB_DIR") != "" { + Config.GeoIPDbDir = os.Getenv("GEOIP_DB_DIR") + } else { + Config.GeoIPDbDir = "." + } + + if os.Getenv("MAXMIND_LICENSE") != "" { + Config.MaxmindLicense = os.Getenv("MAXMIND_LICENSE") + } else { + logs.Logs.Println("[WARNING][ENV] MAXMIND_LICENSE variable is empty") + Config.MaxmindLicense = "" + } + + if os.Getenv("GRAFANA_POSTGRES_PASSWORD") != "" { + Config.GrafanaPostgresPassword = os.Getenv("GRAFANA_POSTGRES_PASSWORD") + } else { + logs.Logs.Println("[CRITICAL][ENV] GRAFANA_POSTGRES_PASSWORD variable is empty") + os.Exit(1) + } + + if os.Getenv("RETENTION_DAYS") != "" { + Config.RetentionDays = os.Getenv("RETENTION_DAYS") + } else { + Config.RetentionDays = "60" + } + + if os.Getenv("ENCRYPTION_KEY") != "" { + Config.EncryptionKey = os.Getenv("ENCRYPTION_KEY") + if len(Config.EncryptionKey) != 32 { + logs.Logs.Println("[CRITICAL][ENV] ENCRYPTION_KEY variable is not 32 bytes") + os.Exit(1) + } + } else { + logs.Logs.Println("[CRITICAL][ENV] ENCRYPTION_KEY variable is empty") + os.Exit(1) + } + + if os.Getenv("PLATFORM_INFO") != "" { + var platformInfo models.PlatformInfo + err := json.Unmarshal([]byte(os.Getenv("PLATFORM_INFO")), &platformInfo) + if err != nil { + logs.Logs.Println("[WARNING][ENV] PLATFORM_INFO variable is not valid JSON:", err) + } + Config.PlatformInfo = platformInfo + } else { + Config.PlatformInfo = models.PlatformInfo{} + } + + if os.Getenv("PROMETHEUS_AUTH_USERNAME") != "" { + Config.PrometheusAuthUsername = os.Getenv("PROMETHEUS_AUTH_USERNAME") + } else { + Config.PrometheusAuthUsername = "prometheus" + } + + if os.Getenv("PROMETHEUS_AUTH_PASSWORD") != "" { + Config.PrometheusAuthPassword = os.Getenv("PROMETHEUS_AUTH_PASSWORD") + } else { + Config.PrometheusAuthPassword = "prometheus" + } + + if v, err := strconv.Atoi(os.Getenv("GLOBAL_RATE_LIMIT_AVERAGE")); err == nil { + Config.GlobalRateLimitAverage = v + } else { + Config.GlobalRateLimitAverage = 25 + } + + if v, err := strconv.Atoi(os.Getenv("GLOBAL_RATE_LIMIT_BURST")); err == nil { + Config.GlobalRateLimitBurst = v + } else { + Config.GlobalRateLimitBurst = 100 + } +} diff --git a/controller/api/configuration/configuration_test.go b/controller/api/configuration/configuration_test.go new file mode 100644 index 00000000..5ed006c9 --- /dev/null +++ b/controller/api/configuration/configuration_test.go @@ -0,0 +1,178 @@ +/* + * Copyright (C) 2025 Nethesis S.r.l. + * http://www.nethesis.it - info@nethesis.it + * + * SPDX-License-Identifier: GPL-2.0-only + * + * author: Giacomo Sanchietti + */ + +package configuration + +import ( + "os" + "testing" + + "github.com/NethServer/nethsecurity-controller/api/logs" + "github.com/stretchr/testify/assert" +) + +func TestInit(t *testing.T) { + logs.Init("test") + // Save original config + originalConfig := Config + defer func() { + Config = originalConfig + }() + + // Clean up env vars + os.Unsetenv("ENCRYPTION_KEY") + os.Unsetenv("GRAFANA_POSTGRES_PASSWORD") + os.Unsetenv("REPORT_DB_URI") + os.Unsetenv("GRAFANA_PATH") + os.Unsetenv("WEBSSH_PATH") + os.Unsetenv("PROMETHEUS_PATH") + os.Unsetenv("PROMTAIL_PORT") + os.Unsetenv("PROMTAIL_ADDRESS") + os.Unsetenv("ISSUER_2FA") + os.Unsetenv("DATA_DIR") + os.Unsetenv("CREDENTIALS_DIR") + os.Unsetenv("REGISTRATION_TOKEN") + os.Unsetenv("SECRET_JWT") + os.Unsetenv("ADMIN_PASSWORD") + os.Unsetenv("ADMIN_USERNAME") + os.Unsetenv("LISTEN_ADDRESS") + os.Unsetenv("SENSITIVE_LIST") + os.Unsetenv("OVPN_DIR") + os.Unsetenv("SECRETS_DIR") + os.Unsetenv("FQDN") + os.Unsetenv("VALID_SUBSCRIPTION") + + // Set required env vars to avoid os.Exit + os.Setenv("ENCRYPTION_KEY", "12345678901234567890123456789012") + os.Setenv("GRAFANA_POSTGRES_PASSWORD", "grafana_pass") + os.Setenv("REPORT_DB_URI", "postgres://user:pass@localhost/db") + os.Setenv("GRAFANA_PATH", "/grafana") + os.Setenv("WEBSSH_PATH", "/webssh") + os.Setenv("PROMETHEUS_PATH", "/prometheus") + os.Setenv("PROMTAIL_PORT", "3100") + os.Setenv("PROMTAIL_ADDRESS", "localhost") + os.Setenv("ISSUER_2FA", "issuer") + os.Setenv("DATA_DIR", "/tmp/data") + os.Setenv("CREDENTIALS_DIR", "/tmp/creds") + os.Setenv("REGISTRATION_TOKEN", "token") + os.Setenv("SECRET_JWT", "secret") + os.Setenv("ADMIN_PASSWORD", "password") + os.Setenv("ADMIN_USERNAME", "admin") + + // Test with custom LISTEN_ADDRESS + os.Setenv("LISTEN_ADDRESS", "127.0.0.1:8080,0.0.0.0:9090") + os.Setenv("SENSITIVE_LIST", "pass,secret,key") + os.Setenv("OVPN_DIR", "/etc/openvpn_custom") + os.Setenv("SECRETS_DIR", "/tmp/secrets") + os.Setenv("FQDN", "example.com") + os.Setenv("VALID_SUBSCRIPTION", "true") + os.Setenv("GLOBAL_RATE_LIMIT_AVERAGE", "50") + os.Setenv("GLOBAL_RATE_LIMIT_BURST", "150") + + Init() + + assert.Equal(t, "12345678901234567890123456789012", Config.EncryptionKey) + assert.Equal(t, "grafana_pass", Config.GrafanaPostgresPassword) + assert.Equal(t, "postgres://user:pass@localhost/db", Config.ReportDbUri) + assert.Equal(t, "/grafana", Config.GrafanaPath) + assert.Equal(t, "/webssh", Config.WebSSHPath) + assert.Equal(t, "/prometheus", Config.PrometheusPath) + assert.Equal(t, "3100", Config.PromtailPort) + assert.Equal(t, "localhost", Config.PromtailAddress) + assert.Equal(t, "issuer", Config.Issuer2FA) + assert.Equal(t, "/tmp/data", Config.DataDir) + assert.Equal(t, "/tmp/creds", Config.CredentialsDir) + assert.Equal(t, "token", Config.RegistrationToken) + assert.Equal(t, "secret", Config.SecretJWT) + assert.Equal(t, "password", Config.AdminPassword) + assert.Equal(t, "admin", Config.AdminUsername) + assert.Equal(t, []string{"127.0.0.1:8080", "0.0.0.0:9090"}, Config.ListenAddress) + assert.Equal(t, []string{"pass", "secret", "key"}, Config.SensitiveList) + assert.Equal(t, "/etc/openvpn_custom", Config.OpenVPNDir) + assert.Equal(t, "/tmp/secrets", Config.SecretsDir) + assert.Equal(t, "example.com", Config.FQDN) + assert.True(t, Config.ValidSubscription) + assert.Equal(t, 50, Config.GlobalRateLimitAverage) + assert.Equal(t, 150, Config.GlobalRateLimitBurst) +} + +func TestInitDefaults(t *testing.T) { + logs.Init("test") + // Save original config + originalConfig := Config + defer func() { + Config = originalConfig + }() + + // Clean up env vars + os.Unsetenv("ENCRYPTION_KEY") + os.Unsetenv("GRAFANA_POSTGRES_PASSWORD") + os.Unsetenv("REPORT_DB_URI") + os.Unsetenv("GRAFANA_PATH") + os.Unsetenv("WEBSSH_PATH") + os.Unsetenv("PROMETHEUS_PATH") + os.Unsetenv("PROMTAIL_PORT") + os.Unsetenv("PROMTAIL_ADDRESS") + os.Unsetenv("ISSUER_2FA") + os.Unsetenv("DATA_DIR") + os.Unsetenv("CREDENTIALS_DIR") + os.Unsetenv("REGISTRATION_TOKEN") + os.Unsetenv("SECRET_JWT") + os.Unsetenv("ADMIN_PASSWORD") + os.Unsetenv("ADMIN_USERNAME") + os.Unsetenv("LISTEN_ADDRESS") + os.Unsetenv("SENSITIVE_LIST") + os.Unsetenv("OVPN_DIR") + os.Unsetenv("SECRETS_DIR") + os.Unsetenv("FQDN") + os.Unsetenv("VALID_SUBSCRIPTION") + os.Unsetenv("PROMETHEUS_AUTH_PASSWORD") + os.Unsetenv("PROMETHEUS_AUTH_USERNAME") + os.Unsetenv("RETENTION_DAYS") + os.Unsetenv("CACHE_TTL") + os.Unsetenv("OVPN_UDP_PORT") + os.Unsetenv("OVPN_NETMASK") + os.Unsetenv("OVPN_NETWORK") + os.Unsetenv("GLOBAL_RATE_LIMIT_AVERAGE") + os.Unsetenv("GLOBAL_RATE_LIMIT_BURST") + + // Set only required env vars + os.Setenv("ENCRYPTION_KEY", "12345678901234567890123456789012") + os.Setenv("GRAFANA_POSTGRES_PASSWORD", "grafana_pass") + os.Setenv("REPORT_DB_URI", "postgres://user:pass@localhost/db") + os.Setenv("GRAFANA_PATH", "/grafana") + os.Setenv("WEBSSH_PATH", "/webssh") + os.Setenv("PROMETHEUS_PATH", "/prometheus") + os.Setenv("PROMTAIL_PORT", "3100") + os.Setenv("PROMTAIL_ADDRESS", "localhost") + os.Setenv("ISSUER_2FA", "issuer") + os.Setenv("DATA_DIR", "/tmp/data") + os.Setenv("CREDENTIALS_DIR", "/tmp/creds") + os.Setenv("REGISTRATION_TOKEN", "token") + os.Setenv("SECRET_JWT", "secret") + os.Setenv("ADMIN_PASSWORD", "password") + os.Setenv("ADMIN_USERNAME", "admin") + + Init() + + // Check defaults + assert.Equal(t, "prometheus", Config.PrometheusAuthPassword) + assert.Equal(t, "prometheus", Config.PrometheusAuthUsername) + assert.Equal(t, "60", Config.RetentionDays) + assert.False(t, Config.ValidSubscription) + assert.Equal(t, "7200", Config.CacheTTL) + assert.Equal(t, "1194", Config.OpenVPNUDPPort) + assert.Equal(t, "255.255.0.0", Config.OpenVPNNetmask) + assert.Equal(t, "172.21.0.0", Config.OpenVPNNetwork) + assert.Equal(t, "/etc/openvpn", Config.OpenVPNDir) + assert.Equal(t, []string{"password", "secret", "token", "passphrase", "private", "key"}, Config.SensitiveList) + assert.Equal(t, []string{"127.0.0.1:5000"}, Config.ListenAddress) + assert.Equal(t, 25, Config.GlobalRateLimitAverage) + assert.Equal(t, 100, Config.GlobalRateLimitBurst) +} diff --git a/controller/api/entrypoint.sh b/controller/api/entrypoint.sh new file mode 100755 index 00000000..c6a845d3 --- /dev/null +++ b/controller/api/entrypoint.sh @@ -0,0 +1,31 @@ +#!/bin/sh + +mkdir -p /etc/openvpn/sockets + +cd /nethsecurity-api + +export ADMIN_USERNAME="${ADMIN_USERNAME:-admin}" +export ADMIN_PASSWORD="${ADMIN_PASSWORD:-8c6976e5b5410415bde908bd4dee15dfb167a9c873fc4bb8a81f6f2ab448a918}" # sha256sum of "admin" +export SECRET_JWT=$(head /dev/urandom | sha256sum) # regenerate SECRET at each restart to invalidate all tokens +export TOKENS_DIR="${TOKENS_DIR:-/nethsecurity-api/tokens}" +export CREDENTIALS_DIR="${CREDENTIALS_DIR:-/nethsecurity-api/credentials}" +export PROMTAIL_ADDRESS="${PROMTAIL_ADDRESS:-127.0.0.1}" +export PROMTAIL_PORT="${PROMTAIL_PORT:-9900}" + +socket=/etc/openvpn/run/mgmt.sock +limit=60 +while [ ! -e "$socket" ]; do + echo "Waiting for $socket to appear ..." + sleep 1 + limit=$((limit - 1)) + if [ "$limit" -le 0 ]; then + echo "Socket not found!" + break + fi +done + +# Create database config for OpenVPN hooks +echo REPORT_DB_URI=$REPORT_DB_URI > /etc/openvpn/conf.env +echo OVPN_NETMASK=$OVPN_NETMASK >> /etc/openvpn/conf.env + +exec "$@" diff --git a/controller/api/go.mod b/controller/api/go.mod new file mode 100644 index 00000000..cdc5c00b --- /dev/null +++ b/controller/api/go.mod @@ -0,0 +1,68 @@ +module github.com/NethServer/nethsecurity-controller/api + +go 1.26.0 + +toolchain go1.26.8 + +require ( + github.com/Jeffail/gabs/v2 v2.7.0 + github.com/Showmax/go-fqdn v1.0.0 + github.com/appleboy/gin-jwt/v2 v2.10.3 + github.com/fatih/structs v1.1.0 + github.com/gin-contrib/cors v1.7.9 + github.com/gin-contrib/gzip v1.2.8 + github.com/gin-gonic/gin v1.12.0 + github.com/golang-jwt/jwt/v5 v5.3.1 + github.com/google/uuid v1.6.0 + github.com/jackc/pgx/v5 v5.9.2 + github.com/mattn/go-sqlite3 v1.14.52 + github.com/nqd/flat v0.2.0 + github.com/oschwald/geoip2-golang v1.13.0 + github.com/pquerna/otp v1.5.0 + github.com/stretchr/testify v1.11.1 + golang.org/x/crypto v0.56.0 + golang.org/x/time v0.15.0 +) + +require ( + github.com/boombuler/barcode v1.0.1-0.20190219062509-6c824513bacc // indirect + github.com/bytedance/gopkg v0.1.4 // indirect + github.com/bytedance/sonic v1.15.2 // indirect + github.com/bytedance/sonic/loader v0.5.1 // indirect + github.com/cloudwego/base64x v0.1.7 // indirect + github.com/davecgh/go-spew v1.1.1 // indirect + github.com/gabriel-vasile/mimetype v1.4.13 // indirect + github.com/gin-contrib/sse v1.1.1 // indirect + github.com/go-playground/locales v0.14.1 // indirect + github.com/go-playground/universal-translator v0.18.1 // indirect + github.com/go-playground/validator/v10 v10.30.3 // indirect + github.com/goccy/go-json v0.10.6 // indirect + github.com/goccy/go-yaml v1.19.2 // indirect + github.com/golang-jwt/jwt/v4 v4.5.2 // indirect + github.com/imdario/mergo v0.3.12 // indirect + github.com/jackc/pgpassfile v1.0.0 // indirect + github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 // indirect + github.com/jackc/puddle/v2 v2.2.2 // indirect + github.com/json-iterator/go v1.1.12 // indirect + github.com/klauspost/cpuid/v2 v2.4.0 // indirect + github.com/leodido/go-urn v1.4.0 // indirect + github.com/mattn/go-isatty v0.0.23 // indirect + github.com/modern-go/concurrent v0.0.0-20180306012644-bacd9c7ef1dd // indirect + github.com/modern-go/reflect2 v1.0.2 // indirect + github.com/oschwald/maxminddb-golang v1.13.0 // indirect + github.com/pelletier/go-toml/v2 v2.4.3 // indirect + github.com/pmezard/go-difflib v1.0.0 // indirect + github.com/quic-go/qpack v0.6.0 // indirect + github.com/quic-go/quic-go v0.60.0 // indirect + github.com/twitchyliquid64/golang-asm v0.15.1 // indirect + github.com/ugorji/go/codec v1.3.1 // indirect + github.com/youmark/pkcs8 v0.0.0-20240726163527-a2c0da244d78 // indirect + go.mongodb.org/mongo-driver/v2 v2.8.0 // indirect + golang.org/x/arch v0.29.0 // indirect + golang.org/x/net v0.57.0 // indirect + golang.org/x/sync v0.22.0 // indirect + golang.org/x/sys v0.47.0 // indirect + golang.org/x/text v0.41.0 // indirect + google.golang.org/protobuf v1.36.11 // indirect + gopkg.in/yaml.v3 v3.0.1 // indirect +) diff --git a/controller/api/go.sum b/controller/api/go.sum new file mode 100644 index 00000000..f9a40264 --- /dev/null +++ b/controller/api/go.sum @@ -0,0 +1,210 @@ +github.com/Jeffail/gabs/v2 v2.7.0 h1:Y2edYaTcE8ZpRsR2AtmPu5xQdFDIthFG0jYhu5PY8kg= +github.com/Jeffail/gabs/v2 v2.7.0/go.mod h1:dp5ocw1FvBBQYssgHsG7I1WYsiLRtkUaB1FEtSwvNUw= +github.com/Showmax/go-fqdn v1.0.0 h1:0rG5IbmVliNT5O19Mfuvna9LL7zlHyRfsSvBPZmF9tM= +github.com/Showmax/go-fqdn v1.0.0/go.mod h1:SfrFBzmDCtCGrnHhoDjuvFnKsWjEQX/Q9ARZvOrJAko= +github.com/appleboy/gin-jwt/v2 v2.10.3 h1:KNcPC+XPRNpuoBh+j+rgs5bQxN+SwG/0tHbIqpRoBGc= +github.com/appleboy/gin-jwt/v2 v2.10.3/go.mod h1:LDUaQ8mF2W6LyXIbd5wqlV2SFebuyYs4RDwqMNgpsp8= +github.com/appleboy/gofight/v2 v2.1.2 h1:VOy3jow4vIK8BRQJoC/I9muxyYlJ2yb9ht2hZoS3rf4= +github.com/appleboy/gofight/v2 v2.1.2/go.mod h1:frW+U1QZEdDgixycTj4CygQ48yLTUhplt43+Wczp3rw= +github.com/boombuler/barcode v1.0.1-0.20190219062509-6c824513bacc h1:biVzkmvwrH8WK8raXaxBx6fRVTlJILwEwQGL1I/ByEI= +github.com/boombuler/barcode v1.0.1-0.20190219062509-6c824513bacc/go.mod h1:paBWMcWSl3LHKBqUq+rly7CNSldXjb2rDl3JlRe0mD8= +github.com/bytedance/gopkg v0.1.3 h1:TPBSwH8RsouGCBcMBktLt1AymVo2TVsBVCY4b6TnZ/M= +github.com/bytedance/gopkg v0.1.3/go.mod h1:576VvJ+eJgyCzdjS+c4+77QF3p7ubbtiKARP3TxducM= +github.com/bytedance/gopkg v0.1.4 h1:oZnQwnX82KAIWb7033bEwtxvTqXcYMxDBaQxo5JJHWM= +github.com/bytedance/gopkg v0.1.4/go.mod h1:v1zWfPm21Fb+OsyXN2VAHdL6TBb2L88anLQgdyje6R4= +github.com/bytedance/sonic v1.15.0 h1:/PXeWFaR5ElNcVE84U0dOHjiMHQOwNIx3K4ymzh/uSE= +github.com/bytedance/sonic v1.15.0/go.mod h1:tFkWrPz0/CUCLEF4ri4UkHekCIcdnkqXw9VduqpJh0k= +github.com/bytedance/sonic v1.15.2 h1:90H+rcF/FwLXwfB1cudOLq/je83n683Utf4Cbp0xHCo= +github.com/bytedance/sonic v1.15.2/go.mod h1:mT2NbXunuaEbnZ+mRIX/vYqKISmgEuHFDI4UzmKx2SA= +github.com/bytedance/sonic/loader v0.5.0 h1:gXH3KVnatgY7loH5/TkeVyXPfESoqSBSBEiDd5VjlgE= +github.com/bytedance/sonic/loader v0.5.0/go.mod h1:AR4NYCk5DdzZizZ5djGqQ92eEhCCcdf5x77udYiSJRo= +github.com/bytedance/sonic/loader v0.5.1 h1:Ygpfa9zwRCCKSlrp5bBP/b/Xzc3VxsAW+5NIYXrOOpI= +github.com/bytedance/sonic/loader v0.5.1/go.mod h1:AR4NYCk5DdzZizZ5djGqQ92eEhCCcdf5x77udYiSJRo= +github.com/cloudwego/base64x v0.1.6 h1:t11wG9AECkCDk5fMSoxmufanudBtJ+/HemLstXDLI2M= +github.com/cloudwego/base64x v0.1.6/go.mod h1:OFcloc187FXDaYHvrNIjxSe8ncn0OOM8gEHfghB2IPU= +github.com/cloudwego/base64x v0.1.7 h1:NppS+Fgzg5ovhn4NkUXaDT3x9jldgH5ToMCqzBSi2zI= +github.com/cloudwego/base64x v0.1.7/go.mod h1:Cu1PV9zfrSf7ET2tIbWbbEy7jO7HHJ13q4X2SQ8aWYg= +github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= +github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= +github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= +github.com/fatih/structs v1.1.0 h1:Q7juDM0QtcnhCpeyLGQKyg4TOIghuNXrkL32pHAUMxo= +github.com/fatih/structs v1.1.0/go.mod h1:9NiDSp5zOcgEDl+j00MP/WkGVPOlPRLejGD8Ga6PJ7M= +github.com/gabriel-vasile/mimetype v1.4.12 h1:e9hWvmLYvtp846tLHam2o++qitpguFiYCKbn0w9jyqw= +github.com/gabriel-vasile/mimetype v1.4.12/go.mod h1:d+9Oxyo1wTzWdyVUPMmXFvp4F9tea18J8ufA774AB3s= +github.com/gabriel-vasile/mimetype v1.4.13 h1:46nXokslUBsAJE/wMsp5gtO500a4F3Nkz9Ufpk2AcUM= +github.com/gabriel-vasile/mimetype v1.4.13/go.mod h1:d+9Oxyo1wTzWdyVUPMmXFvp4F9tea18J8ufA774AB3s= +github.com/gin-contrib/cors v1.7.7 h1:Oh9joP463x7Mw72vhvJ61YQm8ODh9b04YR7vsOErD0Q= +github.com/gin-contrib/cors v1.7.7/go.mod h1:K5tW0RkzJtWSiOdikXloy8VEZlgdVNpHNw8FpjUPNrE= +github.com/gin-contrib/cors v1.7.8 h1:U0jjyXlWXMEx27hFE6hmdg/TOZviYgsOGGT+kLmUZrc= +github.com/gin-contrib/cors v1.7.8/go.mod h1:u3nLI2pP1IlJn7tbvL7iiubDO9quBZ8FXP9LxmqpsPI= +github.com/gin-contrib/cors v1.7.9 h1:69rz5YU6PW7XlQ4VPXacl77ldFiUvR0E0bmEqmbq86w= +github.com/gin-contrib/cors v1.7.9/go.mod h1:KTqiTA5HzlWPgVz7p4tolzAuptzGO7DEJn/CbklG0nU= +github.com/gin-contrib/gzip v1.2.6 h1:OtN8DplD5DNZCSLAnQ5HxRkD2qZ5VU+JhOrcfJrcRvg= +github.com/gin-contrib/gzip v1.2.6/go.mod h1:BQy8/+JApnRjAVUplSGZiVtD2k8GmIE2e9rYu/hLzzU= +github.com/gin-contrib/gzip v1.2.7 h1:eQYOd81DpSU24TYYYNPzATrl7Hv3hGyzQilt3fGkxoc= +github.com/gin-contrib/gzip v1.2.7/go.mod h1:mfl5NDloGODrP2QryKtW37zsWrLkJkp/Y3iHTw0ZDy8= +github.com/gin-contrib/gzip v1.2.8 h1:wDb1thtVsSUe+126xHdgLiWcZqHdWXiBzdZj2yhdUM8= +github.com/gin-contrib/gzip v1.2.8/go.mod h1:OiNBR7FxAwHN3iwoDtjj2RlypuONCk48zmtagvASfrc= +github.com/gin-contrib/sse v1.1.0 h1:n0w2GMuUpWDVp7qSpvze6fAu9iRxJY4Hmj6AmBOU05w= +github.com/gin-contrib/sse v1.1.0/go.mod h1:hxRZ5gVpWMT7Z0B0gSNYqqsSCNIJMjzvm6fqCz9vjwM= +github.com/gin-contrib/sse v1.1.1 h1:uGYpNwTacv5R68bSGMapo62iLTRa9l5zxGCps4hK6ko= +github.com/gin-contrib/sse v1.1.1/go.mod h1:QXzuVkA0YO7o/gun03UI1Q+FTI8ZV/n5t03kIQAI89s= +github.com/gin-gonic/gin v1.12.0 h1:b3YAbrZtnf8N//yjKeU2+MQsh2mY5htkZidOM7O0wG8= +github.com/gin-gonic/gin v1.12.0/go.mod h1:VxccKfsSllpKshkBWgVgRniFFAzFb9csfngsqANjnLc= +github.com/go-playground/assert/v2 v2.2.0 h1:JvknZsQTYeFEAhQwI4qEt9cyV5ONwRHC+lYKSsYSR8s= +github.com/go-playground/assert/v2 v2.2.0/go.mod h1:VDjEfimB/XKnb+ZQfWdccd7VUvScMdVu0Titje2rxJ4= +github.com/go-playground/locales v0.14.1 h1:EWaQ/wswjilfKLTECiXz7Rh+3BjFhfDFKv/oXslEjJA= +github.com/go-playground/locales v0.14.1/go.mod h1:hxrqLVvrK65+Rwrd5Fc6F2O76J/NuW9t0sjnWqG1slY= +github.com/go-playground/universal-translator v0.18.1 h1:Bcnm0ZwsGyWbCzImXv+pAJnYK9S473LQFuzCbDbfSFY= +github.com/go-playground/universal-translator v0.18.1/go.mod h1:xekY+UJKNuX9WP91TpwSH2VMlDf28Uj24BCp08ZFTUY= +github.com/go-playground/validator/v10 v10.30.1 h1:f3zDSN/zOma+w6+1Wswgd9fLkdwy06ntQJp0BBvFG0w= +github.com/go-playground/validator/v10 v10.30.1/go.mod h1:oSuBIQzuJxL//3MelwSLD5hc2Tu889bF0Idm9Dg26cM= +github.com/go-playground/validator/v10 v10.30.3 h1:4MU6YkEwx7GbcPJOZxrtbu+QfF3pJLJuaYTeAH0DYy8= +github.com/go-playground/validator/v10 v10.30.3/go.mod h1:4Axh7oCNGcoGkqLoE4YWt6n20mcEIsPRlB7vPk3lpyc= +github.com/goccy/go-json v0.10.5 h1:Fq85nIqj+gXn/S5ahsiTlK3TmC85qgirsdTP/+DeaC4= +github.com/goccy/go-json v0.10.5/go.mod h1:oq7eo15ShAhp70Anwd5lgX2pLfOS3QCiwU/PULtXL6M= +github.com/goccy/go-json v0.10.6 h1:p8HrPJzOakx/mn/bQtjgNjdTcN+/S6FcG2CTtQOrHVU= +github.com/goccy/go-json v0.10.6/go.mod h1:oq7eo15ShAhp70Anwd5lgX2pLfOS3QCiwU/PULtXL6M= +github.com/goccy/go-yaml v1.19.2 h1:PmFC1S6h8ljIz6gMRBopkjP1TVT7xuwrButHID66PoM= +github.com/goccy/go-yaml v1.19.2/go.mod h1:XBurs7gK8ATbW4ZPGKgcbrY1Br56PdM69F7LkFRi1kA= +github.com/golang-jwt/jwt/v4 v4.5.2 h1:YtQM7lnr8iZ+j5q71MGKkNw9Mn7AjHM68uc9g5fXeUI= +github.com/golang-jwt/jwt/v4 v4.5.2/go.mod h1:m21LjoU+eqJr34lmDMbreY2eSTRJ1cv77w39/MY0Ch0= +github.com/golang-jwt/jwt/v5 v5.3.1 h1:kYf81DTWFe7t+1VvL7eS+jKFVWaUnK9cB1qbwn63YCY= +github.com/golang-jwt/jwt/v5 v5.3.1/go.mod h1:fxCRLWMO43lRc8nhHWY6LGqRcf+1gQWArsqaEUEa5bE= +github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8= +github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU= +github.com/google/gofuzz v1.0.0/go.mod h1:dBl0BpW6vV/+mYPU4Po3pmUjxk6FQPldtuIdl/M65Eg= +github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0= +github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= +github.com/imdario/mergo v0.3.12 h1:b6R2BslTbIEToALKP7LxUvijTsNI9TAe80pLWN2g/HU= +github.com/imdario/mergo v0.3.12/go.mod h1:jmQim1M+e3UYxmgPu/WyfjB3N3VflVyUjjjwH0dnCYA= +github.com/jackc/pgpassfile v1.0.0 h1:/6Hmqy13Ss2zCq62VdNG8tM1wchn8zjSGOBJ6icpsIM= +github.com/jackc/pgpassfile v1.0.0/go.mod h1:CEx0iS5ambNFdcRtxPj5JhEz+xB6uRky5eyVu/W2HEg= +github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 h1:iCEnooe7UlwOQYpKFhBabPMi4aNAfoODPEFNiAnClxo= +github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761/go.mod h1:5TJZWKEWniPve33vlWYSoGYefn3gLQRzjfDlhSJ9ZKM= +github.com/jackc/pgx/v5 v5.9.2 h1:3ZhOzMWnR4yJ+RW1XImIPsD1aNSz4T4fyP7zlQb56hw= +github.com/jackc/pgx/v5 v5.9.2/go.mod h1:mal1tBGAFfLHvZzaYh77YS/eC6IX9OWbRV1QIIM0Jn4= +github.com/jackc/puddle/v2 v2.2.2 h1:PR8nw+E/1w0GLuRFSmiioY6UooMp6KJv0/61nB7icHo= +github.com/jackc/puddle/v2 v2.2.2/go.mod h1:vriiEXHvEE654aYKXXjOvZM39qJ0q+azkZFrfEOc3H4= +github.com/json-iterator/go v1.1.12 h1:PV8peI4a0ysnczrg+LtxykD8LfKY9ML6u2jnxaEnrnM= +github.com/json-iterator/go v1.1.12/go.mod h1:e30LSqwooZae/UwlEbR2852Gd8hjQvJoHmT4TnhNGBo= +github.com/klauspost/cpuid/v2 v2.3.0 h1:S4CRMLnYUhGeDFDqkGriYKdfoFlDnMtqTiI/sFzhA9Y= +github.com/klauspost/cpuid/v2 v2.3.0/go.mod h1:hqwkgyIinND0mEev00jJYCxPNVRVXFQeu1XKlok6oO0= +github.com/klauspost/cpuid/v2 v2.4.0 h1:S6Hrbc7+ywsr0r+RLapfGBHfyefhCTwEh3A0tV913Dw= +github.com/klauspost/cpuid/v2 v2.4.0/go.mod h1:19jmZ9mjzoF//ddRSUsv0zfBTJWh3QJh9FNxZTMrGxU= +github.com/kr/pretty v0.3.1 h1:flRD4NNwYAUpkphVc1HcthR4KEIFJ65n8Mw5qdRn3LE= +github.com/kr/pretty v0.3.1/go.mod h1:hoEshYVHaxMs3cyo3Yncou5ZscifuDolrwPKZanG3xk= +github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY= +github.com/kr/text v0.2.0/go.mod h1:eLer722TekiGuMkidMxC/pM04lWEeraHUUmBw8l2grE= +github.com/leodido/go-urn v1.4.0 h1:WT9HwE9SGECu3lg4d/dIA+jxlljEa1/ffXKmRjqdmIQ= +github.com/leodido/go-urn v1.4.0/go.mod h1:bvxc+MVxLKB4z00jd1z+Dvzr47oO32F/QSNjSBOlFxI= +github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWEY= +github.com/mattn/go-isatty v0.0.20/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y= +github.com/mattn/go-isatty v0.0.23 h1:cYwCQTQf3HB6xUC+BtyCLZNr7IzbOmoZbmssVNzSyiQ= +github.com/mattn/go-isatty v0.0.23/go.mod h1:nMCL3Zebbrt45jsMDgnfIwz6ydEQApk5oEI3HqDio6A= +github.com/mattn/go-sqlite3 v1.14.47 h1:jOBI62gS7nKeZv+as1oGEy0+1qISgXwH/QBlR6KbfIo= +github.com/mattn/go-sqlite3 v1.14.47/go.mod h1:6JTjA44L93a0QCyJef5YvlPoKXntQPjzWv5gtm9sB6w= +github.com/mattn/go-sqlite3 v1.14.48 h1:7XHIgl0a8HwOaiK4E47ozLkST78rR9+OtNGx27D/TFs= +github.com/mattn/go-sqlite3 v1.14.48/go.mod h1:6JTjA44L93a0QCyJef5YvlPoKXntQPjzWv5gtm9sB6w= +github.com/mattn/go-sqlite3 v1.14.49 h1:B8jBHC3xhxZgxztrgruTuLucebnULQnx4W7cF7SAE9w= +github.com/mattn/go-sqlite3 v1.14.49/go.mod h1:6JTjA44L93a0QCyJef5YvlPoKXntQPjzWv5gtm9sB6w= +github.com/mattn/go-sqlite3 v1.14.50 h1:dmdFvo1XG4MPzA4IkAmE9upVz/Nj31uRoM5+jC8hYbY= +github.com/mattn/go-sqlite3 v1.14.50/go.mod h1:6JTjA44L93a0QCyJef5YvlPoKXntQPjzWv5gtm9sB6w= +github.com/mattn/go-sqlite3 v1.14.52 h1:wVbm2Qnf4OXkqhBTSPuCRZDRnxfbVrrmiCEroVdog8U= +github.com/mattn/go-sqlite3 v1.14.52/go.mod h1:6JTjA44L93a0QCyJef5YvlPoKXntQPjzWv5gtm9sB6w= +github.com/modern-go/concurrent v0.0.0-20180228061459-e0a39a4cb421/go.mod h1:6dJC0mAP4ikYIbvyc7fijjWJddQyLn8Ig3JB5CqoB9Q= +github.com/modern-go/concurrent v0.0.0-20180306012644-bacd9c7ef1dd h1:TRLaZ9cD/w8PVh93nsPXa1VrQ6jlwL5oN8l14QlcNfg= +github.com/modern-go/concurrent v0.0.0-20180306012644-bacd9c7ef1dd/go.mod h1:6dJC0mAP4ikYIbvyc7fijjWJddQyLn8Ig3JB5CqoB9Q= +github.com/modern-go/reflect2 v1.0.2 h1:xBagoLtFs94CBntxluKeaWgTMpvLxC4ur3nMaC9Gz0M= +github.com/modern-go/reflect2 v1.0.2/go.mod h1:yWuevngMOJpCy52FWWMvUC8ws7m/LJsjYzDa0/r8luk= +github.com/nqd/flat v0.2.0 h1:g6lXtMxsxrz6PZOO+rNnAJUn/GGRrK4FgVEhy/v+cHI= +github.com/nqd/flat v0.2.0/go.mod h1:FOuslZmNY082wVfVUUb7qAGWKl8z8Nor9FMg+Xj2Nss= +github.com/oschwald/geoip2-golang v1.13.0 h1:Q44/Ldc703pasJeP5V9+aFSZFmBN7DKHbNsSFzQATJI= +github.com/oschwald/geoip2-golang v1.13.0/go.mod h1:P9zG+54KPEFOliZ29i7SeYZ/GM6tfEL+rgSn03hYuUo= +github.com/oschwald/maxminddb-golang v1.13.0 h1:R8xBorY71s84yO06NgTmQvqvTvlS/bnYZrrWX1MElnU= +github.com/oschwald/maxminddb-golang v1.13.0/go.mod h1:BU0z8BfFVhi1LQaonTwwGQlsHUEu9pWNdMfmq4ztm0o= +github.com/pelletier/go-toml/v2 v2.2.4 h1:mye9XuhQ6gvn5h28+VilKrrPoQVanw5PMw/TB0t5Ec4= +github.com/pelletier/go-toml/v2 v2.2.4/go.mod h1:2gIqNv+qfxSVS7cM2xJQKtLSTLUE9V8t9Stt+h56mCY= +github.com/pelletier/go-toml/v2 v2.4.3 h1:GTRvJQutkOSftxIFD5xw9aepkYNuPWmVJpffdDPYVpY= +github.com/pelletier/go-toml/v2 v2.4.3/go.mod h1:2gIqNv+qfxSVS7cM2xJQKtLSTLUE9V8t9Stt+h56mCY= +github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= +github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= +github.com/pquerna/otp v1.5.0 h1:NMMR+WrmaqXU4EzdGJEE1aUUI0AMRzsp96fFFWNPwxs= +github.com/pquerna/otp v1.5.0/go.mod h1:dkJfzwRKNiegxyNb54X/3fLwhCynbMspSyWKnvi1AEg= +github.com/quic-go/qpack v0.6.0 h1:g7W+BMYynC1LbYLSqRt8PBg5Tgwxn214ZZR34VIOjz8= +github.com/quic-go/qpack v0.6.0/go.mod h1:lUpLKChi8njB4ty2bFLX2x4gzDqXwUpaO1DP9qMDZII= +github.com/quic-go/quic-go v0.59.1 h1:0Gmua0HW1Tv7ANR7hUYwRyD0MG5OJfgvYSZasGZzBic= +github.com/quic-go/quic-go v0.59.1/go.mod h1:upnsH4Ju1YkqpLXC305eW3yDZ4NfnNbmQRCMWS58IKU= +github.com/quic-go/quic-go v0.60.0 h1:xcQioE8OM66UQLeUMHltK1CCcOu3JbVB4JAQdDQSB+0= +github.com/quic-go/quic-go v0.60.0/go.mod h1:wpKpjmPpftl30sL6pFh7REVpjbcCVy4zt2vDyK1TuJk= +github.com/rogpeppe/go-internal v1.10.0 h1:TMyTOH3F/DB16zRVcYyreMH6GnZZrwQVAoYjRBZyWFQ= +github.com/rogpeppe/go-internal v1.10.0/go.mod h1:UQnix2H7Ngw/k4C5ijL5+65zddjncjaFoBhdsK/akog= +github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME= +github.com/stretchr/objx v0.4.0/go.mod h1:YvHI0jy2hoMjB+UWwv71VJQ9isScKT/TqJzVSSt89Yw= +github.com/stretchr/objx v0.5.0/go.mod h1:Yh+to48EsGEfYuaHDzXPcE3xhTkx73EhmCGUpEOglKo= +github.com/stretchr/objx v0.5.2/go.mod h1:FRsXN1f5AsAjCGJKqEizvkpNtU+EGNCLh3NxZ/8L+MA= +github.com/stretchr/testify v1.3.0/go.mod h1:M5WIy9Dh21IEIfnGCwXGc5bZfKNJtfHm1UVUgZn+9EI= +github.com/stretchr/testify v1.7.0/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg= +github.com/stretchr/testify v1.7.1/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg= +github.com/stretchr/testify v1.8.0/go.mod h1:yNjHg4UonilssWZ8iaSj1OCr/vHnekPRkoO+kdMU+MU= +github.com/stretchr/testify v1.8.4/go.mod h1:sz/lmYIOXD/1dqDmKjjqLyZ2RngseejIcXlSw2iwfAo= +github.com/stretchr/testify v1.10.0/go.mod h1:r2ic/lqez/lEtzL7wO/rwa5dbSLXVDPFyf8C91i36aY= +github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U= +github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U= +github.com/tidwall/gjson v1.17.1 h1:wlYEnwqAHgzmhNUFfw7Xalt2JzQvsMx2Se4PcoFCT/U= +github.com/tidwall/gjson v1.17.1/go.mod h1:/wbyibRr2FHMks5tjHJ5F8dMZh3AcwJEMf5vlfC0lxk= +github.com/tidwall/match v1.1.1 h1:+Ho715JplO36QYgwN9PGYNhgZvoUSc9X2c80KVTi+GA= +github.com/tidwall/match v1.1.1/go.mod h1:eRSPERbgtNPcGhD8UCthc6PmLEQXEWd3PRB5JTxsfmM= +github.com/tidwall/pretty v1.2.0 h1:RWIZEg2iJ8/g6fDDYzMpobmaoGh5OLl4AXtGUGPcqCs= +github.com/tidwall/pretty v1.2.0/go.mod h1:ITEVvHYasfjBbM0u2Pg8T2nJnzm8xPwvNhhsoaGGjNU= +github.com/twitchyliquid64/golang-asm v0.15.1 h1:SU5vSMR7hnwNxj24w34ZyCi/FmDZTkS4MhqMhdFk5YI= +github.com/twitchyliquid64/golang-asm v0.15.1/go.mod h1:a1lVb/DtPvCB8fslRZhAngC2+aY1QWCk3Cedj/Gdt08= +github.com/ugorji/go/codec v1.3.1 h1:waO7eEiFDwidsBN6agj1vJQ4AG7lh2yqXyOXqhgQuyY= +github.com/ugorji/go/codec v1.3.1/go.mod h1:pRBVtBSKl77K30Bv8R2P+cLSGaTtex6fsA2Wjqmfxj4= +github.com/youmark/pkcs8 v0.0.0-20240726163527-a2c0da244d78 h1:ilQV1hzziu+LLM3zUTJ0trRztfwgjqKnBWNtSRkbmwM= +github.com/youmark/pkcs8 v0.0.0-20240726163527-a2c0da244d78/go.mod h1:aL8wCCfTfSfmXjznFBSZNN13rSJjlIOI1fUNAtF7rmI= +go.mongodb.org/mongo-driver/v2 v2.5.0 h1:yXUhImUjjAInNcpTcAlPHiT7bIXhshCTL3jVBkF3xaE= +go.mongodb.org/mongo-driver/v2 v2.5.0/go.mod h1:yOI9kBsufol30iFsl1slpdq1I0eHPzybRWdyYUs8K/0= +go.mongodb.org/mongo-driver/v2 v2.8.0 h1:CxWDGQYY8QQwNjAl/aq2sfWakdnWZynnqJ9F4DhHbP8= +go.mongodb.org/mongo-driver/v2 v2.8.0/go.mod h1:yOI9kBsufol30iFsl1slpdq1I0eHPzybRWdyYUs8K/0= +go.uber.org/mock v0.6.0 h1:hyF9dfmbgIX5EfOdasqLsWD6xqpNZlXblLB/Dbnwv3Y= +go.uber.org/mock v0.6.0/go.mod h1:KiVJ4BqZJaMj4svdfmHM0AUx4NJYO8ZNpPnZn1Z+BBU= +golang.org/x/arch v0.23.0 h1:lKF64A2jF6Zd8L0knGltUnegD62JMFBiCPBmQpToHhg= +golang.org/x/arch v0.23.0/go.mod h1:dNHoOeKiyja7GTvF9NJS1l3Z2yntpQNzgrjh1cU103A= +golang.org/x/arch v0.29.0 h1:8sSET5wB0+exBm0FGmOtdHMqjlRdV2DRD3/IV6OZgho= +golang.org/x/arch v0.29.0/go.mod h1:0X+GdSIP+kL5wPmpK7sdkEVTt2XoYP0cSjQSbZBwOi8= +golang.org/x/crypto v0.53.0 h1:QZ4Muo8THX6CizN2vPPd5fBGHyogrdK9fG4wLPFUsto= +golang.org/x/crypto v0.53.0/go.mod h1:DNLU434OwVakk9PzuwV8w62mAJpRJL3vsgcfp4Qnsio= +golang.org/x/crypto v0.55.0 h1:+KWHjbgOaAQ66dh/YlkZKHlz9ZUlq61AFirAR9ntP8M= +golang.org/x/crypto v0.55.0/go.mod h1:uq0V9dE/fzQuJtbnL+2EhWOE63vo164FY8xqEnV9xis= +golang.org/x/crypto v0.56.0 h1:GUh5Ii4J5jtcseSMiRqr1jXCNHoxjeV9Fmekc2oLy6Y= +golang.org/x/crypto v0.56.0/go.mod h1:OMW5y6CY9l38uPLmxU6l6pwcXp1obtLo3e6gT7gQR2I= +golang.org/x/net v0.55.0 h1:bcvxaJn3e1U6InsFWt1JUq1aSjnRxLzT2rtD2KfkDF8= +golang.org/x/net v0.55.0/go.mod h1:L5U2KuzuOe1lY7Z+aWVIKK6qEeJXnXV9yzGA+WCHJww= +golang.org/x/net v0.57.0 h1:K5+3DljvIuDG9/Jv9rvyMywYNFCQ9RSUY6OOTTkT+tE= +golang.org/x/net v0.57.0/go.mod h1:KpXc8iv+r3XplLAG/f7Jsf9RPszJzdR0f58q9vGOuEU= +golang.org/x/sync v0.21.0 h1:HLII4xRRTtCRkxYp4HNFF0Js/Og6q2i++KXbg0gHCwM= +golang.org/x/sync v0.21.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0= +golang.org/x/sync v0.22.0 h1:SZjpbeLmrCk4xhRSZFNZW5gFUeCeFgjekvI/+gfScek= +golang.org/x/sync v0.22.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0= +golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= +golang.org/x/sys v0.46.0 h1:noSf2Fq6F8DBgS+LysIkx7rIExoNHJsxOAtPp4rthXw= +golang.org/x/sys v0.46.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= +golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs= +golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= +golang.org/x/text v0.38.0 h1:sXmwo9DwP3OK9EZ7PqAdaooSGozfl/3a6/xJcbzPRhE= +golang.org/x/text v0.38.0/go.mod h1:YXZt3QhHUKYT53r2lLKFIVi6Ao1jdzrTR/KQ09qyxF4= +golang.org/x/text v0.41.0 h1:vz/seA0lnX87Othu2f/0L24RcgrXD9/YFTSuGjj3rH8= +golang.org/x/text v0.41.0/go.mod h1:jvf1O8ajNzZqhSrQBPbutR/EB83Cc0CFrezNQIwbb5M= +golang.org/x/time v0.15.0 h1:bbrp8t3bGUeFOx08pvsMYRTCVSMk89u4tKbNOZbp88U= +golang.org/x/time v0.15.0/go.mod h1:Y4YMaQmXwGQZoFaVFk4YpCt4FLQMYKZe9oeV/f4MSno= +google.golang.org/protobuf v1.36.11 h1:fV6ZwhNocDyBLK0dj+fg8ektcVegBBuEolpbTQyBNVE= +google.golang.org/protobuf v1.36.11/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco= +gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= +gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c h1:Hei/4ADfdWqJk1ZMxUNpqntNwaWcugrBjAiHlqqRiVk= +gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c/go.mod h1:JHkPIbrfpd72SG/EVd6muEfDQjcINNoR0C8j2r3qZ4Q= +gopkg.in/yaml.v2 v2.3.0 h1:clyUAQHOM3G0M3f5vQj7LuJrETvjVot3Z5el9nffUtU= +gopkg.in/yaml.v2 v2.3.0/go.mod h1:hI93XBmqTisBFMUTm0b8Fm+jr3Dg1NNxqwp+5A1VGuI= +gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= +gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= +gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= diff --git a/controller/api/logs/logs.go b/controller/api/logs/logs.go new file mode 100644 index 00000000..9cf14dfe --- /dev/null +++ b/controller/api/logs/logs.go @@ -0,0 +1,25 @@ +/* + * Copyright (C) 2023 Nethesis S.r.l. + * http://www.nethesis.it - info@nethesis.it + * + * SPDX-License-Identifier: GPL-2.0-only + * + * author: Edoardo Spadoni + */ + +package logs + +import ( + "log" + "os" +) + +var Logs *log.Logger + +func Init(name string) { + // init syslog writer + logger := log.New(os.Stderr, name+" ", log.Ldate|log.Ltime|log.Lshortfile) + + // assign writer to Logs var + Logs = logger +} diff --git a/controller/api/logs/logs_test.go b/controller/api/logs/logs_test.go new file mode 100644 index 00000000..781ee62c --- /dev/null +++ b/controller/api/logs/logs_test.go @@ -0,0 +1,24 @@ +/* + * Copyright (C) 2025 Nethesis S.r.l. + * http://www.nethesis.it - info@nethesis.it + * + * SPDX-License-Identifier: GPL-2.0-only + * + * author: Giacomo Sanchietti + */ + +package logs + +import ( + "testing" + + "github.com/stretchr/testify/assert" +) + +func TestInit(t *testing.T) { + // Test that Init doesn't panic + assert.NotPanics(t, func() { + Init("test") + }) + assert.NotNil(t, Logs) +} diff --git a/controller/api/main.go b/controller/api/main.go new file mode 100644 index 00000000..b2ca42a9 --- /dev/null +++ b/controller/api/main.go @@ -0,0 +1,277 @@ +/* + * Copyright (C) 2024 Nethesis S.r.l. + * http://www.nethesis.it - info@nethesis.it + * + * SPDX-License-Identifier: GPL-2.0-only + * + * author: Edoardo Spadoni + */ + +package main + +import ( + "context" + "io" + "log" + "net/http" + "os" + "os/signal" + "syscall" + "time" + + "github.com/fatih/structs" + "github.com/gin-contrib/cors" + "github.com/gin-contrib/gzip" + "github.com/gin-gonic/gin" + "golang.org/x/time/rate" + + "github.com/NethServer/nethsecurity-controller/api/response" + "github.com/NethServer/nethsecurity-controller/api/configuration" + "github.com/NethServer/nethsecurity-controller/api/logs" + "github.com/NethServer/nethsecurity-controller/api/methods" + "github.com/NethServer/nethsecurity-controller/api/middleware" + "github.com/NethServer/nethsecurity-controller/api/routines" + "github.com/NethServer/nethsecurity-controller/api/socket" + "github.com/NethServer/nethsecurity-controller/api/storage" + "github.com/NethServer/nethsecurity-controller/api/utils" +) + +// @title NethSecurity Controller API Server +// @version 1.0 +// @description NethSecurity Controller API Server is used to manage multiple stand-alone NethSecurity instances +// @termsOfService https://nethserver.org/terms/ + +// @contact.name NethServer Developer Team +// @contact.url https://nethserver.org/support + +// @license.name GNU GENERAL PUBLIC LICENSE + +// @host localhost:5000 +// @schemes http +// @BasePath /api + +func setup() *gin.Engine { + // init logs with syslog + logs.Init("nethsecurity_controller") + + // init configuration + configuration.Init() + + // init storage + storage.Init() + + // init socket connection + socket.Init() + + // init geoip + utils.InitGeoIP() + // start geoip refresh loop + go routines.RefreshGeoIPDatabase() + + // starts remote info loop + go routines.RefreshRemoteInfoLoop() + + // disable log to stdout when running in release mode + if gin.Mode() == gin.ReleaseMode { + gin.DefaultWriter = io.Discard + } + + // init routers + router := gin.Default() + + // the app is only ever reached via the local reverse proxy (module Traefik, + // itself behind the NS8 core Traefik), both on loopback: trust only that + // hop's X-Forwarded-For so ClientIP() resolves the real client IP for + // RateLimiter instead of bucketing all traffic under 127.0.0.1 + router.SetTrustedProxies([]string{"127.0.0.1", "::1"}) + + // Generous global per-IP rate limit as a coarse safety net across all + // routes (a looser second layer behind the tighter per-route limits on + // the pre-auth routes). Set GLOBAL_RATE_LIMIT_AVERAGE=0 to disable. + if configuration.Config.GlobalRateLimitAverage > 0 { + router.Use(middleware.RateLimiter( + rate.Limit(configuration.Config.GlobalRateLimitAverage), + configuration.Config.GlobalRateLimitBurst, + )) + } + + // add default compression + router.Use(gzip.Gzip(gzip.DefaultCompression)) + + // cors configuration only in debug mode GIN_MODE=debug (default) + if gin.Mode() == gin.DebugMode { + // gin gonic cors conf + corsConf := cors.DefaultConfig() + corsConf.AllowHeaders = []string{"Authorization", "Content-Type", "Accept"} + corsConf.AllowAllOrigins = true + router.Use(cors.New(corsConf)) + } + + // define api group + api := router.Group("/") + + // define login and logout endpoint. BodyLimit caps each request's size; + // the global per-IP rate limiter (see above) bounds request frequency + // across every route, including these pre-authentication ones. + api.POST("/login", middleware.BodyLimit(32<<10), middleware.InstanceJWT().LoginHandler) + api.POST("/logout", middleware.BodyLimit(1<<10), middleware.InstanceJWT().LogoutHandler) + + // define healthcheck endpoint + api.GET("/health", func(c *gin.Context) { + c.JSON(http.StatusOK, gin.H{"status": "ok"}) + }) + + // 2FA APIs + api.POST("/2fa/otp-verify", middleware.BodyLimit(32<<10), methods.OTPVerify) + + // define server registration + api.POST("/units/register", middleware.BodyLimit(8<<10), methods.RegisterUnit) + + // define JWT middleware + api.Use(middleware.InstanceJWT().MiddlewareFunc()) + { + // refresh handler + api.GET("/refresh", middleware.InstanceJWT().RefreshHandler) + + // 2FA APIs + api.GET("/2fa", methods.Get2FAStatus) + api.DELETE("/2fa", methods.Del2FAStatus) + api.GET("/2fa/qr-code", methods.QRCode) + + // accounts APIs + accounts := api.Group("/accounts") + { + // accounts CRUD + accounts.GET("", methods.GetAccounts) + accounts.GET("/:account_id", methods.GetAccount) + accounts.POST("", methods.AddAccount) + accounts.PUT("/:account_id", methods.UpdateAccount) + accounts.DELETE("/:account_id", methods.DeleteAccount) + + // account password change + accounts.PUT("/password", methods.UpdatePassword) + + // ssh keys read and write + accounts.GET("/ssh-keys", methods.GetSSHKeys) + accounts.POST("/ssh-keys", methods.AddSSHKeys) + accounts.DELETE("/ssh-keys", methods.DeleteSSHKeys) + } + + // default APIs + defaults := api.Group("/defaults") + { + defaults.GET("", methods.GetDefaults) + } + + // units APIs + units := api.Group("/units") + { + units.GET("", methods.GetUnits) + units.GET("/:unit_id", methods.GetUnit) + units.GET("/:unit_id/info", methods.GetUnitInfo) + units.GET("/:unit_id/token", methods.GetToken) + units.POST("", methods.AddUnit) + units.DELETE("/:unit_id", methods.DeleteUnit) + } + + // unit_groups APIs + unitGroups := api.Group("/unit_groups") + { + unitGroups.GET("", methods.ListUnitGroups) + unitGroups.GET("/:group_id", methods.GetUnitGroup) + unitGroups.POST("", methods.AddUnitGroup) + unitGroups.PUT("/:group_id", methods.UpdateUnitGroup) + unitGroups.DELETE("/:group_id", methods.DeleteUnitGroup) + } + + // platforms APIs + api.GET("/platform", methods.GetPlatformInfo) + } + + // Ingest APIs: receive data from firewalls + authorized := router.Group("/ingest", middleware.BasicUnitAuth(), middleware.BodyLimit(8<<20)) // 8 Mib limit + authorized.POST("/info", methods.AddInfo) + authorized.POST("/:firewall_api", methods.HandelMonitoring) + + // Forwarded authentication middleware + forwarded := router.Group("/auth", middleware.BasicUserAuth()) + forwarded.GET("", func(c *gin.Context) { + c.Status(http.StatusOK) + }) + forwarded.GET("/:unit_id", func(c *gin.Context) { + c.Status(http.StatusOK) + }) + + // Prometheus metrics endpoint + prometheus := router.Group("/prometheus", gin.BasicAuth(gin.Accounts{ + configuration.Config.PrometheusAuthUsername: configuration.Config.PrometheusAuthPassword, + })) + prometheus.GET("/targets", methods.GetPrometheusTargets) + + // handle missing endpoint + router.NoRoute(func(c *gin.Context) { + c.JSON(http.StatusNotFound, structs.Map(response.StatusNotFound{ + Code: 404, + Message: "API not found", + Data: nil, + })) + }) + + return router +} + +func main() { + router := setup() + + // Create HTTP servers for each listen address + servers := make([]*http.Server, len(configuration.Config.ListenAddress)) + for i, addr := range configuration.Config.ListenAddress { + servers[i] = &http.Server{ + Addr: addr, + Handler: router, + } + } + + // Start servers in goroutines + for _, srv := range servers { + go func(server *http.Server) { + if err := server.ListenAndServe(); err != nil && err != http.ErrServerClosed { + log.Fatalf("listen: %s\n", err) + } + }(srv) + } + + // Wait for interrupt signal + // - SIGINT and SIGTERM for graceful shutdown + // - SIGUSR1 to reload ACLs + quit := make(chan os.Signal, 1) + signal.Notify(quit, syscall.SIGINT, syscall.SIGTERM, syscall.SIGUSR1) + + for { + sig := <-quit + + switch sig { + case syscall.SIGUSR1: + logs.Logs.Println("[INFO][API] Received SIGUSR1, reloading ACLs...") + storage.ReloadACLs() + // Continue running after ACL reload + case syscall.SIGINT, syscall.SIGTERM: + logs.Logs.Println("[INFO][API] Shutdown signal received, shutting down servers...") + + // The context is used to inform the server it has 5 seconds to finish + // the request it is currently handling + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + + // Shutdown all servers gracefully + for _, srv := range servers { + if err := srv.Shutdown(ctx); err != nil { + log.Fatal("Server forced to shutdown:", err) + } + } + + logs.Logs.Println("[INFO][API] Servers exiting") + return + } + } +} diff --git a/controller/api/main_test.go b/controller/api/main_test.go new file mode 100644 index 00000000..35cd30b2 --- /dev/null +++ b/controller/api/main_test.go @@ -0,0 +1,1240 @@ +/* + * Copyright (C) 2025 Nethesis S.r.l. + * http://www.nethesis.it - info@nethesis.it + * + * SPDX-License-Identifier: GPL-2.0-only + * + * author: Giacomo Sanchietti + */ + +package main + +import ( + "bytes" + "crypto/rand" + "encoding/json" + "fmt" + "net/http" + "net/http/httptest" + "os" + "strings" + "testing" + "time" + + mathrand "math/rand" + + "github.com/NethServer/nethsecurity-controller/api/configuration" + "github.com/NethServer/nethsecurity-controller/api/methods" + "github.com/NethServer/nethsecurity-controller/api/models" + "github.com/NethServer/nethsecurity-controller/api/storage" + "github.com/NethServer/nethsecurity-controller/api/utils" + "github.com/gin-gonic/gin" + "github.com/google/uuid" + "github.com/pquerna/otp/totp" + "github.com/stretchr/testify/assert" +) + +var router *gin.Engine + +// TestAESGCMEncryption tests AES-GCM encryption and decryption. +func TestAESGCMEncryption(t *testing.T) { + key := []byte("12345678901234567890123456789012") // AES-256, 32 byte + _, err := rand.Read(key) + if err != nil { + t.Fatalf("failed to generate random key: %v", err) + } + plaintext := []byte("Hello, AES-GCM encryption!") + + ciphertext, err := utils.EncryptAESGCM(plaintext, key) + if err != nil { + t.Fatalf("encryption failed: %v", err) + } + if bytes.Equal(ciphertext, plaintext) { + t.Error("ciphertext should not match plaintext") + } + + decrypted, err := utils.DecryptAESGCM(ciphertext, key) + if err != nil { + t.Fatalf("decryption failed: %v", err) + } + if !bytes.Equal(decrypted, plaintext) { + t.Errorf("decrypted text does not match original. got: %s, want: %s", decrypted, plaintext) + } + + // Test with wrong key + wrongKey := make([]byte, 32) + _, err = rand.Read(wrongKey) + if err != nil { + t.Fatalf("failed to generate wrong key: %v", err) + } + _, err = utils.DecryptAESGCM(ciphertext, wrongKey) + if err == nil { + t.Error("decryption should fail with wrong key") + } +} + +// TestAESGCMToString tests EncryptAESGCMToString and DecryptAESGCMFromString helpers. +func TestAESGCMToString(t *testing.T) { + key := []byte("12345678901234567890123456789012") // 32 bytes + + plaintext := []byte("Store this in DB as base64!") + + ciphertextB64, err := utils.EncryptAESGCMToString(plaintext, key) + if err != nil { + t.Fatalf("EncryptAESGCMToString failed: %v", err) + } + if ciphertextB64 == "" { + t.Error("ciphertextB64 should not be empty") + } + + decrypted, err := utils.DecryptAESGCMFromString(ciphertextB64, key) + if err != nil { + t.Fatalf("DecryptAESGCMFromString failed: %v", err) + } + if !bytes.Equal(decrypted, plaintext) { + t.Errorf("decrypted text does not match original. got: %s, want: %s", decrypted, plaintext) + } + + // Test with wrong key + wrongKey := []byte("abcdefghabcdefghabcdefghabcdefgh") // 32 bytes + _, err = utils.DecryptAESGCMFromString(ciphertextB64, wrongKey) + if err == nil { + t.Error("decryption should fail with wrong key") + } +} + +func TestMultipleListenAddresses(t *testing.T) { + gin.SetMode(gin.TestMode) + router = setupRouter() + + // Start two servers on different listeners + if len(configuration.Config.ListenAddress) < 2 { + t.Fatalf("expected at least 2 listen addresses, got %d", len(configuration.Config.ListenAddress)) + } + + servers := make([]*httptest.Server, 0, len(configuration.Config.ListenAddress)) + for range configuration.Config.ListenAddress { + // Use httptest.Server to simulate listening on multiple addresses + ts := httptest.NewServer(router) + servers = append(servers, ts) + } + defer func() { + for _, ts := range servers { + ts.Close() + } + }() + + // Test /health endpoint on all servers + for i, ts := range servers { + resp, err := http.Get(ts.URL + "/health") + if err != nil { + t.Fatalf("server %d: failed to GET /health: %v", i, err) + } + defer resp.Body.Close() + if resp.StatusCode != http.StatusOK { + t.Errorf("server %d: expected status 200, got %d", i, resp.StatusCode) + } + } +} + +// TestHealthEndpoint tests the /health endpoint. +func TestHealthEndpoint(t *testing.T) { + gin.SetMode(gin.TestMode) + router = setupRouter() + w := httptest.NewRecorder() + req, _ := http.NewRequest("GET", "/health", nil) + router.ServeHTTP(w, req) + if w.Code != http.StatusOK { + t.Fatalf("expected status 200, got %d", w.Code) + } + var resp map[string]interface{} + if err := json.NewDecoder(w.Body).Decode(&resp); err != nil { + t.Fatalf("failed to decode response: %v", err) + } + if resp["status"] != "ok" { + t.Errorf("expected status 'ok', got %v", resp["status"]) + } +} + +func TestMainEndpoints(t *testing.T) { + // Tests assume to run on a clean database, otherwise 2FA tests will fail + gin.SetMode(gin.TestMode) + router = setupRouter() + var token string + + t.Run("TestLoginEndpoint", func(t *testing.T) { + // Remove 2FA config from previous tests + os.RemoveAll(configuration.Config.SecretsDir + "/" + "admin") + w := httptest.NewRecorder() + var jsonResponse map[string]interface{} + body := `{"username": "admin", "password": "admin"}` + req, _ := http.NewRequest("POST", "/login", bytes.NewBuffer([]byte(body))) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(w, req) + json.NewDecoder(w.Body).Decode(&jsonResponse) + token = jsonResponse["token"].(string) + assert.Equal(t, http.StatusOK, w.Code) + assert.NotEmpty(t, token) + assert.True(t, methods.CheckTokenValidation("admin", token)) + }) + + t.Run("TestRefreshEndpoint", func(t *testing.T) { + w := httptest.NewRecorder() + req, _ := http.NewRequest("GET", "/refresh", nil) + req.Header.Set("Authorization", "Bearer "+token) + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusOK, w.Code) + }) + + t.Run("TestGet2FAStatusEndpoint", func(t *testing.T) { + w := httptest.NewRecorder() + req, _ := http.NewRequest("GET", "/2fa", nil) + req.Header.Set("Authorization", "Bearer "+token) + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusOK, w.Code) + }) + + t.Run("TestLogoutEndpoint", func(t *testing.T) { + w := httptest.NewRecorder() + req, _ := http.NewRequest("POST", "/logout", nil) + req.Header.Set("Authorization", "Bearer "+token) + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusOK, w.Code) + assert.False(t, methods.CheckTokenValidation("admin", token)) + }) + + t.Run("TestGetAccountsEndpoint", func(t *testing.T) { + // Login again + var jsonResponse map[string]interface{} + w := httptest.NewRecorder() + body := `{"username": "admin", "password": "admin"}` + req, _ := http.NewRequest("POST", "/login", bytes.NewBuffer([]byte(body))) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(w, req) + json.NewDecoder(w.Body).Decode(&jsonResponse) + token = jsonResponse["token"].(string) + + req, _ = http.NewRequest("GET", "/accounts", nil) + req.Header.Set("Authorization", "Bearer "+token) + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusOK, w.Code, w.Body.String()) + // response is: gin.H{"accounts": accounts, "total": len(accounts)}, + json.NewDecoder(w.Body).Decode(&jsonResponse) + data := jsonResponse["data"].(map[string]interface{}) + assert.Equal(t, data["accounts"].([]interface{})[0].(map[string]interface{})["username"], "admin") + assert.Equal(t, data["accounts"].([]interface{})[0].(map[string]interface{})["display_name"], "Administrator") + assert.Equal(t, data["accounts"].([]interface{})[0].(map[string]interface{})["two_fa"], false) + }) + + t.Run("TestAddUpdateDeleteAccount", func(t *testing.T) { + w := httptest.NewRecorder() + // Add account + addBody := `{"username": "testuser", "password": "testpass", "admin": false, "display_name": "Test User"}` + addReq, _ := http.NewRequest("POST", "/accounts", bytes.NewBuffer([]byte(addBody))) + addReq.Header.Set("Content-Type", "application/json") + addReq.Header.Set("Authorization", "Bearer "+token) + router.ServeHTTP(w, addReq) + assert.Equal(t, http.StatusCreated, w.Code) + var addResp map[string]interface{} + json.NewDecoder(w.Body).Decode(&addResp) + id := fmt.Sprintf("%v", addResp["data"].(map[string]interface{})["id"]) + assert.NotEmpty(t, id) + // Get accounts to find the new account's ID + w = httptest.NewRecorder() + getReq, _ := http.NewRequest("GET", "/accounts", nil) + getReq.Header.Set("Authorization", "Bearer "+token) + router.ServeHTTP(w, getReq) + var getResp map[string]interface{} + json.NewDecoder(w.Body).Decode(&getResp) + accounts := getResp["data"].(map[string]interface{})["accounts"].([]interface{}) + var testAccountID string + for _, acc := range accounts { + accMap := acc.(map[string]interface{}) + if accMap["username"] == "testuser" { + testAccountID = fmt.Sprintf("%v", accMap["id"]) + } + } + assert.NotEmpty(t, testAccountID) + // Update account display name + updateBody := `{"display_name": "Updated User", "unit_groups": [], "admin": false}` + updateReq, _ := http.NewRequest("PUT", "/accounts/"+testAccountID, bytes.NewBuffer([]byte(updateBody))) + updateReq.Header.Set("Content-Type", "application/json") + updateReq.Header.Set("Authorization", "Bearer "+token) + w = httptest.NewRecorder() + router.ServeHTTP(w, updateReq) + assert.Equal(t, http.StatusOK, w.Code) + // Get account and check display name + w = httptest.NewRecorder() + getOneReq, _ := http.NewRequest("GET", "/accounts/"+testAccountID, nil) + getOneReq.Header.Set("Authorization", "Bearer "+token) + router.ServeHTTP(w, getOneReq) + var getOneResp map[string]interface{} + json.NewDecoder(w.Body).Decode(&getOneResp) + accData := getOneResp["data"].(map[string]interface{})["account"].(map[string]interface{}) + assert.Equal(t, "Updated User", accData["display_name"]) + // Delete account + w = httptest.NewRecorder() + deleteReq, _ := http.NewRequest("DELETE", "/accounts/"+testAccountID, nil) + deleteReq.Header.Set("Authorization", "Bearer "+token) + router.ServeHTTP(w, deleteReq) + assert.Equal(t, http.StatusOK, w.Code) + // Ensure account is deleted + w = httptest.NewRecorder() + getOneReq, _ = http.NewRequest("GET", "/accounts/"+testAccountID, nil) + getOneReq.Header.Set("Authorization", "Bearer "+token) + router.ServeHTTP(w, getOneReq) + assert.Equal(t, http.StatusNotFound, w.Code) + }) + + t.Run("TestRegisterUnitEndpoint", func(t *testing.T) { + // create credentials directory + if _, err := os.Stat(configuration.Config.CredentialsDir); os.IsNotExist(err) { + if err := os.MkdirAll(configuration.Config.CredentialsDir, 0755); err != nil { + t.Fatalf("failed to create directory: %v", err) + } + } + // make sure configuration.Config.OpenVPNPKIDir does not exists + os.RemoveAll(configuration.Config.OpenVPNPKIDir) + unitID := "88860838-63bd-4717-a6c3-cbc351010843" + body := `{"unit_id": "` + unitID + `", "username": "myuser", "unit_name": "myname", "password": "mypassword"}` + req, _ := http.NewRequest("POST", "/units/register", bytes.NewBuffer([]byte(body))) + req.Header.Set("Content-Type", "application/json") + req.Header.Set("RegistrationToken", "1234") + w := httptest.NewRecorder() + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusForbidden, w.Code) + + // create OpenVPN directory + if _, err := os.Stat(configuration.Config.OpenVPNPKIDir + "/issued"); os.IsNotExist(err) { + if err := os.MkdirAll(configuration.Config.OpenVPNPKIDir+"/issued", 0755); err != nil { + t.Fatalf("failed to create directory: %v", err) + } + if err := os.MkdirAll(configuration.Config.OpenVPNPKIDir+"/private", 0755); err != nil { + t.Fatalf("failed to create directory: %v", err) + } + } + // create fake certificate file and key file + if _, err := os.Create(configuration.Config.OpenVPNPKIDir + "/issued/" + unitID + ".crt"); err != nil { + t.Fatalf("failed to create file: %v", err) + } + if _, err := os.Create(configuration.Config.OpenVPNPKIDir + "/private/" + unitID + ".key"); err != nil { + t.Fatalf("failed to create file: %v", err) + } + // create face ca.crt file + if _, err := os.Create(configuration.Config.OpenVPNPKIDir + "/ca.crt"); err != nil { + t.Fatalf("failed to create file: %v", err) + } + req, _ = http.NewRequest("POST", "/units/register", bytes.NewBuffer([]byte(body))) + req.Header.Set("Content-Type", "application/json") + req.Header.Set("RegistrationToken", "1234") + w = httptest.NewRecorder() + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusOK, w.Code, w.Body.String()) + + // Check password retrieval at lower level + user, pass, err := storage.GetUnitCredentials(unitID) // should return empty credentials + assert.NoError(t, err, "GetUnitCredentials should not return an error") + assert.Equal(t, "myuser", user) + assert.Equal(t, "mypassword", pass) + }) + + t.Run("TestNoRoute", func(t *testing.T) { + w := httptest.NewRecorder() + req, _ := http.NewRequest("GET", "/nonexistent", nil) + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusNotFound, w.Code) + }) + + // 2FA test: enable, verify with OTP, verify with recovery code, remove + t.Run("Test2FAEnableVerifyRemove", func(t *testing.T) { + w := httptest.NewRecorder() + // Execute login to get token + var jsonResponse map[string]interface{} + body := `{"username": "admin", "password": "admin"}` + req, _ := http.NewRequest("POST", "/login", bytes.NewBuffer([]byte(body))) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(w, req) + json.NewDecoder(w.Body).Decode(&jsonResponse) + token = jsonResponse["token"].(string) + assert.Equal(t, http.StatusOK, w.Code) + assert.NotEmpty(t, token) + + // Enable 2FA (get QR code and secret) + qrReq, _ := http.NewRequest("GET", "/2fa/qr-code", nil) + qrReq.Header.Set("Authorization", "Bearer "+token) + router.ServeHTTP(w, qrReq) + assert.Equal(t, http.StatusOK, w.Code) + var qrResp map[string]interface{} + json.NewDecoder(w.Body).Decode(&qrResp) + secret := qrResp["data"].(map[string]interface{})["key"].(string) + assert.NotEmpty(t, secret) + + otp, err := totp.GenerateCode(secret, time.Now()) + assert.NoError(t, err) + assert.NotEmpty(t, otp) + + // Verify 2FA login with OTP code + otpBody := map[string]string{"username": "admin", "token": token, "otp": otp} + otpBodyBytes, _ := json.Marshal(otpBody) + otpReq, _ := http.NewRequest("POST", "/2fa/otp-verify", bytes.NewBuffer(otpBodyBytes)) + otpReq.Header.Set("Content-Type", "application/json") + w = httptest.NewRecorder() + router.ServeHTTP(w, otpReq) + assert.Equal(t, http.StatusOK, w.Code) + + // Get recovery codes + w = httptest.NewRecorder() + statusReq, _ := http.NewRequest("GET", "/2fa", nil) + statusReq.Header.Set("Authorization", "Bearer "+token) + router.ServeHTTP(w, statusReq) + assert.Equal(t, http.StatusOK, w.Code) + var statusResp map[string]interface{} + json.NewDecoder(w.Body).Decode(&statusResp) + recoveryCodes := statusResp["data"].(map[string]interface{})["recovery_codes"].([]interface{}) + assert.NotEmpty(t, recoveryCodes) + recoveryCode := recoveryCodes[0].(string) + + // Verify 2FA login with recovery code + recBody := map[string]string{"username": "admin", "token": token, "otp": recoveryCode} + recBodyBytes, _ := json.Marshal(recBody) + recReq, _ := http.NewRequest("POST", "/2fa/otp-verify", bytes.NewBuffer(recBodyBytes)) + recReq.Header.Set("Content-Type", "application/json") + w = httptest.NewRecorder() + router.ServeHTTP(w, recReq) + assert.Equal(t, http.StatusOK, w.Code) + + // Remove 2FA + w = httptest.NewRecorder() + delReq, _ := http.NewRequest("DELETE", "/2fa", nil) + delReq.Header.Set("Authorization", "Bearer "+token) + router.ServeHTTP(w, delReq) + assert.Equal(t, http.StatusOK, w.Code) + + // Check 2FA is disabled + w = httptest.NewRecorder() + statusReq, _ = http.NewRequest("GET", "/2fa", nil) + statusReq.Header.Set("Authorization", "Bearer "+token) + router.ServeHTTP(w, statusReq) + assert.Equal(t, http.StatusOK, w.Code) + var statusResp2fa map[string]interface{} + json.NewDecoder(w.Body).Decode(&statusResp2fa) + assert.Equal(t, false, statusResp2fa["data"].(map[string]interface{})["status"]) + // recovery codes should be empty + assert.Equal(t, []interface{}{}, statusResp2fa["data"].(map[string]interface{})["recovery_codes"]) + }) + + // 2FA test partial setup (issue #1376) + t.Run("Test2FAPartialSetup", func(t *testing.T) { + w := httptest.NewRecorder() + // Execute login to get token + var jsonResponse map[string]interface{} + body := `{"username": "admin", "password": "admin"}` + req, _ := http.NewRequest("POST", "/login", bytes.NewBuffer([]byte(body))) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(w, req) + json.NewDecoder(w.Body).Decode(&jsonResponse) + token = jsonResponse["token"].(string) + assert.Equal(t, http.StatusOK, w.Code) + assert.NotEmpty(t, token) + + // Enable 2FA (get QR code and secret) + qrReq, _ := http.NewRequest("GET", "/2fa/qr-code", nil) + qrReq.Header.Set("Authorization", "Bearer "+token) + router.ServeHTTP(w, qrReq) + assert.Equal(t, http.StatusOK, w.Code) + var qrResp map[string]interface{} + json.NewDecoder(w.Body).Decode(&qrResp) + secret := qrResp["data"].(map[string]interface{})["key"].(string) + assert.NotEmpty(t, secret) + + otp, err := totp.GenerateCode(secret, time.Now()) + assert.NoError(t, err) + assert.NotEmpty(t, otp) + + // Check 2FA is disabled + w = httptest.NewRecorder() + statusReq, _ := http.NewRequest("GET", "/2fa", nil) + statusReq.Header.Set("Authorization", "Bearer "+token) + router.ServeHTTP(w, statusReq) + assert.Equal(t, http.StatusOK, w.Code) + var statusResp2fa map[string]interface{} + json.NewDecoder(w.Body).Decode(&statusResp2fa) + assert.Equal(t, false, statusResp2fa["data"].(map[string]interface{})["status"]) + // recovery codes should be empty + assert.Equal(t, []interface{}{}, statusResp2fa["data"].(map[string]interface{})["recovery_codes"]) + }) +} + +func addUnit(t *testing.T) string { + // Generate a UUID v4 and convert it to string using the uuid package + unitID := uuid.New().String() + + if _, err := os.Stat(configuration.Config.CredentialsDir); os.IsNotExist(err) { + os.MkdirAll(configuration.Config.CredentialsDir, 0755) + } + if _, err := os.Stat(configuration.Config.OpenVPNCCDDir); os.IsNotExist(err) { + os.MkdirAll(configuration.Config.OpenVPNCCDDir, 0755) + } + if _, err := os.Stat(configuration.Config.OpenVPNPKIDir); os.IsNotExist(err) { + os.MkdirAll(configuration.Config.OpenVPNPKIDir, 0755) + } + if _, err := os.Stat(configuration.Config.OpenVPNStatusDir); os.IsNotExist(err) { + os.MkdirAll(configuration.Config.OpenVPNStatusDir, 0755) + } + + // Create fake credentials, ccd and cr files, otherwise GetUnit will fail + creds := map[string]string{"username": "testuser", "password": "testpass"} + credsBytes, _ := json.Marshal(creds) + werr := os.WriteFile(configuration.Config.CredentialsDir+"/"+unitID, credsBytes, 0644) + assert.NoError(t, werr, "failed to write credentials file") + assert.NoError(t, werr, "failed to write ccd file") + if _, err := os.Stat(configuration.Config.OpenVPNPKIDir + "/issued/" + unitID + ".crt"); os.IsNotExist(err) { + if _, err := os.Create(configuration.Config.OpenVPNPKIDir + "/issued/" + unitID + ".crt"); err != nil { + t.Fatalf("failed to create certificate file: %v", err) + } + } + // Manually add to the database: we can't call /units POST endpoint because + // it requires the presence of easyrsa binary and configuration files + newIp := storage.GetFreeIP() + storage.AddUnit(unitID, newIp) + + return unitID +} + +func TestAddInfoAndGetRemoteInfo(t *testing.T) { + gin.SetMode(gin.TestMode) + router = setupRouter() + + // Simulate an add unit + unitID := addUnit(t) + + // AddInfo: POST /ingest/info (simulate BasicAuth middleware) + w := httptest.NewRecorder() + info := models.UnitInfo{ + UnitName: "my-test-unit", + Version: "1.0.0", + VersionUpdate: "1.0.1", + ScheduledUpdate: 0, + SubscriptionType: "test-subscription", + SystemID: "test-system-id", + SSHPort: 22, + FQDN: "test.example.com", + APIVersion: "v1", + } + infoBytes, _ := json.Marshal(info) + req := httptest.NewRequest("POST", "/ingest/info", bytes.NewBuffer(infoBytes)) + req.Header.Set("Content-Type", "application/json") + req.SetBasicAuth(unitID, "1234") + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusOK, w.Code, "AddInfo should return 200 OK") + + w = httptest.NewRecorder() + var jsonResponse map[string]interface{} + body := `{"username": "admin", "password": "admin"}` + req, _ = http.NewRequest("POST", "/login", bytes.NewBuffer([]byte(body))) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(w, req) + json.NewDecoder(w.Body).Decode(&jsonResponse) + token := jsonResponse["token"].(string) + + // Call /units/:unit_id to retrieve unit info + req = httptest.NewRequest("GET", "/units/"+unitID, nil) + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Authorization", "Bearer "+token) + router.ServeHTTP(w, req) + json.NewDecoder(w.Body).Decode(&jsonResponse) + infoResp := jsonResponse["data"].(map[string]interface{})["info"].(map[string]interface{}) + assert.Equal(t, info.UnitName, infoResp["unit_name"]) + assert.Equal(t, info.Version, infoResp["version"]) + assert.Equal(t, info.VersionUpdate, infoResp["version_update"]) + assert.Equal(t, float64(info.ScheduledUpdate), infoResp["scheduled_update"]) + assert.Equal(t, info.SubscriptionType, infoResp["subscription_type"]) + assert.Equal(t, info.SystemID, infoResp["system_id"]) + assert.Equal(t, float64(info.SSHPort), infoResp["ssh_port"]) + assert.Equal(t, info.FQDN, infoResp["fqdn"]) + assert.Equal(t, info.APIVersion, infoResp["api_version"]) + ipaddress := jsonResponse["data"].(map[string]interface{})["ipaddress"].(string) + assert.True(t, strings.HasPrefix(ipaddress, "172.21.0"), "ipaddress should start with 172.21.0, got: %v", ipaddress) + netmask := jsonResponse["data"].(map[string]interface{})["netmask"].(string) + assert.Equal(t, configuration.Config.OpenVPNNetmask, netmask, "OpenVPNNetmask should match the one in configuration, got: %v", netmask) +} + +func TestForwardedAuthMiddleware(t *testing.T) { + gin.SetMode(gin.TestMode) + router = setupRouter() + + // Test with valid credentials + w := httptest.NewRecorder() + req, _ := http.NewRequest("GET", "/auth", nil) + req.SetBasicAuth("admin", "admin") // Use BasicAuth for testing + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusOK, w.Code, w.Body.String()) + // Check if X-Auth-User header is set + authUser := w.Header().Get("X-Auth-User") + assert.Equal(t, "admin", authUser, "X-Auth-User header should be set to 'admin'") + + // Test with invalid credentials + w = httptest.NewRecorder() + req, _ = http.NewRequest("GET", "/auth", nil) + req.SetBasicAuth("admin", "wrongpassword") // Use BasicAuth for testing + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusUnauthorized, w.Code, w.Body.String()) +} + +func TestGetPlatformInfo(t *testing.T) { + router = setupRouter() + + // Step 1: Login and get token + loginBody := []byte(`{"username":"admin","password":"admin"}`) + w := httptest.NewRecorder() + req, _ := http.NewRequest("POST", "/login", bytes.NewBuffer(loginBody)) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusOK, w.Code) + var loginResp map[string]interface{} + _ = json.NewDecoder(w.Body).Decode(&loginResp) + token, ok := loginResp["token"].(string) + assert.True(t, ok) + assert.NotEmpty(t, token) + + // Step 2: Call GET /platform with token + w = httptest.NewRecorder() + req, _ = http.NewRequest("GET", "/platform", nil) + req.Header.Set("Authorization", "Bearer "+token) + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusOK, w.Code) + platformInfoEnv := os.Getenv("PLATFORM_INFO") + var platformInfo map[string]interface{} + err := json.Unmarshal([]byte(platformInfoEnv), &platformInfo) + assert.NoError(t, err) + // Step 3: Check response + var resp map[string]interface{} + _ = json.NewDecoder(w.Body).Decode(&resp) + assert.Equal(t, float64(200), resp["code"]) + assert.Equal(t, "success", resp["message"]) + data, ok := resp["data"].(map[string]interface{}) + assert.True(t, ok) + assert.Equal(t, "1194", data["vpn_port"]) + assert.Equal(t, "192.168.100.0/24", data["vpn_network"]) + assert.Equal(t, "1.0.0", data["controller_version"]) + assert.Equal(t, float64(30), data["metrics_retention_days"]) + assert.Equal(t, float64(90), data["logs_retention_days"]) +} + +func TestUnitGroupsAPI(t *testing.T) { + router = setupRouter() + unitId_1 := addUnit(t) + unitId_2 := addUnit(t) + randCounter := fmt.Sprintf("%d", mathrand.New(mathrand.NewSource(time.Now().UnixNano())).Intn(10000)) + + // Login to get token + w := httptest.NewRecorder() + loginBody := []byte(`{"username":"admin","password":"admin"}`) + req, _ := http.NewRequest("POST", "/login", bytes.NewBuffer(loginBody)) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusOK, w.Code) + var loginResp map[string]interface{} + _ = json.NewDecoder(w.Body).Decode(&loginResp) + token, ok := loginResp["token"].(string) + assert.True(t, ok) + assert.NotEmpty(t, token) + + // Create an empty unit group + w = httptest.NewRecorder() + // generate a random group name composed by testgroups + random number from 1 to 1000 + groupName := fmt.Sprintf("testgroups%s", randCounter) + groupBody := []byte(`{"name":"` + groupName + `","description":"desc"}`) + req, _ = http.NewRequest("POST", "/unit_groups", bytes.NewBuffer(groupBody)) + req.Header.Set("Authorization", "Bearer "+token) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusCreated, w.Code, w.Body.String()) + var groupResp map[string]interface{} + _ = json.NewDecoder(w.Body).Decode(&groupResp) + groupData := groupResp["data"].(map[string]interface{}) + groupID := fmt.Sprintf("%v", groupData["id"]) + assert.NotEmpty(t, groupID) + + // Update the group with units + w = httptest.NewRecorder() + updateBody := []byte(`{"name":"` + groupName + `","description":"desc", "units":["` + unitId_1 + `", "` + unitId_2 + `"]}`) + req, _ = http.NewRequest("PUT", "/unit_groups/"+groupID, bytes.NewBuffer([]byte(updateBody))) + req.Header.Set("Authorization", "Bearer "+token) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusOK, w.Code, w.Body.String()) + + // List unit groups + w = httptest.NewRecorder() + req, _ = http.NewRequest("GET", "/unit_groups", nil) + req.Header.Set("Authorization", "Bearer "+token) + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusOK, w.Code) + + // Get the created unit group + w = httptest.NewRecorder() + req, _ = http.NewRequest("GET", "/unit_groups/"+groupID, nil) + req.Header.Set("Authorization", "Bearer "+token) + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusOK, w.Code) + var getGroupResp map[string]interface{} + _ = json.NewDecoder(w.Body).Decode(&getGroupResp) + groupData = getGroupResp["data"].(map[string]interface{}) + assert.Equal(t, groupName, groupData["name"]) + assert.Equal(t, "desc", groupData["description"]) + units := groupData["units"].([]interface{}) + assert.Len(t, units, 2) + assert.Contains(t, units, unitId_1) + assert.Contains(t, units, unitId_2) + + // Update the unit group + w = httptest.NewRecorder() + updateBody = []byte(`{"name":"updatedgroup` + randCounter + `","description":"updated desc", "units":["` + unitId_1 + `", "` + unitId_2 + `"]}}`) + req, _ = http.NewRequest("PUT", "/unit_groups/"+groupID, bytes.NewBuffer(updateBody)) + req.Header.Set("Authorization", "Bearer "+token) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusOK, w.Code) + // Get the updated unit group and check the new name and description + w = httptest.NewRecorder() + req, _ = http.NewRequest("GET", "/unit_groups/"+groupID, nil) + req.Header.Set("Authorization", "Bearer "+token) + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusOK, w.Code) + var updatedGroupResp map[string]interface{} + _ = json.NewDecoder(w.Body).Decode(&updatedGroupResp) + updatedGroupData := updatedGroupResp["data"].(map[string]interface{}) + assert.Equal(t, "updatedgroup"+randCounter, updatedGroupData["name"]) + assert.Equal(t, "updated desc", updatedGroupData["description"]) + + // Try to update the group with a non-existing unit, expect failure (400) + w = httptest.NewRecorder() + nonExistingUnitID := uuid.New().String() + updateBody = []byte(`{"name":"` + groupName + `","description":"desc", "units":["` + unitId_1 + `", "` + nonExistingUnitID + `"]}`) + req, _ = http.NewRequest("PUT", "/unit_groups/"+groupID, bytes.NewBuffer(updateBody)) + req.Header.Set("Authorization", "Bearer "+token) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusBadRequest, w.Code, "should fail when adding non-existing unit to group") + + // Add limited user account + w = httptest.NewRecorder() + limitedUserName := fmt.Sprintf("limited%s", randCounter) + addBody := `{"username": "` + limitedUserName + `", "password": "limited", "display_name": "Limited user"}` + addReq, _ := http.NewRequest("POST", "/accounts", bytes.NewBuffer([]byte(addBody))) + addReq.Header.Set("Content-Type", "application/json") + addReq.Header.Set("Authorization", "Bearer "+token) + router.ServeHTTP(w, addReq) + assert.Equal(t, http.StatusCreated, w.Code) + var addUserResp map[string]interface{} + json.NewDecoder(w.Body).Decode(&addUserResp) + testAccountID := fmt.Sprintf("%v", addUserResp["data"].(map[string]interface{})["id"]) + assert.NotEmpty(t, testAccountID) + + // Try to add a non-existing group ID to the user account, expect failure + w = httptest.NewRecorder() + nonExistingGroupID := "9999999" + addUserBody := []byte(`{"username":"` + limitedUserName + `","unit_groups":[` + nonExistingGroupID + `]}`) + req, _ = http.NewRequest("PUT", "/accounts/"+testAccountID, bytes.NewBuffer(addUserBody)) + req.Header.Set("Authorization", "Bearer "+token) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusBadRequest, w.Code, "should fail when adding non-existing group ID") + + // Delete unitId_1, it must be removed from the group + w = httptest.NewRecorder() + req, _ = http.NewRequest("DELETE", "/units/"+unitId_1, nil) + req.Header.Set("Authorization", "Bearer "+token) + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusOK, w.Code, "deleting unitId_1 should succeed") + // Get the unit group again and check that unitId_1 is removed + w = httptest.NewRecorder() + req, _ = http.NewRequest("GET", "/unit_groups/"+groupID, nil) + req.Header.Set("Authorization", "Bearer "+token) + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusOK, w.Code) + var afterDeleteUpdatedGroupResp map[string]interface{} + _ = json.NewDecoder(w.Body).Decode(&afterDeleteUpdatedGroupResp) + afterDeleteUpdatedGroupData := afterDeleteUpdatedGroupResp["data"].(map[string]interface{}) + units = afterDeleteUpdatedGroupData["units"].([]interface{}) + assert.Len(t, units, 1, "unitId_1 should be removed from the group") + assert.Contains(t, units, unitId_2, "unitId_2 should still be in the group") + assert.NotContains(t, units, unitId_1, "unitId_1 should not be in the group anymore") + + // Add unit group to user account + w = httptest.NewRecorder() + addUserBody = []byte(`{"username":"` + limitedUserName + `","unit_groups":[` + groupID + `]}`) + req, _ = http.NewRequest("PUT", "/accounts/"+testAccountID, bytes.NewBuffer(addUserBody)) + req.Header.Set("Authorization", "Bearer "+token) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusOK, w.Code, w.Body.String()) + + // Delete the unit group: should fail because it is associated with an account + w = httptest.NewRecorder() + req, _ = http.NewRequest("DELETE", "/unit_groups/"+groupID, nil) + req.Header.Set("Authorization", "Bearer "+token) + router.ServeHTTP(w, req) + assert.NotEqual(t, http.StatusOK, w.Code, "should not allow deleting a group associated with an account") + assert.True(t, w.Code == http.StatusBadRequest, "expected 400 when deleting a group in use") + + // Remove the group from the user account's unit_groups + w = httptest.NewRecorder() + removeGroupBody := []byte(`{"username":"` + limitedUserName + `","unit_groups":[]}`) + req, _ = http.NewRequest("PUT", "/accounts/"+testAccountID, bytes.NewBuffer(removeGroupBody)) + req.Header.Set("Authorization", "Bearer "+token) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusOK, w.Code, w.Body.String()) + + // Now delete the unit group again, should succeed + w = httptest.NewRecorder() + req, _ = http.NewRequest("DELETE", "/unit_groups/"+groupID, nil) + req.Header.Set("Authorization", "Bearer "+token) + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusOK, w.Code, "should allow deleting a group not associated with any account") + + // Get the account again and check that unit_groups does not contain groupID + w = httptest.NewRecorder() + req, _ = http.NewRequest("GET", "/accounts/"+testAccountID, nil) + req.Header.Set("Authorization", "Bearer "+token) + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusOK, w.Code) + var getAccountAfterDeleteResp map[string]interface{} + _ = json.NewDecoder(w.Body).Decode(&getAccountAfterDeleteResp) + accountAfterDeleteData := getAccountAfterDeleteResp["data"].(map[string]interface{})["account"].(map[string]interface{}) + unitGroups := accountAfterDeleteData["unit_groups"].([]interface{}) + for _, v := range unitGroups { + assert.NotEqual(t, groupID, fmt.Sprintf("%.0f", v), "unit_groups should not contain deleted groupID") + } + + // Create a new group and add unitId_1 to it + w = httptest.NewRecorder() + groupName2 := fmt.Sprintf("group2_%s", randCounter) + groupBody2 := []byte(`{"name":"` + groupName2 + `","description":"desc2", "units":["` + unitId_1 + `"]}`) + req, _ = http.NewRequest("POST", "/unit_groups", bytes.NewBuffer(groupBody2)) + req.Header.Set("Authorization", "Bearer "+token) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusCreated, w.Code, w.Body.String()) + var group2Resp map[string]interface{} + _ = json.NewDecoder(w.Body).Decode(&group2Resp) + group2Data := group2Resp["data"].(map[string]interface{}) + group2ID := fmt.Sprintf("%v", group2Data["id"]) + assert.NotEmpty(t, group2ID) + + // Get the group and check that unitId_1 is present + w = httptest.NewRecorder() + req, _ = http.NewRequest("GET", "/unit_groups/"+group2ID, nil) + req.Header.Set("Authorization", "Bearer "+token) + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusOK, w.Code) + var getGroup3Resp map[string]interface{} + _ = json.NewDecoder(w.Body).Decode(&getGroup3Resp) + group2Data = getGroup3Resp["data"].(map[string]interface{}) + units2 := group2Data["units"].([]interface{}) + assert.Len(t, units2, 1, "group2 should contain exactly one unit") + assert.Equal(t, unitId_1, fmt.Sprintf("%v", units2[0]), "unitId_1 should be present in group2") + + // DELETE "/units/"+unitId_1 can't be tested because it requires easy-rsa binary and configuration file +} + +func TestUnitAuthorization(t *testing.T) { + router = setupRouter() + unitId_1 := addUnit(t) + unitId_2 := addUnit(t) + + randCounter := fmt.Sprintf("%d", mathrand.New(mathrand.NewSource(time.Now().UnixNano())).Intn(10000)) + + // Login to get token + w := httptest.NewRecorder() + loginBody := []byte(`{"username":"admin","password":"admin"}`) + req, _ := http.NewRequest("POST", "/login", bytes.NewBuffer(loginBody)) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusOK, w.Code) + var loginResp map[string]interface{} + _ = json.NewDecoder(w.Body).Decode(&loginResp) + token, ok := loginResp["token"].(string) + assert.True(t, ok) + assert.NotEmpty(t, token) + + // Create unit group with unitId_1 + w = httptest.NewRecorder() + groupName := fmt.Sprintf("authgroup%s", randCounter) + groupBody := []byte(`{"name":"` + groupName + `","description":"auth test group", "units":["` + unitId_1 + `"]}`) + req, _ = http.NewRequest("POST", "/unit_groups", bytes.NewBuffer(groupBody)) + req.Header.Set("Authorization", "Bearer "+token) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusCreated, w.Code, w.Body.String()) + var groupResp map[string]interface{} + _ = json.NewDecoder(w.Body).Decode(&groupResp) + groupData := groupResp["data"].(map[string]interface{}) + groupID := fmt.Sprintf("%v", groupData["id"]) + assert.NotEmpty(t, groupID) + + // Create limited account associated to the unit group + w = httptest.NewRecorder() + limitedUserName := fmt.Sprintf("limited%s", randCounter) + addBody := `{"username": "` + limitedUserName + `", "password": "limited", "display_name": "Limited user", "unit_groups": [` + groupID + `]}` + addReq, _ := http.NewRequest("POST", "/accounts", bytes.NewBuffer([]byte(addBody))) + addReq.Header.Set("Content-Type", "application/json") + addReq.Header.Set("Authorization", "Bearer "+token) + router.ServeHTTP(w, addReq) + assert.Equal(t, http.StatusCreated, w.Code) + var addUserResp map[string]interface{} + json.NewDecoder(w.Body).Decode(&addUserResp) + limitedAccountID := fmt.Sprintf("%v", addUserResp["data"].(map[string]interface{})["id"]) + assert.NotEmpty(t, limitedAccountID) + + // Test /auth/ with limited user - should return 200 OK + w = httptest.NewRecorder() + req, _ = http.NewRequest("GET", "/auth/"+unitId_1, nil) + req.SetBasicAuth(limitedUserName, "limited") + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusOK, w.Code, "limited user should have access to unitId_1") + authUser := w.Header().Get("X-Auth-User") + assert.Equal(t, limitedUserName, authUser, "X-Auth-User header should be set to limited user") + + // Test /auth/ with limited user - should return 403 Forbidden + w = httptest.NewRecorder() + req, _ = http.NewRequest("GET", "/auth/"+unitId_2, nil) + req.SetBasicAuth(limitedUserName, "limited") + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusForbidden, w.Code, "limited user should not have access to unitId_2") + + // Login with limited user to get their token + w = httptest.NewRecorder() + limitedLoginBody := []byte(`{"username":"` + limitedUserName + `","password":"limited"}`) + req, _ = http.NewRequest("POST", "/login", bytes.NewBuffer(limitedLoginBody)) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusOK, w.Code) + var limitedLoginResp map[string]interface{} + _ = json.NewDecoder(w.Body).Decode(&limitedLoginResp) + limitedToken, ok := limitedLoginResp["token"].(string) + assert.True(t, ok) + assert.NotEmpty(t, limitedToken) + + // // Test GET /units/ with limited user - should return 200 OK + w = httptest.NewRecorder() + req, _ = http.NewRequest("GET", "/units/"+unitId_1, nil) + req.Header.Set("Authorization", "Bearer "+limitedToken) + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusOK, w.Code, "limited user should be able to get unitId_1") + + // Test GET /units/ with limited user - should return 403 Forbidden + w = httptest.NewRecorder() + req, _ = http.NewRequest("GET", "/units/"+unitId_2, nil) + req.Header.Set("Authorization", "Bearer "+limitedToken) + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusForbidden, w.Code, "limited user should not be able to get unitId_2") + + // Test GET /units with limited user - should only return unitId_1 + w = httptest.NewRecorder() + req, _ = http.NewRequest("GET", "/units", nil) + req.Header.Set("Authorization", "Bearer "+limitedToken) + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusOK, w.Code) + var unitsResp map[string]interface{} + _ = json.NewDecoder(w.Body).Decode(&unitsResp) + unitsData := unitsResp["data"].([]interface{}) + assert.Len(t, unitsData, 1, "limited user should only see one unit") + unit := unitsData[0].(map[string]interface{}) + assert.Equal(t, unitId_1, unit["id"], "limited user should only see unitId_1") + + // Limited user tries to delete unitId_1 (should fail with 403) + w = httptest.NewRecorder() + req, _ = http.NewRequest("DELETE", "/units/"+unitId_1, nil) + req.Header.Set("Authorization", "Bearer "+limitedToken) + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusForbidden, w.Code, "limited user should not be able to delete unitId_1") + + // Limited user tries to add a new unit (should fail with 403) + w = httptest.NewRecorder() + addUnitBody := []byte(`{"unit_id":"shouldfail","username":"failuser","unit_name":"Should Fail Unit","password":"failpass"}`) + req, _ = http.NewRequest("POST", "/units", bytes.NewBuffer(addUnitBody)) + req.Header.Set("Authorization", "Bearer "+limitedToken) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusForbidden, w.Code, "limited user should not be able to add a unit") +} + +func TestToCIDR(t *testing.T) { + tests := []struct { + ip string + mask string + want string + wantErr bool + }{ + {"192.168.1.10", "255.255.255.0", "192.168.1.10/24", false}, + {"172.16.5.4", "255.255.0.0", "172.16.5.4/16", false}, + {"192.168.1.10", "255.255.255.255", "192.168.1.10/32", false}, + {"192.168.1.10", "255.255.0", "", true}, // invalid mask + {"notanip", "255.255.255.0", "", true}, // invalid ip + } + for _, tt := range tests { + got := utils.ToCIDR(tt.ip, tt.mask) + if got == "" { + assert.Error(t, fmt.Errorf("invalid input"), "expected error for input: %v/%v", tt.ip, tt.mask) + } else { + assert.Equal(t, tt.want, got, "unexpected CIDR for input: %v/%v", tt.ip, tt.mask) + } + } +} + +func TestToIpMask(t *testing.T) { + tests := []struct { + cidr string + wantIP string + wantNet string + wantErr bool + }{ + {"192.168.1.10/24", "192.168.1.10", "255.255.255.0", false}, + {"172.16.5.4/16", "172.16.5.4", "255.255.0.0", false}, + {"10.0.0.1/32", "10.0.0.1", "255.255.255.255", false}, + {"192.168.1.10/33", "", "", true}, // invalid mask + {"notanip/24", "", "", true}, // invalid ip + {"", "", "", true}, // empty input + } + for _, tt := range tests { + ip, mask := utils.ToIpMask(tt.cidr) + if tt.wantErr { + assert.Equal(t, "", ip, "expected empty ip for input: %v", tt.cidr) + assert.Equal(t, "", mask, "expected empty mask for input: %v", tt.cidr) + } else { + assert.Equal(t, tt.wantIP, ip, "unexpected ip for input: %v", tt.cidr) + assert.Equal(t, tt.wantNet, mask, "unexpected mask for input: %v", tt.cidr) + } + } +} + +func TestPrometheusEndpoints(t *testing.T) { + router = setupRouter() + + // Add unit for testing + unitId := addUnit(t) + + // Test without authentication - should fail + w := httptest.NewRecorder() + req, _ := http.NewRequest("GET", "/prometheus/targets", nil) + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusUnauthorized, w.Code, "should require authentication") + + // Test with wrong credentials - should fail + w = httptest.NewRecorder() + req, _ = http.NewRequest("GET", "/prometheus/targets", nil) + req.SetBasicAuth("wrong", "credentials") + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusUnauthorized, w.Code, "should reject wrong credentials") + + // Test with correct credentials + w = httptest.NewRecorder() + req, _ = http.NewRequest("GET", "/prometheus/targets", nil) + req.SetBasicAuth(configuration.Config.PrometheusAuthUsername, configuration.Config.PrometheusAuthPassword) + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusOK, w.Code, "should allow access with correct credentials") + + var resp []map[string]interface{} + err := json.NewDecoder(w.Body).Decode(&resp) + assert.NoError(t, err, "response should be valid JSON") + assert.NotEmpty(t, resp, "response should not be empty") + + // Verify that the returned list contains the unit we just added + found := false + for _, item := range resp { + labels, ok := item["labels"].(map[string]interface{}) + if !ok { + continue + } + if unit, ok := labels["unit"].(string); ok && unit == unitId { + found = true + // Optionally check targets field + targets, ok := item["targets"].([]interface{}) + assert.True(t, ok, "targets should be a slice") + assert.NotEmpty(t, targets, "targets should not be empty") + break + } + } + assert.True(t, found, "should find the added unit in prometheus targets list") +} + +// TestPasswordHashingConsistency tests that password hashing is consistent and secure. +func TestPasswordHashingConsistency(t *testing.T) { + password := "TestPassword123!" + + // Hash the same password twice + hash1 := utils.HashPassword(password) + hash2 := utils.HashPassword(password) + + // Hashes should be different (bcrypt uses salt) + assert.NotEqual(t, hash1, hash2, "bcrypt should generate different hashes due to salt") + + // Both hashes should verify the original password + assert.True(t, utils.CheckPasswordHash(password, hash1), "first hash should verify password") + assert.True(t, utils.CheckPasswordHash(password, hash2), "second hash should verify password") + + // Different password should not verify + wrongPassword := "WrongPassword123!" + assert.False(t, utils.CheckPasswordHash(wrongPassword, hash1), "wrong password should not verify") + assert.False(t, utils.CheckPasswordHash(wrongPassword, hash2), "wrong password should not verify") +} + +// TestNetworkUtilitiesListIPs tests the ListIPs utility function for various network sizes. +func TestNetworkUtilitiesListIPs(t *testing.T) { + tests := []struct { + name string + ip string + netmask string + minCount int + wantErr bool + }{ + {"Small network /28", "192.168.1.0", "255.255.255.240", 13, false}, + {"Medium network /24", "10.0.0.0", "255.255.255.0", 250, false}, + {"Large network /16", "172.16.0.0", "255.255.0.0", 65000, false}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + ips, err := utils.ListIPs(tt.ip, tt.netmask) + + if tt.wantErr { + assert.Error(t, err, "expected error for invalid input") + } else { + assert.NoError(t, err, "should not return error for valid input") + assert.GreaterOrEqual(t, len(ips), tt.minCount, "should return expected number of IPs") + + // Verify all IPs are strings + for _, ip := range ips { + assert.NotEmpty(t, ip, "IP should not be empty") + // Basic IP format check + octets := 0 + for i := 0; i < len(ip); i++ { + if ip[i] == '.' { + octets++ + } + } + assert.Equal(t, 3, octets, "IP should have valid format") + } + } + }) + } +} + +// TestListIPsEdgeCases tests edge cases for ListIPs utility. +func TestListIPsEdgeCases(t *testing.T) { + // Test single host network (/32) + ips, err := utils.ListIPs("192.168.1.1", "255.255.255.255") + assert.NoError(t, err, "should handle /32 network") + assert.Equal(t, 1, len(ips), "/32 network should return the single host IP") + + // Test /31 network (point-to-point) + ips, err = utils.ListIPs("192.168.1.0", "255.255.255.254") + assert.NoError(t, err, "should handle /31 network") + assert.Equal(t, 0, len(ips), "/31 network excludes network and broadcast addresses") + + // Test /30 network (smallest typical network) + ips, err = utils.ListIPs("192.168.1.0", "255.255.255.252") + assert.NoError(t, err, "should handle /30 network") + assert.Equal(t, 2, len(ips), "/30 network should return usable IPs (excludes network and broadcast)") +} + +// TestEncryptionKeyRotation tests encryption and decryption with key management. +func TestEncryptionKeyRotation(t *testing.T) { + key1 := []byte("key12345678901234567890123456789") // 32 bytes for AES-256 + key2 := []byte("key22345678901234567890123456789") // 32 bytes for AES-256, different key + plaintext := []byte("sensitive data to encrypt") + + // Encrypt with key1 + ciphertext, err := utils.EncryptAESGCM(plaintext, key1) + assert.NoError(t, err, "encryption should succeed") + assert.NotEmpty(t, ciphertext, "ciphertext should not be empty") + + // Decrypt with key1 should work + decrypted, err := utils.DecryptAESGCM(ciphertext, key1) + assert.NoError(t, err, "decryption with same key should succeed") + assert.Equal(t, plaintext, decrypted, "decrypted text should match original") + + // Decrypt with key2 should fail + _, err = utils.DecryptAESGCM(ciphertext, key2) + assert.Error(t, err, "decryption with different key should fail") +} + +// TestEncryptionRobustness tests encryption with various input sizes. +func TestEncryptionRobustness(t *testing.T) { + key := []byte("12345678901234567890123456789012") // 32 bytes + + tests := []struct { + name string + plaintext []byte + }{ + {"Empty data", []byte("")}, + {"Single byte", []byte("A")}, + {"Small data", []byte("Hello")}, + {"Medium data", []byte("This is a longer message with special chars: !@#$%^&*()")}, + {"Large data", make([]byte, 10000)}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + ciphertext, err := utils.EncryptAESGCM(tt.plaintext, key) + assert.NoError(t, err, "encryption should succeed for %s", tt.name) + + decrypted, err := utils.DecryptAESGCM(ciphertext, key) + assert.NoError(t, err, "decryption should succeed for %s", tt.name) + + // For empty data, both plaintext and decrypted could be empty slice or nil + if len(tt.plaintext) == 0 { + assert.True(t, len(decrypted) == 0, "decrypted data should be empty for %s", tt.name) + } else { + assert.Equal(t, tt.plaintext, decrypted, "decrypted data should match original for %s", tt.name) + } + }) + } +} + +func setupRouter() *gin.Engine { + // Singleton + if router != nil { + return router + } + os.Setenv("LISTEN_ADDRESS", "0.0.0.0:8000,127.0.0.1:5000") + os.Setenv("ADMIN_USERNAME", "admin") + // default password is "password" + os.Setenv("ADMIN_PASSWORD", "admin") + os.Setenv("SECRET_JWT", "secret") + os.Setenv("CREDENTIALS_DIR", "./credentials") + os.Setenv("PROMTAIL_ADDRESS", "127.0.0.1") + os.Setenv("PROMTAIL_PORT", "6565") + os.Setenv("PROMETHEUS_PATH", "/prometheus") + os.Setenv("WEBSSH_PATH", "webssh") + os.Setenv("GRAFANA_PATH", "/grafana") + os.Setenv("REGISTRATION_TOKEN", "1234") + os.Setenv("DATA_DIR", "./data") + os.Setenv("OVPN_DIR", "./ovpn") + os.Setenv("REPORT_DB_URI", "postgres://report:password@127.0.0.1:5432/report") + os.Setenv("GRAFANA_POSTGRES_PASSWORD", "password") + os.Setenv("ISSUER_2FA", "test") + os.Setenv("SECRETS_DIR", "./secrets") + os.Setenv("ENCRYPTION_KEY", "12345678901234567890123456789012") + os.Setenv("PLATFORM_INFO", `{"vpn_port":"1194","vpn_network":"192.168.100.0/24", "controller_version":"1.0.0", "metrics_retention_days":30, "logs_retention_days":90}`) + // disable the global rate limiter in tests: the suite reuses this singleton + // router and fires many requests from one client IP, which would otherwise + // trip the limiter and cause cross-test flakiness (RateLimiter logic itself + // is covered by TestRateLimiter) + os.Setenv("GLOBAL_RATE_LIMIT_AVERAGE", "0") + + // create directory configuration directory + if _, err := os.Stat(os.Getenv("DATA_DIR")); os.IsNotExist(err) { + if err := os.MkdirAll(os.Getenv("DATA_DIR"), 0755); err != nil { + fmt.Printf("failed to create directory: %v\n", err) + os.Exit(1) + } + } + + router := setup() + + return router +} diff --git a/controller/api/methods/account.go b/controller/api/methods/account.go new file mode 100644 index 00000000..8c503e4c --- /dev/null +++ b/controller/api/methods/account.go @@ -0,0 +1,401 @@ +/* + * Copyright (C) 2024 Nethesis S.r.l. + * http://www.nethesis.it - info@nethesis.it + * + * SPDX-License-Identifier: GPL-2.0-only + * + * author: Edoardo Spadoni + */ + +package methods + +import ( + "net/http" + "os" + "os/exec" + "strings" + "time" + + "github.com/NethServer/nethsecurity-controller/api/response" + "github.com/NethServer/nethsecurity-controller/api/configuration" + "github.com/NethServer/nethsecurity-controller/api/models" + "github.com/NethServer/nethsecurity-controller/api/storage" + "github.com/NethServer/nethsecurity-controller/api/utils" + "github.com/fatih/structs" + + jwt "github.com/appleboy/gin-jwt/v2" + "github.com/gin-gonic/gin" +) + +func GetAccounts(c *gin.Context) { + // check auth for not admin users + isAdmin := storage.IsAdmin(jwt.ExtractClaims(c)["id"].(string)) + if !isAdmin { + c.JSON(http.StatusForbidden, structs.Map(response.StatusForbidden{ + Code: 403, + Message: "can't access this resource", + Data: nil, + })) + return + } + + // execute query + accounts, err := storage.GetAccounts() + + if err != nil { + c.JSON(http.StatusInternalServerError, structs.Map(response.StatusInternalServerError{ + Code: 500, + Message: "get accounts error", + Data: err.Error(), + })) + return + } + + // check results + if len(accounts) == 0 { + c.JSON(http.StatusNotFound, structs.Map(response.StatusNotFound{ + Code: 404, + Message: "not found", + Data: nil, + })) + return + } + + // return ok + c.JSON(http.StatusOK, structs.Map(response.StatusOK{ + Code: 200, + Message: "success", + Data: gin.H{"accounts": accounts, "total": len(accounts)}, + })) +} + +func GetAccount(c *gin.Context) { + // check auth for not admin users + isAdmin := storage.IsAdmin(jwt.ExtractClaims(c)["id"].(string)) + if !isAdmin { + c.JSON(http.StatusForbidden, structs.Map(response.StatusForbidden{ + Code: 403, + Message: "can't access this resource", + Data: nil, + })) + return + } + + // get account id + accountID := c.Param("account_id") + + // execute query + accounts, err := storage.GetAccount(accountID) + + if err != nil { + c.JSON(http.StatusInternalServerError, structs.Map(response.StatusInternalServerError{ + Code: 500, + Message: "get account error", + Data: err.Error(), + })) + return + } + + // check results + if len(accounts) == 0 { + c.JSON(http.StatusNotFound, structs.Map(response.StatusNotFound{ + Code: 404, + Message: "not found", + Data: nil, + })) + return + } + + // return ok + c.JSON(http.StatusOK, structs.Map(response.StatusOK{ + Code: 200, + Message: "success", + Data: gin.H{"account": accounts[0]}, + })) +} + +func AddAccount(c *gin.Context) { + // check auth for not admin users + isAdmin := storage.IsAdmin(jwt.ExtractClaims(c)["id"].(string)) + if !isAdmin { + c.JSON(http.StatusForbidden, structs.Map(response.StatusForbidden{ + Code: 403, + Message: "can't access this resource", + Data: nil, + })) + return + } + + // get account fields + var json models.Account + if err := c.BindJSON(&json); err != nil { + c.JSON(http.StatusBadRequest, gin.H{"message": "request fields malformed", "error": err.Error()}) + return + } + + // create account + json.Created = time.Now() + id, err := storage.AddAccount(json) + + // check results + if err != nil { + c.JSON(http.StatusInternalServerError, structs.Map(response.StatusInternalServerError{ + Code: 500, + Message: "add account error", + Data: err.Error(), + })) + return + } + + // return ok + c.JSON(http.StatusCreated, structs.Map(response.StatusCreated{ + Code: 201, + Message: "success", + Data: gin.H{"id": id}, + })) +} + +func UpdateAccount(c *gin.Context) { + // get account id + accountID := c.Param("account_id") + + // check auth for not admin users + isAdmin := storage.IsAdmin(jwt.ExtractClaims(c)["id"].(string)) + if !isAdmin { + c.JSON(http.StatusForbidden, structs.Map(response.StatusForbidden{ + Code: 403, + Message: "can't access this resource", + Data: nil, + })) + return + } + + // get account fields + var json models.AccountUpdate + if err := c.BindJSON(&json); err != nil { + c.JSON(http.StatusBadRequest, gin.H{"message": "request fields malformed", "error": err.Error()}) + return + } + + // update account + err := storage.UpdateAccount(accountID, json) + + // check if all unit_groups exist + for _, groupID := range json.UnitGroups { + exists, err := storage.UnitGroupExists(groupID) + if err != nil { + c.JSON(http.StatusInternalServerError, structs.Map(response.StatusInternalServerError{ + Code: 500, + Message: "error checking group existence", + Data: err.Error(), + })) + return + } + if !exists { + c.JSON(http.StatusBadRequest, structs.Map(response.StatusBadRequest{ + Code: 400, + Message: "unit group does not exist", + Data: gin.H{"unit_group": groupID}, + })) + return + } + } + + // check for groupid_not_found error + if err != nil { + c.JSON(http.StatusInternalServerError, structs.Map(response.StatusInternalServerError{ + Code: 500, + Message: "add account error", + Data: err.Error(), + })) + return + } + + // return ok + c.JSON(http.StatusOK, structs.Map(response.StatusOK{ + Code: 200, + Message: "success", + Data: nil, + })) +} + +func DeleteAccount(c *gin.Context) { + // check auth + isAdmin := storage.IsAdmin(jwt.ExtractClaims(c)["id"].(string)) + if !isAdmin { + c.JSON(http.StatusForbidden, structs.Map(response.StatusForbidden{ + Code: 403, + Message: "can't access this resource", + Data: nil, + })) + return + } + + // get account id + accountID := c.Param("account_id") + + // prevent deleting the admin account + if accountID == "1" { + c.JSON(http.StatusForbidden, structs.Map(response.StatusForbidden{ + Code: 403, + Message: "can't remove super-admin user", + Data: nil, + })) + return + } + + // execute query + err := storage.DeleteAccount(accountID) + + // check results + if err != nil { + c.JSON(http.StatusInternalServerError, structs.Map(response.StatusInternalServerError{ + Code: 500, + Message: "delete account error", + Data: err.Error(), + })) + return + } + + // return ok + c.JSON(http.StatusOK, structs.Map(response.StatusOK{ + Code: 200, + Message: "success", + Data: nil, + })) +} + +func UpdatePassword(c *gin.Context) { + // get passwords fields + var json models.PasswordChange + if err := c.BindJSON(&json); err != nil { + c.JSON(http.StatusBadRequest, gin.H{"message": "request fields malformed", "error": err.Error()}) + return + } + + // get current password + currentPassword := storage.GetPassword(jwt.ExtractClaims(c)["id"].(string)) + + // check if current password is equal with passed one + equal := utils.CheckPasswordHash(json.OldPassword, currentPassword) + + // return err if not equal + if !equal { + c.JSON(http.StatusBadRequest, structs.Map(response.StatusBadRequest{ + Code: 400, + Message: "current password mismatch with passed one", + Data: nil, + })) + return + } + + // update password + err := storage.UpdatePassword(jwt.ExtractClaims(c)["id"].(string), json.NewPassword) + + // check results + if err != nil { + c.JSON(http.StatusInternalServerError, structs.Map(response.StatusInternalServerError{ + Code: 500, + Message: "change password account error", + Data: err.Error(), + })) + return + } + + // return ok + c.JSON(http.StatusOK, structs.Map(response.StatusOK{ + Code: 200, + Message: "success", + Data: nil, + })) + +} + +func GetSSHKeys(c *gin.Context) { + // get username + username := jwt.ExtractClaims(c)["id"].(string) + + // get path for ssh keys + keysPath := configuration.Config.DataDir + "/" + username + ".key" + + // read key + keyPrivate, _ := os.ReadFile(keysPath) + keyPub, _ := os.ReadFile(keysPath + ".pub") + + // return ok + c.JSON(http.StatusOK, structs.Map(response.StatusOK{ + Code: 200, + Message: "success", + Data: gin.H{ + "key_pub": strings.TrimSuffix(string(keyPub), "\n"), + "key": strings.TrimSuffix(string(keyPrivate), "\n"), + }, + })) +} + +func AddSSHKeys(c *gin.Context) { + // get passphrase field + var json models.SSHGenerate + if err := c.BindJSON(&json); err != nil { + c.JSON(http.StatusBadRequest, gin.H{"message": "request fields malformed", "error": err.Error()}) + return + } + + // get username + username := jwt.ExtractClaims(c)["id"].(string) + + // create path for key and key.pub + keysPath := configuration.Config.DataDir + "/" + username + ".key" + + // execute command + args := []string{"-t", "rsa", "-q", "-f", keysPath, "-N", json.Passphrase} + cmd := exec.Command("/usr/bin/ssh-keygen", args...) + + // check error + err := cmd.Run() + if err != nil { + c.JSON(http.StatusInternalServerError, structs.Map(response.StatusInternalServerError{ + Code: 500, + Message: "generate ssh pair failed", + Data: err.Error(), + })) + return + } + + // read key.pub + keyPub, err := os.ReadFile(keysPath + ".pub") + if err != nil { + c.JSON(http.StatusBadRequest, structs.Map(response.StatusBadRequest{ + Code: 400, + Message: "access ssh directory keys file failed", + Data: err.Error(), + })) + return + } + + // return ok + c.JSON(http.StatusOK, structs.Map(response.StatusOK{ + Code: 200, + Message: "success", + Data: gin.H{"key_pub": strings.TrimSuffix(string(keyPub), "\n")}, + })) +} + +func DeleteSSHKeys(c *gin.Context) { + // get username + username := jwt.ExtractClaims(c)["id"].(string) + + // get path for ssh keys + keysPath := configuration.Config.DataDir + "/" + username + ".key" + + // remove both keys + _ = os.Remove(keysPath) + _ = os.Remove(keysPath + ".pub") + + // return ok + c.JSON(http.StatusOK, structs.Map(response.StatusOK{ + Code: 200, + Message: "success", + Data: nil, + })) +} diff --git a/controller/api/methods/auth.go b/controller/api/methods/auth.go new file mode 100644 index 00000000..d4467014 --- /dev/null +++ b/controller/api/methods/auth.go @@ -0,0 +1,423 @@ +/* + * Copyright (C) 2023 Nethesis S.r.l. + * http://www.nethesis.it - info@nethesis.it + * + * SPDX-License-Identifier: GPL-2.0-only + * + * author: Edoardo Spadoni + */ + +package methods + +import ( + "crypto/rand" + "encoding/base32" + "fmt" + "math/big" + "net/http" + "net/url" + "slices" + "sync" + "time" + + "github.com/Jeffail/gabs/v2" + "github.com/NethServer/nethsecurity-controller/api/logs" + "github.com/NethServer/nethsecurity-controller/api/models" + "github.com/NethServer/nethsecurity-controller/api/response" + "github.com/NethServer/nethsecurity-controller/api/storage" + "github.com/NethServer/nethsecurity-controller/api/utils" + jwt "github.com/appleboy/gin-jwt/v2" + "github.com/fatih/structs" + "github.com/gin-gonic/gin" + "github.com/gin-gonic/gin/binding" + jwtl "github.com/golang-jwt/jwt/v5" + + "github.com/NethServer/nethsecurity-controller/api/configuration" + "github.com/pquerna/otp" + "github.com/pquerna/otp/totp" +) + +var activeTokens sync.Map +var tempOtpSecrets sync.Map + +func CheckTokenValidation(username string, token string) bool { + value, ok := activeTokens.Load(username) + if !ok { + return false + } + tokens := value.([]string) + return slices.Contains(tokens, token) +} + +func SetTokenValidation(username string, token string) bool { + value, _ := activeTokens.LoadOrStore(username, []string{}) + tokens := value.([]string) + + // Avoid duplicates + if !slices.Contains(tokens, token) { + tokens = append(tokens, token) + activeTokens.Store(username, tokens) + } + return true +} + +func DelTokenValidation(username string, token string) bool { + value, ok := activeTokens.Load(username) + if !ok { + return false + } + tokens := value.([]string) + + // Remove the token from the user's tokens slice + newTokens := slices.DeleteFunc(tokens, func(s string) bool { + return s == token + }) + if len(newTokens) == 0 { + activeTokens.Delete(username) + } else { + activeTokens.Store(username, newTokens) + } + return true +} + +func OTPVerify(c *gin.Context) { + // get payload + var jsonOTP models.OTPJson + if err := c.ShouldBindBodyWith(&jsonOTP, binding.JSON); err != nil { + c.JSON(http.StatusBadRequest, structs.Map(response.StatusBadRequest{ + Code: 400, + Message: "request fields malformed", + Data: err.Error(), + })) + return + } + + // verify JWT + if !ValidateAuth(jsonOTP.Token, false) { + c.JSON(http.StatusBadRequest, structs.Map(response.StatusBadRequest{ + Code: 400, + Message: "JWT token invalid", + Data: "", + })) + return + } + + // get secret for the user + // - first, search inside the database + // - then, search inside the temporary secrets (for new 2FA setup) + isTempSecret := false + secret := storage.GetUserOtpSecret(jsonOTP.Username) + + // check secret + if len(secret) == 0 { + // search inside the temporary secrets + value, ok := tempOtpSecrets.Load(jsonOTP.Username) + if !ok { + c.JSON(http.StatusNotFound, structs.Map(response.StatusNotFound{ + Code: 404, + Message: "user secret not found", + Data: "", + })) + return + } + secret = value.(string) + isTempSecret = true + } + + // verifiy OTP + valid := false + err := error(nil) + // compose validation error + jsonParsed, _ := gabs.ParseJSON([]byte(`{ + "validation": { + "errors": [ + { + "message": "invalid_otp", + "parameter": "otp", + "value": "" + } + ] + } + }`)) + + valid, err = totp.ValidateCustom(jsonOTP.OTP, secret, time.Now(), totp.ValidateOpts{ + Period: 30, + Skew: 3, // window size + Digits: otp.DigitsSix, + Algorithm: otp.AlgorithmSHA1, + }) + // fail if temp secret and not valid + if isTempSecret { + if err != nil || !valid { + // return validation error + c.JSON(http.StatusBadRequest, structs.Map(response.StatusBadRequest{ + Code: 400, + Message: "validation_failed", + Data: jsonParsed, + })) + return + } else { + // move secret from temp to permanent + ok, _ := SetUserSecret(jsonOTP.Username, secret) + if !ok { + c.JSON(http.StatusBadRequest, structs.Map(response.StatusBadRequest{ + Code: 400, + Message: "user secret set error", + Data: "", + })) + return + } + storage.SetUserRecoveryCodes(jsonOTP.Username, generateRecoveryCodes()) + + // remove from temp + tempOtpSecrets.Delete(jsonOTP.Username) + + // response + c.JSON(http.StatusOK, structs.Map(response.StatusOK{ + Code: 200, + Message: "2FA enabled successfully", + Data: jsonOTP.Token, + })) + return + + } + } + // check for OTP recovery codes only for permanent secrets + if err != nil || !valid { + recoveryCodes := storage.GetRecoveryCodes(jsonOTP.Username) + + if !utils.Contains(jsonOTP.OTP, recoveryCodes) { + // return validation error + c.JSON(http.StatusBadRequest, structs.Map(response.StatusBadRequest{ + Code: 400, + Message: "validation_failed", + Data: jsonParsed, + })) + return + } + + // remove used recovery OTP + recoveryCodes = utils.Remove(jsonOTP.OTP, recoveryCodes) + + // update recovery codes file + if !UpdateRecoveryCodes(jsonOTP.Username, recoveryCodes) { + c.JSON(http.StatusBadRequest, structs.Map(response.StatusBadRequest{ + Code: 400, + Message: "OTP recovery codes not updated", + Data: "", + })) + return + } + + } + + // Just fail if 2FA is not enabled + if !storage.Is2FAEnabled(jsonOTP.Username) { + c.JSON(http.StatusBadRequest, structs.Map(response.StatusBadRequest{ + Code: 400, + Message: "2fa_disabled", + Data: "", + })) + return + } + + // set auth token to valid + if !SetTokenValidation(jsonOTP.Username, jsonOTP.Token) { + c.JSON(http.StatusBadRequest, structs.Map(response.StatusBadRequest{ + Code: 400, + Message: "token validation set error", + Data: "", + })) + return + } + + // response + c.JSON(http.StatusOK, structs.Map(response.StatusOK{ + Code: 200, + Message: "OTP verified", + Data: jsonOTP.Token, + })) +} + +func ValidateAuth(tokenString string, ensureTokenExists bool) bool { + // convert token string and validate it + if tokenString != "" { + token, err := jwtl.Parse(tokenString, func(token *jwtl.Token) (interface{}, error) { + // validate the alg + if _, ok := token.Method.(*jwtl.SigningMethodHMAC); !ok { + return nil, fmt.Errorf("unexpected signing method: %v", token.Header["alg"]) + } + + // return secret + return []byte(configuration.Config.SecretJWT), nil + }) + + if err != nil { + logs.Logs.Println("[ERR][JWT] error in JWT token validation: " + err.Error()) + return false + } + + if claims, ok := token.Claims.(jwtl.MapClaims); ok && token.Valid { + if claims["id"] != nil { + if ensureTokenExists { + username := claims["id"].(string) + + if !CheckTokenValidation(username, tokenString) { + logs.Logs.Println("[ERR][JWT] error JWT token not found") + return false + } + } + return true + } + } else { + logs.Logs.Println("[ERR][JWT] error in JWT token claims") + return false + } + } + return false +} + +func UpdateRecoveryCodes(username string, codes []string) bool { + err := storage.SetUserRecoveryCodes(username, codes) + // check error + return err == nil +} + +func Get2FAStatus(c *gin.Context) { + // get claims from token + claims := jwt.ExtractClaims(c) + var message string + var recoveryCodes []string + + twofa_enabled := storage.Is2FAEnabled(claims["id"].(string)) + if twofa_enabled { + message = "2FA set for this user" + recoveryCodes = storage.GetRecoveryCodes(claims["id"].(string)) + } else { + message = "2FA not set for this user" + recoveryCodes = []string{} + } + + // return response + c.JSON(http.StatusOK, structs.Map(response.StatusOK{ + Code: 200, + Message: message, + Data: gin.H{"status": twofa_enabled, "recovery_codes": recoveryCodes}, + })) +} + +func Del2FAStatus(c *gin.Context) { + // get claims from token + claims := jwt.ExtractClaims(c) + + // revoke 2FA secret + err := storage.SetUserOtpSecret(claims["id"].(string), "") + if err != nil { + c.JSON(http.StatusBadRequest, structs.Map(response.StatusBadRequest{ + Code: 400, + Message: "error in revoke 2FA for user", + Data: nil, + })) + return + } + + // revoke 2FA recovery codes + err = storage.SetUserRecoveryCodes(claims["id"].(string), []string{}) + if err != nil { + c.JSON(http.StatusBadRequest, structs.Map(response.StatusBadRequest{ + Code: 400, + Message: "error in revoke 2FA recovery codes for user", + Data: nil, + })) + return + } + + // response + c.JSON(http.StatusOK, structs.Map(response.StatusOK{ + Code: 200, + Message: "2FA revocate successfully", + Data: "", + })) +} + +func QRCode(c *gin.Context) { + // generate random secret + secret := make([]byte, 20) + _, err := rand.Read(secret) + if err != nil { + logs.Logs.Println("[ERR][2FA] Failed to generate random secret for QRCode: " + err.Error()) + } + + // convert to string + secretBase32 := base32.StdEncoding.EncodeToString(secret) + + // get claims from token + claims := jwt.ExtractClaims(c) + + // define issuer + account := claims["id"].(string) + issuer := configuration.Config.Issuer2FA + + // set temporary secret for user + tempOtpSecrets.Store(account, secretBase32) + + // define URL + URL, err := url.Parse("otpauth://totp") + if err != nil { + logs.Logs.Println("[ERR][2FA] Failed to parse URL for QRCode: " + err.Error()) + } + + // add params + URL.Path += "/" + issuer + ":" + account + params := url.Values{} + params.Add("secret", secretBase32) + params.Add("issuer", issuer) + params.Add("algorithm", "SHA1") + params.Add("digits", "6") + params.Add("period", "30") + + // print url + URL.RawQuery = params.Encode() + + // response + c.JSON(http.StatusOK, structs.Map(response.StatusOK{ + Code: 200, + Message: "QR code string", + Data: gin.H{"url": URL.String(), "key": secretBase32}, + })) +} + +func SetUserSecret(username string, secret string) (bool, string) { + err := storage.SetUserOtpSecret(username, secret) + return err == nil, secret +} + +func generateRecoveryCodes() []string { + recoveryCodes := make([]string, 10) + for i := 0; i < 10; i++ { + num, err := rand.Int(rand.Reader, big.NewInt(1000000)) + if err != nil { + recoveryCodes[i] = "000000" // fallback in case of error + continue + } + recoveryCodes[i] = fmt.Sprintf("%06d", num.Int64()) + } + return recoveryCodes +} + +func UserCanAccessUnit(user string, unitID string) bool { + if storage.IsAdmin(user) { + return true + } + userUnits := storage.GetUserUnits() + units, ok := userUnits[user] + if !ok { + return false + } + for _, u := range units { + if u == unitID { + return true + } + } + return false +} diff --git a/controller/api/methods/defaults.go b/controller/api/methods/defaults.go new file mode 100644 index 00000000..487d904a --- /dev/null +++ b/controller/api/methods/defaults.go @@ -0,0 +1,42 @@ +/* + * Copyright (C) 2024 Nethesis S.r.l. + * http://www.nethesis.it - info@nethesis.it + * + * SPDX-License-Identifier: GPL-2.0-only + * + * author: Edoardo Spadoni + */ + +package methods + +import ( + "net/http" + + "github.com/NethServer/nethsecurity-controller/api/response" + "github.com/NethServer/nethsecurity-controller/api/configuration" + "github.com/fatih/structs" + "github.com/gin-gonic/gin" +) + +func GetDefaults(c *gin.Context) { + // read and return defaults path + c.JSON(http.StatusOK, structs.Map(response.StatusOK{ + Code: 200, + Message: "success", + Data: gin.H{ + "fqdn": configuration.Config.FQDN, + "webssh_path": configuration.Config.WebSSHPath, + "grafana_path": configuration.Config.GrafanaPath, + "valid_subscription": configuration.Config.ValidSubscription, + }, + })) +} + +func GetPlatformInfo(c *gin.Context) { + // read and return platform info + c.JSON(http.StatusOK, structs.Map(response.StatusOK{ + Code: 200, + Message: "success", + Data: structs.Map(configuration.Config.PlatformInfo), + })) +} diff --git a/controller/api/methods/report.go b/controller/api/methods/report.go new file mode 100644 index 00000000..7f48c296 --- /dev/null +++ b/controller/api/methods/report.go @@ -0,0 +1,326 @@ +/* + * Copyright (C) 2024 Nethesis S.r.l. + * http://www.nethesis.it - info@nethesis.it + * + * SPDX-License-Identifier: GPL-2.0-only + * + * author: Giacomo Sanchietti + */ + +package methods + +import ( + "context" + "errors" + "net" + "net/http" + "time" + + "github.com/NethServer/nethsecurity-controller/api/response" + "github.com/NethServer/nethsecurity-controller/api/logs" + "github.com/NethServer/nethsecurity-controller/api/models" + "github.com/NethServer/nethsecurity-controller/api/storage" + "github.com/NethServer/nethsecurity-controller/api/utils" + "github.com/fatih/structs" + "github.com/gin-gonic/gin" + "github.com/jackc/pgx/v5" + "github.com/jackc/pgx/v5/pgxpool" +) + +func setUnitName(dbpool *pgxpool.Pool, dbctx context.Context, c *gin.Context) (int, error) { + // bind json + var req models.UnitNameRequest + if err := c.ShouldBindJSON(&req); err != nil { + return 400, errors.New("invalid request") + } + + // check if uuid is valid + unitId := c.MustGet("UnitId").(string) + var uuid string + err := dbpool.QueryRow(dbctx, "SELECT uuid FROM units WHERE uuid = $1", unitId).Scan(&uuid) + if err != nil || uuid != unitId { + // insert a new unit and return the id + _, err := dbpool.Exec(dbctx, "INSERT INTO units (uuid, name) VALUES ($1, $2)", unitId, req.Name) + if err != nil { + return 500, errors.New("error inserting unit name: " + err.Error()) + } + } else { + // update the unit name + _, err := dbpool.Exec(dbctx, "UPDATE units SET name = $1 WHERE uuid = $2", req.Name, unitId) + if err != nil { + return 500, errors.New("error updating unit name: " + err.Error()) + } + } + return 200, nil +} + +func setUnitOvpnConfig(dbpool *pgxpool.Pool, dbctx context.Context, c *gin.Context) (int, error) { + var req models.UnitOpenVPNRWRequest + if err := c.ShouldBindJSON(&req); err != nil { + return 400, errors.New("invalid request") + } + + // Remove all previous data + _, err := dbpool.Exec(dbctx, "DELETE FROM openvpn_config WHERE uuid = $1", c.MustGet("UnitId").(string)) + if err != nil { + logs.Logs.Println("[ERR][UNITOVPNCONFIG] error deleting previous data: " + err.Error()) + return 500, errors.New("error deleting previous data") + } + + // insert inside OpenVPN table + for _, server := range req.Data { + _, err := dbpool.Exec(dbctx, "INSERT INTO openvpn_config (uuid, instance, name, device, type) VALUES ($1, $2, $3, $4, $5)", c.MustGet("UnitId").(string), server.Instance, server.Name, server.Device, server.Type) + if err != nil { + logs.Logs.Println("[ERR][UNITOVPNCONFIG] error inserting data: " + err.Error()) + return 500, errors.New("error inserting data") + } + } + return 200, nil +} + +func setUnitWan(dbpool *pgxpool.Pool, dbctx context.Context, c *gin.Context) (int, error) { + // bind json + var req models.UnitWanRequest + if err := c.ShouldBindJSON(&req); err != nil { + return 400, errors.New("invalid request") + } + + // Remove all previous data + _, err := dbpool.Exec(dbctx, "DELETE FROM wan_config WHERE uuid = $1", c.MustGet("UnitId").(string)) + if err != nil { + logs.Logs.Println("[ERR][UNITWAN] error deleting previous data: " + err.Error()) + return 500, errors.New("error deleting previous data") + } + // Insert inside WAN table + for _, wan := range req.Data { + _, err := dbpool.Exec(dbctx, "INSERT INTO wan_config (uuid, interface, device, status) VALUES ($1, $2, $3, $4)", c.MustGet("UnitId").(string), wan.Interface, wan.Device, wan.Status) + if err != nil { + logs.Logs.Println("[ERR][UNITWAN] error inserting data: " + err.Error()) + return 500, errors.New("error inserting data") + } + } + return 200, nil +} + +func updateMwanSeries(dbpool *pgxpool.Pool, dbctx context.Context, c *gin.Context) (int, error) { + // bind json + var req models.MwanEventRequest + if err := c.ShouldBindJSON(&req); err != nil { + return 400, errors.New("invalid request") + } + + // To prevent performance issues, do not use single insert + // CopyFrom can't handle conflict resolution, so use batch insert instead + batch := &pgx.Batch{} + for _, event := range req.Data { + // skip invalid objects + if event.Timestamp == 0 || event.Wan == "" || event.Event == "" { + logs.Logs.Println("[WARN][MWANEVENTS] skipping invalid object") + continue + } + batch.Queue("INSERT INTO mwan_events (time, uuid, wan, event, interface) VALUES ($1, $2, $3, $4, $5) ON CONFLICT DO NOTHING", time.Unix(event.Timestamp, 0), c.MustGet("UnitId").(string), event.Wan, event.Event, event.Interface) + } + if batch.Len() != 0 { + err := dbpool.SendBatch(dbctx, batch).Close() + if err != nil { + logs.Logs.Println("[ERR][MWANEVENTS] error inserting data: " + err.Error()) + return 500, errors.New("error inserting data") + } + } + + return 200, nil +} + +func updateTsAttacks(dbpool *pgxpool.Pool, dbctx context.Context, c *gin.Context) (int, error) { + // bind json + var req models.TsAttackRequest + if err := c.ShouldBindJSON(&req); err != nil { + return 400, errors.New("invalid request") + } + + // To prevent performance issues, do not use single insert + // CopyFrom can't handle conflict resolution, so use batch insert instead + batch := &pgx.Batch{} + + for _, attack := range req.Data { + country := "" + // skip invalid objects + if attack.Timestamp == 0 || attack.Ip == "" { + logs.Logs.Println("[WARN][TSATTACKS] skipping invalid object") + continue + } + country = utils.GetCountryShort(attack.Ip) + batch.Queue("INSERT INTO ts_attacks (time, uuid, ip, country) VALUES ($1, $2, $3, $4) ON CONFLICT DO NOTHING", time.Unix(attack.Timestamp, 0), c.MustGet("UnitId").(string), attack.Ip, country) + } + if batch.Len() != 0 { + err := dbpool.SendBatch(dbctx, batch).Close() + if err != nil { + logs.Logs.Println("[ERR][TSATTACKS] error inserting data: " + err.Error()) + return 500, errors.New("error inserting data") + } + } + + return 200, nil +} + +func updateTsMalware(dbpool *pgxpool.Pool, dbctx context.Context, c *gin.Context) (int, error) { + // bind json + var req models.TsMalwareRequest + if err := c.ShouldBindJSON(&req); err != nil { + return 400, errors.New("invalid request") + } + + // To prevent performance issues, do not use single insert + // CopyFrom can't handle conflict resolution, so use batch insert instead + batch := &pgx.Batch{} + for _, malware := range req.Data { + country := "" + // skip invalid objects + if malware.Timestamp == 0 || malware.Src == "" || malware.Dst == "" || malware.Category == "" || malware.Chain == "" { + logs.Logs.Println("[WARN][TSMALWARE] skipping invalid object") + continue + } + + // GeoIP info + if malware.Chain == "inp-wan" { + // Retrieve GeoIP country code for source when traffic is destined to the WAN + country = utils.GetCountryShort(malware.Src) + } else { + // Retrieve GeoIP country code for non-private IP when traffic is forwarded + if !net.ParseIP(malware.Dst).IsPrivate() { + country = utils.GetCountryShort(malware.Dst) + } else if !net.ParseIP(malware.Src).IsPrivate() { + country = utils.GetCountryShort(malware.Src) + } + } + + batch.Queue("INSERT INTO ts_malware (time, uuid, src, dst, category, chain, country) VALUES ($1, $2, $3, $4, $5, $6, $7) ON CONFLICT DO NOTHING", time.Unix(malware.Timestamp, 0), c.MustGet("UnitId").(string), malware.Src, malware.Dst, malware.Category, malware.Chain, country) + } + if batch.Len() != 0 { + err := dbpool.SendBatch(dbctx, batch).Close() + if err != nil { + logs.Logs.Println("[ERR][TSMALWARE] error inserting data: " + err.Error()) + return 500, errors.New("error inserting data") + } + } + + return 200, nil +} + +func updateOvpnConnections(dbpool *pgxpool.Pool, dbctx context.Context, c *gin.Context) (int, error) { + // bind json + var req models.OvpnRwConnectionsRequest + if err := c.ShouldBindJSON(&req); err != nil { + return 400, errors.New("invalid request") + } + + // To prevent performance issues, do not use single insert + // CopyFrom can't handle conflict resolution, so use batch insert instead + batch := &pgx.Batch{} + for _, connection := range req.Data { + country := "" + // skip invalid objects + if connection.Timestamp == 0 || connection.Instance == "" || connection.CommonName == "" || connection.StartTime == 0 { + logs.Logs.Println("[WARN][OVPNCONNECTIONS] skipping invalid object") + continue + } + // GeoIP info for the remote IP + country = utils.GetCountryShort(connection.RemoteIpAddr) + batch.Queue("INSERT INTO ovpnrw_connections (time, uuid, instance, common_name, virtual_ip_addr, remote_ip_addr, start_time, duration, bytes_received, bytes_sent, country) VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11) ON CONFLICT (time, uuid, instance, common_name) DO UPDATE SET duration=EXCLUDED.duration, bytes_received=EXCLUDED.bytes_received, bytes_sent=EXCLUDED.bytes_sent", time.Unix(connection.Timestamp, 0), c.MustGet("UnitId").(string), connection.Instance, connection.CommonName, connection.VirtualIpAddr, connection.RemoteIpAddr, connection.StartTime, connection.Duration, connection.BytesReceived, connection.BytesSent, country) + } + if batch.Len() != 0 { + err := dbpool.SendBatch(dbctx, batch).Close() + if err != nil { + logs.Logs.Println("[ERR][OVPNCONNECTIONS] error inserting data: " + err.Error()) + return 500, errors.New("error inserting data") + } + } + return 200, nil +} + +func updateDpiStats(dbpool *pgxpool.Pool, dbctx context.Context, c *gin.Context) (int, error) { + // bind json + var req models.DpiStatsRequest + if err := c.ShouldBindJSON(&req); err != nil { + return 400, errors.New("invalid request") + } + + // To prevent performance issues, do not use single insert + // CopyFrom can't handle conflict resolution, so use batch insert instead + // do not use CopyFrom + batch := &pgx.Batch{} + for _, dpi := range req.Data { + // skip invalid objects + if dpi.Timestamp == 0 || dpi.ClientAddress == "" || dpi.Bytes == 0 { + logs.Logs.Println("[WARN][DPISTATS] skipping invalid object") + continue + } + batch.Queue("INSERT INTO dpi_stats (time, uuid, client_address, client_name, protocol, host, application, bytes) VALUES ($1, $2, $3, $4, $5, $6, $7, $8) ON CONFLICT (time, uuid, client_address, protocol, host, application) DO UPDATE SET bytes = EXCLUDED.bytes", time.Unix(dpi.Timestamp, 0), c.MustGet("UnitId").(string), dpi.ClientAddress, dpi.ClientName, dpi.Protocol, dpi.Host, dpi.Application, dpi.Bytes) + } + if batch.Len() != 0 { + err := dbpool.SendBatch(dbctx, batch).Close() + if err != nil { + logs.Logs.Println("[ERR][DPISTATS] error inserting data: " + err.Error()) + return 500, errors.New("error inserting data") + } + } + + return 200, nil +} + +func HandelMonitoring(c *gin.Context) { + var err error + var code int + unitId := c.MustGet("UnitId").(string) + + dbpool, dbctx := storage.ReportInstance() + + firewall_api := c.Param("firewall_api") + // setting unit name and creating a unit if it does not exist + if firewall_api == "dump-nsplug-config" { + code, err = setUnitName(dbpool, dbctx, c) + } else { + // for all other metrics, check if the unit exists + var uuid string + err = dbpool.QueryRow(dbctx, "SELECT uuid FROM units WHERE uuid = $1", unitId).Scan(&uuid) + if err != nil { + err = errors.New("unit not found") + code = 404 + } else { + // the unit exists, handle the metric + switch firewall_api { + case "dump-ovpn-config": + code, err = setUnitOvpnConfig(dbpool, dbctx, c) + case "dump-wan-config": + code, err = setUnitWan(dbpool, dbctx, c) + case "dump-ts-malware": + code, err = updateTsMalware(dbpool, dbctx, c) + case "dump-ts-attacks": + code, err = updateTsAttacks(dbpool, dbctx, c) + case "dump-mwan-events": + code, err = updateMwanSeries(dbpool, dbctx, c) + case "dump-dpi-stats": + code, err = updateDpiStats(dbpool, dbctx, c) + case "dump-ovpn-connections": + code, err = updateOvpnConnections(dbpool, dbctx, c) + default: + code = 404 + err = errors.New("metric not found") + } + } + } + + if err != nil { + c.JSON(http.StatusBadRequest, structs.Map(response.StatusBadRequest{ + Code: code, + Message: err.Error(), + Data: nil, + })) + } else { + c.JSON(http.StatusOK, structs.Map(response.StatusOK{ + Code: code, + Message: "success", + Data: nil, + })) + } +} diff --git a/controller/api/methods/unit.go b/controller/api/methods/unit.go new file mode 100644 index 00000000..33001a34 --- /dev/null +++ b/controller/api/methods/unit.go @@ -0,0 +1,1002 @@ +/* + * Copyright (C) 2024 Nethesis S.r.l. + * http://www.nethesis.it - info@nethesis.it + * + * SPDX-License-Identifier: GPL-2.0-only + * + * author: Edoardo Spadoni + */ + +package methods + +import ( + "bytes" + "encoding/json" + "errors" + "net/http" + "os" + "os/exec" + "strconv" + "strings" + "time" + + "github.com/NethServer/nethsecurity-controller/api/configuration" + "github.com/NethServer/nethsecurity-controller/api/logs" + "github.com/NethServer/nethsecurity-controller/api/models" + "github.com/NethServer/nethsecurity-controller/api/response" + + "github.com/NethServer/nethsecurity-controller/api/socket" + "github.com/NethServer/nethsecurity-controller/api/storage" + "github.com/NethServer/nethsecurity-controller/api/utils" + jwt "github.com/appleboy/gin-jwt/v2" + + "github.com/fatih/structs" + "github.com/gin-gonic/gin" +) + +func GetUnits(c *gin.Context) { + // extract user from JWT claims + user := jwt.ExtractClaims(c)["id"].(string) + + units, err := storage.ListUnits() + if err != nil { + c.JSON(http.StatusBadRequest, structs.Map(response.StatusBadRequest{ + Code: 400, + Message: "can't list units", + Data: err.Error(), + })) + return + } + + // loop through units + var results []gin.H + for _, unit := range units { + unitId, ok := unit["id"].(string) + if !ok || !UserCanAccessUnit(user, unitId) { + continue + } + // append to array + results = append(results, unit) + } + + // return 200 OK with data + c.JSON(http.StatusOK, structs.Map(response.StatusOK{ + Code: 200, + Message: "units listed successfully", + Data: results, + })) +} + +func GetUnit(c *gin.Context) { + // get unit id + unitId := c.Param("unit_id") + user := jwt.ExtractClaims(c)["id"].(string) + if !UserCanAccessUnit(user, unitId) { + c.JSON(http.StatusForbidden, structs.Map(response.StatusForbidden{ + Code: 403, + Message: "user does not have access to this unit", + Data: nil, + })) + return + } + + // parse unit file + result, err := storage.GetUnit(unitId) + + if err != nil { + c.JSON(http.StatusBadRequest, structs.Map(response.StatusBadRequest{ + Code: 400, + Message: "can't get unit info for: " + unitId, + Data: err.Error(), + })) + } else { + // return 200 OK with data + c.JSON(http.StatusOK, structs.Map(response.StatusOK{ + Code: 200, + Message: "unit listed successfully", + Data: result, + })) + } +} + +func GetToken(c *gin.Context) { + // extract user + user := jwt.ExtractClaims(c)["id"].(string) + + // get unit id + unitId := c.Param("unit_id") + + // Gating access only if the user can actually access the unit + if !UserCanAccessUnit(user, unitId) { + c.JSON(http.StatusForbidden, structs.Map(response.StatusForbidden{ + Code: 403, + Message: "user does not have access to this unit", + Data: nil, + })) + return + } + + token, expire, err := getUnitToken(unitId, user) + + if err != nil { + c.JSON(http.StatusBadRequest, structs.Map(response.StatusBadRequest{ + Code: 400, + Message: err.Error(), + Data: "", + })) + return + } + + // return 200 OK with data + c.JSON(http.StatusOK, structs.Map(response.StatusOK{ + Code: 200, + Message: "unit token retrieved successfully", + Data: gin.H{ + "token": token, + "expire": expire, + }, + })) +} + +func GetUnitInfo(c *gin.Context) { + // extract user from JWT claims + user := jwt.ExtractClaims(c)["id"].(string) + + // get unit id + unitId := c.Param("unit_id") + + if !UserCanAccessUnit(user, unitId) { + c.JSON(http.StatusForbidden, structs.Map(response.StatusForbidden{ + Code: 403, + Message: "user does not have access to this unit", + Data: nil, + })) + return + } + + // get unit info and store it + info, err := GetRemoteInfo(unitId, user) + + // check errors + if err != nil { + c.JSON(http.StatusBadRequest, structs.Map(response.StatusBadRequest{ + Code: 400, + Message: "can't get unit info for: " + unitId, + Data: err.Error(), + })) + return + } + + // return 200 OK + c.JSON(http.StatusOK, structs.Map(response.StatusOK{ + Code: 200, + Message: "unit info retrieved successfully", + Data: info, + })) + +} + +func AddInfo(c *gin.Context) { + unitId := c.MustGet("UnitId").(string) + var jsonRequest models.UnitInfo + if err := c.ShouldBindJSON(&jsonRequest); err != nil { + c.JSON(http.StatusBadRequest, structs.Map(response.StatusBadRequest{ + Code: 400, + Message: "request fields malformed", + Data: err.Error(), + })) + return + } + + _, err := json.Marshal(jsonRequest) + if err != nil { + c.JSON(http.StatusBadRequest, structs.Map(response.StatusBadRequest{ + Code: 400, + Message: "can't marshal unit info for: " + unitId, + Data: err.Error(), + })) + return + } + storage.SetUnitInfo(unitId, jsonRequest) + + // return 200 OK + c.JSON(http.StatusOK, structs.Map(response.StatusOK{ + Code: 200, + Message: "unit info retrieved successfully", + })) +} + +func AddUnit(c *gin.Context) { + isAdmin := storage.IsAdmin(jwt.ExtractClaims(c)["id"].(string)) + if !isAdmin { + c.JSON(http.StatusForbidden, structs.Map(response.StatusForbidden{ + Code: 403, + Message: "can't access this resource", + Data: nil, + })) + return + } + + // parse request fields + var jsonRequest models.AddRequest + if err := c.ShouldBindJSON(&jsonRequest); err != nil { + c.JSON(http.StatusBadRequest, structs.Map(response.StatusBadRequest{ + Code: 400, + Message: "request fields malformed", + Data: err.Error(), + })) + return + } + + // if the controller does not have a subscription, limit the number of units to 3 + if !configuration.Config.ValidSubscription { + units, err := storage.ListUnits() + if err != nil { + c.JSON(http.StatusBadRequest, structs.Map(response.StatusBadRequest{ + Code: 400, + Message: "can't list units", + Data: err.Error(), + })) + return + } + if len(units) >= 3 { + c.JSON(http.StatusForbidden, structs.Map(response.StatusBadRequest{ + Code: 403, + Message: "subscription limit reached", + Data: "", + })) + return + } + } + + // check duplicates + _, err := storage.GetUnit(jsonRequest.UnitId) + if err == nil { + c.JSON(http.StatusConflict, structs.Map(response.StatusConflict{ + Code: 409, + Message: "duplicated unit id", + Data: "", + })) + return + } + + // get free ip of a network + freeIP := storage.GetFreeIP() + + if freeIP == "" { + c.JSON(http.StatusBadRequest, structs.Map(response.StatusBadRequest{ + Code: 400, + Message: "no IP available for new unit", + Data: nil, + })) + return + } + + // generate certificate request + cmdGenerateGenReq := exec.Command(configuration.Config.EasyRSAPath, "gen-req", jsonRequest.UnitId, "nopass") + cmdGenerateGenReq.Env = append(os.Environ(), + "EASYRSA_BATCH=1", + "EASYRSA_REQ_CN="+jsonRequest.UnitId, + "EASYRSA_PKI="+configuration.Config.OpenVPNPKIDir, + ) + + // Print the executed command for debug + cmdStr := configuration.Config.EasyRSAPath + " gen-req " + jsonRequest.UnitId + " nopass" + logs.Logs.Println("[DEBUG][AddUnit] Executing command: " + cmdStr) + + // Capture stdout and stderr + var stdout, stderr bytes.Buffer + cmdGenerateGenReq.Stdout = &stdout + cmdGenerateGenReq.Stderr = &stderr + + // Print stdout and stderr after execution + if err := cmdGenerateGenReq.Run(); err != nil { + logs.Logs.Println("[ERROR][AddUnit] Command execution failed: "+err.Error(), " Stdout: "+stdout.String(), " Stderr: "+stderr.String()) + c.JSON(http.StatusBadRequest, structs.Map(response.StatusBadRequest{ + Code: 400, + Message: "cannot generate request certificate for: " + jsonRequest.UnitId, + Data: err.Error(), + })) + return + } + + // generate certificate sign + cmdGenerateSignReq := exec.Command(configuration.Config.EasyRSAPath, "sign-req", "client", jsonRequest.UnitId) + cmdGenerateSignReq.Env = append(os.Environ(), + "EASYRSA_BATCH=1", + "EASYRSA_REQ_CN="+jsonRequest.UnitId, + "EASYRSA_PKI="+configuration.Config.OpenVPNPKIDir, + "EASYRSA_CERT_EXPIRE=3650", + ) + if err := cmdGenerateSignReq.Run(); err != nil { + c.JSON(http.StatusBadRequest, structs.Map(response.StatusBadRequest{ + Code: 400, + Message: "cannot sign request certificate for: " + jsonRequest.UnitId, + Data: err.Error(), + })) + return + } + + // create record inside units table + errCreate := storage.AddUnit(jsonRequest.UnitId, freeIP) + if errCreate != nil { + c.JSON(http.StatusBadRequest, structs.Map(response.StatusBadRequest{ + Code: 400, + Message: "cannot store unit record inside database for: " + jsonRequest.UnitId, + Data: errCreate.Error(), + })) + return + } + + // return 200 OK with data + c.JSON(http.StatusOK, structs.Map(response.StatusOK{ + Code: 200, + Message: "unit added successfully", + Data: gin.H{ + "join_code": utils.GetJoinCode(jsonRequest.UnitId), + }, + })) +} + +func RegisterUnit(c *gin.Context) { + token := c.GetHeader("RegistrationToken") + + // check if token exists + if token == "" { + c.JSON(http.StatusUnauthorized, structs.Map(response.StatusBadRequest{ + Code: 403, + Message: "registration token required", + })) + logs.Logs.Println("[ERROR][RegisterUnit] registration token required") + return + } + + // validate token + if token != configuration.Config.RegistrationToken { + c.JSON(http.StatusUnauthorized, structs.Map(response.StatusBadRequest{ + Code: 403, + Message: "invalid registration token", + })) + logs.Logs.Println("[ERROR][RegisterUnit] invalid registration token") + return + } + + // parse request fields + var jsonRequest models.RegisterRequest + if err := c.ShouldBindJSON(&jsonRequest); err != nil { + c.JSON(http.StatusBadRequest, structs.Map(response.StatusBadRequest{ + Code: 400, + Message: "request fields malformed", + Data: err.Error(), + })) + logs.Logs.Println("[ERROR][RegisterUnit] request fields malformed:", err.Error()) + return + } + + // if the controller has a subscription, the unit must have a valid subscription too + if configuration.Config.ValidSubscription && jsonRequest.SubscriptionType == "" { + c.JSON(http.StatusForbidden, structs.Map(response.StatusBadRequest{ + Code: 403, + Message: "unit subscription is required", + Data: "", + })) + logs.Logs.Println("[ERROR][RegisterUnit] unit subscription is required") + return + } + + // if the controller does not have a subscription, the unit must NOT have a valid subscription too + if !configuration.Config.ValidSubscription && jsonRequest.SubscriptionType != "" { + c.JSON(http.StatusForbidden, structs.Map(response.StatusBadRequest{ + Code: 403, + Message: "unit with subscription is not allowed", + Data: "", + })) + logs.Logs.Println("[ERROR][RegisterUnit] unit with subscription is not allowed") + return + } + + // check openvpn conf exists + if _, err := os.Stat(configuration.Config.OpenVPNPKIDir + "/issued/" + jsonRequest.UnitId + ".crt"); err == nil { + // read ca + ca, errCa := os.ReadFile(configuration.Config.OpenVPNPKIDir + "/" + "ca.crt") + caS := strings.TrimSpace(string(ca[:])) + + // check error + if errCa != nil { + c.JSON(http.StatusBadRequest, structs.Map(response.StatusBadRequest{ + Code: 400, + Message: "cannot retrieve openvpn config: ca.crt read failed", + Data: errCa.Error(), + })) + return + } + + // read cert + crt, errCrt := os.ReadFile(configuration.Config.OpenVPNPKIDir + "/issued/" + jsonRequest.UnitId + ".crt") + crtS := strings.TrimSpace(string(crt[:])) + + // check error + if errCrt != nil { + c.JSON(http.StatusBadRequest, structs.Map(response.StatusBadRequest{ + Code: 400, + Message: "cannot retrieve openvpn config: crt read failed", + Data: errCrt.Error(), + })) + return + } + + // read key + key, errKey := os.ReadFile(configuration.Config.OpenVPNPKIDir + "/private/" + jsonRequest.UnitId + ".key") + keyS := strings.TrimSpace(string(key[:])) + + // check error + if errKey != nil { + c.JSON(http.StatusBadRequest, structs.Map(response.StatusBadRequest{ + Code: 400, + Message: "cannot retrieve openvpn config: key read failed", + Data: errKey.Error(), + })) + return + } + + // extract API port from listen address + addressParts := strings.Split(configuration.Config.ListenAddress[0], ":") + apiPort := addressParts[len(addressParts)-1] + // calculate server address from OpenVPNNetwork + openvpnNetwork := strings.TrimSuffix(configuration.Config.OpenVPNNetwork, ".0") + vpnAddress := openvpnNetwork + ".1" + + // compose config + config := gin.H{ + "host": configuration.Config.FQDN, + "port": configuration.Config.OpenVPNUDPPort, + "ca": caS, + "cert": crtS, + "key": keyS, + "promtail_address": configuration.Config.PromtailAddress, + "promtail_port": configuration.Config.PromtailPort, + "api_port": apiPort, + "vpn_address": vpnAddress, + } + + // read credentials from database + curUsername, _, errRead := storage.GetUnitCredentials(jsonRequest.UnitId) + + var errWrite error + // credentials exists, update only if username matches + if errRead == nil { + if curUsername == jsonRequest.Username { + errWrite = storage.SetUnitCredentials(jsonRequest.UnitId, curUsername, jsonRequest.Password) + } + } else { + // create credentials + errWrite = storage.SetUnitCredentials(jsonRequest.UnitId, jsonRequest.Username, jsonRequest.Password) + } + + // save new credentials + if errWrite != nil { + c.JSON(http.StatusBadRequest, structs.Map(response.StatusBadRequest{ + Code: 400, + Message: "cannot write credentials file for: " + jsonRequest.UnitId, + Data: errWrite.Error(), + })) + return + } + + // return 200 OK with data + c.JSON(http.StatusOK, structs.Map(response.StatusOK{ + Code: 200, + Message: "unit registered successfully", + Data: config, + })) + } else { + // return forbidden state + c.JSON(http.StatusForbidden, structs.Map(response.StatusForbidden{ + Code: 403, + Message: "unit not allowed, no certificate found", + Data: "", + })) + logs.Logs.Println("[ERROR][RegisterUnit] unit not allowed, no certificate found for: " + jsonRequest.UnitId) + } +} + +func DeleteUnit(c *gin.Context) { + isAdmin := storage.IsAdmin(jwt.ExtractClaims(c)["id"].(string)) + if !isAdmin { + c.JSON(http.StatusForbidden, structs.Map(response.StatusForbidden{ + Code: 403, + Message: "can't access this resource", + Data: nil, + })) + return + } + + // get unit id + unitId := c.Param("unit_id") + + // kill vpn connection + _ = socket.Write("kill " + unitId) + + // revoke certificate + cmdRevoke := exec.Command(configuration.Config.EasyRSAPath, "revoke", unitId) + cmdRevoke.Env = append(os.Environ(), + "EASYRSA_BATCH=1", + "EASYRSA_PKI="+configuration.Config.OpenVPNPKIDir, + ) + if err := cmdRevoke.Run(); err != nil { + logs.Logs.Println("[ERROR][DeleteUnit] cannot revoke certificate for: " + unitId + " - " + err.Error()) + } + + // renew certificate revocation list + cmdGen := exec.Command(configuration.Config.EasyRSAPath, "gen-crl") + cmdGen.Env = append(os.Environ(), + "EASYRSA_BATCH=1", + "EASYRSA_PKI="+configuration.Config.OpenVPNPKIDir, + "EASYRSA_CRL_DAYS=3650", + ) + if err := cmdGen.Run(); err != nil { + logs.Logs.Println("[ERROR][DeleteUnit] cannot renew certificate revocation list for: " + unitId + " - " + err.Error()) + } + + // delete traefik conf + if _, err := os.Stat(configuration.Config.OpenVPNProxyDir + "/" + unitId + ".yaml"); err == nil { + errDeleteProxy := os.Remove(configuration.Config.OpenVPNProxyDir + "/" + unitId + ".yaml") + if errDeleteProxy != nil { + logs.Logs.Println("[ERROR][DeleteUnit] cannot delete proxy file for: " + unitId + " - " + errDeleteProxy.Error()) + } + } + + deleteError := storage.DeleteUnit(unitId) + if deleteError != nil { + c.JSON(http.StatusInternalServerError, structs.Map(response.StatusInternalServerError{ + Code: 500, + Message: "error in deletion unit record for: " + unitId, + Data: deleteError.Error(), + })) + return + } + + // return 200 OK + c.JSON(http.StatusOK, structs.Map(response.StatusOK{ + Code: 200, + Message: "unit deleted successfully", + Data: "", + })) +} + +func ListConnectedUnits() ([]string, error) { + return storage.ListConnectedUnits() +} + +func getUnitToken(unitId string, onBehalfOf string) (string, string, error) { + + // read credentials + username, password, err := storage.GetUnitCredentials(unitId) + if err != nil { + return "", "", errors.New("cannot read credentials for: " + unitId) + } + + // compose request URL + postURL := configuration.Config.ProxyProtocol + configuration.Config.ProxyHost + ":" + configuration.Config.ProxyPort + "/" + unitId + configuration.Config.LoginEndpoint + + // create request action + credentials := models.LoginRequest{ + Username: username, + Password: password, + OnBehalfOf: onBehalfOf, + } + body, err := json.Marshal(credentials) + if err != nil { + return "", "", errors.New("cannot marshal credentials for: " + unitId) + } + r, err := http.NewRequest("POST", postURL, bytes.NewBuffer(body)) + if err != nil { + return "", "", errors.New("cannot make request for: " + unitId) + } + + // set request header + r.Header.Add("Content-Type", "application/json") + + // make request, 10 seconds timeout + client := &http.Client{Timeout: 10 * time.Second} + res, err := client.Do(r) + if err != nil { + return "", "", errors.New("request failed for: " + unitId) + } + + // close response + defer res.Body.Close() + + // convert response to struct + loginResponse := &models.LoginResponse{} + err = json.NewDecoder(res.Body).Decode(loginResponse) + if err != nil { + return "", "", errors.New("cannot convert response to struct for: " + unitId) + } + + // check if token is not empty + if len(loginResponse.Token) == 0 { + return "", "", errors.New("invalid token response for: " + unitId) + } + + return loginResponse.Token, loginResponse.Expire, nil +} + +// GetRemoteInfo takes an empty onBehalfOf when called by a routine, so that the +// unit logs the request as the controller itself +func GetRemoteInfo(unitId string, onBehalfOf string) (models.UnitInfo, error) { + // get the unit token and execute the request + token, _, _ := getUnitToken(unitId, onBehalfOf) + if token == "" { + return models.UnitInfo{}, errors.New("error getting token") + } + + // compose request URL + postURL := configuration.Config.ProxyProtocol + configuration.Config.ProxyHost + ":" + configuration.Config.ProxyPort + "/" + unitId + "/api/ubus/call" + payload := models.UbusCommand{ + Path: "ns.controller", + Method: "info", + Payload: map[string]interface{}{}, + } + + // convert payload to JSON byte array + payloadBytes, err := json.Marshal(payload) + if err != nil { + return models.UnitInfo{}, errors.New("error marshalling payload") + } + + // create request action + r, err := http.NewRequest("POST", postURL, bytes.NewBuffer(payloadBytes)) + if err != nil { + return models.UnitInfo{}, errors.New("error creating request") + } + + // set request headers + r.Header.Add("Content-Type", "application/json") + r.Header.Add("Authorization", "Bearer "+token) + + // make request, with 10 seconds timeout + client := &http.Client{Timeout: 10 * time.Second} + res, err := client.Do(r) + if err != nil { + return models.UnitInfo{}, errors.New("error making request") + } + defer res.Body.Close() + + // convert response to struct + unitInfo := &models.UbusResponse[models.UnitInfo]{} + err = json.NewDecoder(res.Body).Decode(unitInfo) + if err != nil { + return models.UnitInfo{}, errors.New("error decoding response") + } + + // ask additional info to the unit + systemUpdatePayload := models.UbusCommand{ + Path: "ns.update", + Method: "check-system-update", + Payload: map[string]interface{}{}, + } + + // convert payload to JSON byte array + systemUpdateBytes, _ := json.Marshal(systemUpdatePayload) + systemUpdateRequest, err := http.NewRequest("POST", postURL, bytes.NewBuffer(systemUpdateBytes)) + if err != nil { + return models.UnitInfo{}, errors.New("error creating request") + } + + // set request headers + systemUpdateRequest.Header.Add("Content-Type", "application/json") + systemUpdateRequest.Header.Add("Authorization", "Bearer "+token) + + // make request, with 10 seconds timeout + systemUpdateResponse, err := client.Do(systemUpdateRequest) + if err != nil { + return models.UnitInfo{}, errors.New("error making request") + } + defer systemUpdateResponse.Body.Close() + + // convert response to struct + systemUpdateInfo := &models.UbusResponse[models.CheckSystemUpdate]{} + err = json.NewDecoder(systemUpdateResponse.Body).Decode(systemUpdateInfo) + if err != nil { + return models.UnitInfo{}, errors.New("error decoding response") + } + + unitInfo.Data.ScheduledUpdate = systemUpdateInfo.Data.ScheduledAt + unitInfo.Data.VersionUpdate = systemUpdateInfo.Data.LastVersion + + // write json to database + storage.SetUnitInfo(unitId, unitInfo.Data) + + return unitInfo.Data, nil +} + +func AddUnitGroup(c *gin.Context) { + isAdmin := storage.IsAdmin(jwt.ExtractClaims(c)["id"].(string)) + if !isAdmin { + c.JSON(http.StatusForbidden, structs.Map(response.StatusForbidden{ + Code: 403, + Message: "admin privileges required", + })) + return + } + + var req models.UnitGroup + if err := c.ShouldBindJSON(&req); err != nil { + c.JSON(http.StatusBadRequest, structs.Map(response.StatusBadRequest{ + Code: 400, + Message: "request fields malformed", + Data: err.Error(), + })) + return + } + + id, err := storage.AddUnitGroup(req) + if err != nil { + c.JSON(http.StatusBadRequest, structs.Map(response.StatusBadRequest{ + Code: 400, + Message: "cannot add unit group", + Data: err.Error(), + })) + return + } + + c.JSON(http.StatusCreated, structs.Map(response.StatusCreated{ + Code: 201, + Message: "unit group added successfully", + Data: gin.H{"id": id}, + })) +} + +func UpdateUnitGroup(c *gin.Context) { + isAdmin := storage.IsAdmin(jwt.ExtractClaims(c)["id"].(string)) + if !isAdmin { + c.JSON(http.StatusForbidden, structs.Map(response.StatusForbidden{ + Code: 403, + Message: "admin privileges required", + })) + return + } + + groupId := c.Param("group_id") + if groupId == "" { + c.JSON(http.StatusBadRequest, structs.Map(response.StatusBadRequest{ + Code: 400, + Message: "group_id is required", + })) + return + } + groupIntId, err := strconv.Atoi(groupId) + if err != nil { + c.JSON(http.StatusBadRequest, structs.Map(response.StatusBadRequest{ + Code: 400, + Message: "group_id must be an integer", + Data: err.Error(), + })) + return + } + + var req models.UnitGroup + if err := c.ShouldBindJSON(&req); err != nil { + c.JSON(http.StatusBadRequest, structs.Map(response.StatusBadRequest{ + Code: 400, + Message: "request fields malformed", + Data: err.Error(), + })) + return + } + + for _, unit := range req.Units { + exists, err := storage.UnitExists(unit) + if err != nil { + c.JSON(http.StatusInternalServerError, structs.Map(response.StatusInternalServerError{ + Code: 500, + Message: "error checking unit existence", + Data: err.Error(), + })) + return + } + if !exists { + c.JSON(http.StatusBadRequest, structs.Map(response.StatusBadRequest{ + Code: 400, + Message: "unit does not exist", + Data: unit, + })) + return + } + } + + if err := storage.UpdateUnitGroup(groupIntId, req); err != nil { + c.JSON(http.StatusBadRequest, structs.Map(response.StatusBadRequest{ + Code: 400, + Message: "cannot edit unit group", + Data: err.Error(), + })) + return + } + + c.JSON(http.StatusOK, structs.Map(response.StatusOK{ + Code: 200, + Message: "unit group edited successfully", + })) +} + +func DeleteUnitGroup(c *gin.Context) { + isAdmin := storage.IsAdmin(jwt.ExtractClaims(c)["id"].(string)) + if !isAdmin { + c.JSON(http.StatusForbidden, structs.Map(response.StatusForbidden{ + Code: 403, + Message: "admin privileges required", + })) + return + } + + groupId := c.Param("group_id") + if groupId == "" { + c.JSON(http.StatusBadRequest, structs.Map(response.StatusBadRequest{ + Code: 400, + Message: "group_id is required", + })) + return + } + groupIdInt, err := strconv.Atoi(groupId) + if err != nil { + c.JSON(http.StatusBadRequest, structs.Map(response.StatusBadRequest{ + Code: 400, + Message: "group_id must be an integer", + Data: err.Error(), + })) + return + } + + // check if the unit group is used + used, err := storage.IsUnitGroupUsed(groupIdInt) + if err != nil { + c.JSON(http.StatusInternalServerError, structs.Map(response.StatusInternalServerError{ + Code: 500, + Message: "error checking if unit group is used", + Data: err.Error(), + })) + return + } + if used { + c.JSON(http.StatusBadRequest, structs.Map(response.StatusBadRequest{ + Code: 400, + Message: "unit group is used and cannot be deleted", + })) + return + } + + if err := storage.DeleteUnitGroup(groupIdInt); err != nil { + c.JSON(http.StatusBadRequest, structs.Map(response.StatusBadRequest{ + Code: 400, + Message: "cannot delete unit group", + Data: err.Error(), + })) + return + } + + c.JSON(http.StatusOK, structs.Map(response.StatusOK{ + Code: 200, + Message: "unit group deleted successfully", + })) +} + +func ListUnitGroups(c *gin.Context) { + isAdmin := storage.IsAdmin(jwt.ExtractClaims(c)["id"].(string)) + if !isAdmin { + c.JSON(http.StatusForbidden, structs.Map(response.StatusForbidden{ + Code: 403, + Message: "admin privileges required", + })) + return + } + + groups, err := storage.ListUnitGroups() + if err != nil { + c.JSON(http.StatusBadRequest, structs.Map(response.StatusBadRequest{ + Code: 400, + Message: "cannot list unit groups", + Data: err.Error(), + })) + return + } + + c.JSON(http.StatusOK, structs.Map(response.StatusOK{ + Code: 200, + Message: "unit groups listed successfully", + Data: groups, + })) +} + +func GetUnitGroup(c *gin.Context) { + isAdmin := storage.IsAdmin(jwt.ExtractClaims(c)["id"].(string)) + if !isAdmin { + c.JSON(http.StatusForbidden, structs.Map(response.StatusForbidden{ + Code: 403, + Message: "admin privileges required", + })) + return + } + + groupId := c.Param("group_id") + if groupId == "" { + c.JSON(http.StatusBadRequest, structs.Map(response.StatusBadRequest{ + Code: 400, + Message: "group_id is required", + })) + return + } + groupIdInt, err := strconv.Atoi(groupId) + if err != nil { + c.JSON(http.StatusBadRequest, structs.Map(response.StatusBadRequest{ + Code: 400, + Message: "group_id must be an integer", + Data: err.Error(), + })) + return + } + group, err := storage.GetUnitGroup(groupIdInt) + if err != nil { + c.JSON(http.StatusBadRequest, structs.Map(response.StatusBadRequest{ + Code: 400, + Message: "cannot get unit group", + Data: err.Error(), + })) + return + } + + c.JSON(http.StatusOK, structs.Map(response.StatusOK{ + Code: 200, + Message: "unit group retrieved successfully", + Data: group, + })) +} + +func GetPrometheusTargets(c *gin.Context) { + // Get all units + units, err := storage.ListUnits() + if err != nil { + c.JSON(http.StatusBadRequest, structs.Map(response.StatusBadRequest{ + Code: 400, + Message: "can't list units", + Data: err.Error(), + })) + return + } + + // Create Prometheus target format + // Format reference: https://prometheus.io/docs/prometheus/latest/http_sd/ + var targets []gin.H + for _, unit := range units { + unitId, ok1 := unit["id"].(string) + unitIp, ok2 := unit["ipaddress"].(string) + + if ok1 && ok2 { + // Netdata target (port 19999) + netdataTarget := gin.H{ + "targets": []string{unitIp + ":19999"}, + "labels": gin.H{ + "node": unitIp, + "unit": unitId, + "metrics_type": "netdata", + "__metrics_path__": "/api/v1/allmetrics?format=prometheus&help=no"}, + } + targets = append(targets, netdataTarget) + + // Telegraf target (port 9273) + telegrafTarget := gin.H{ + "targets": []string{unitIp + ":9273"}, + "labels": gin.H{ + "node": unitIp, + "unit": unitId, + "metrics_type": "telegraf", + "__metrics_path__": "/metrics"}, + } + targets = append(targets, telegrafTarget) + } + } + + // Return targets in Prometheus HTTP SD format + c.JSON(http.StatusOK, targets) +} diff --git a/controller/api/middleware/middleware.go b/controller/api/middleware/middleware.go new file mode 100644 index 00000000..832d292f --- /dev/null +++ b/controller/api/middleware/middleware.go @@ -0,0 +1,467 @@ +/* + * Copyright (C) 2024 Nethesis S.r.l. + * http://www.nethesis.it - info@nethesis.it + * + * SPDX-License-Identifier: GPL-2.0-only + * + * author: Edoardo Spadoni + */ + +package middleware + +import ( + "bytes" + "encoding/json" + "io" + "net/http" + "os" + "strings" + "sync" + "time" + + "github.com/fatih/structs" + "github.com/gin-gonic/gin" + "github.com/nqd/flat" + "golang.org/x/time/rate" + + jwt "github.com/appleboy/gin-jwt/v2" + + "github.com/NethServer/nethsecurity-controller/api/models" + "github.com/NethServer/nethsecurity-controller/api/response" + "github.com/NethServer/nethsecurity-controller/api/configuration" + "github.com/NethServer/nethsecurity-controller/api/logs" + "github.com/NethServer/nethsecurity-controller/api/methods" + "github.com/NethServer/nethsecurity-controller/api/storage" + "github.com/NethServer/nethsecurity-controller/api/utils" +) + +type login struct { + Username string `form:"username" json:"username" binding:"required"` + Password string `form:"password" json:"password" binding:"required"` +} + +const cookieName = "ns_jwt" + +var jwtMiddleware *jwt.GinJWTMiddleware +var identityKey = "id" + +func InstanceJWT() *jwt.GinJWTMiddleware { + if jwtMiddleware == nil { + jwtMiddleware := InitJWT() + return jwtMiddleware + } + return jwtMiddleware +} + +func InitJWT() *jwt.GinJWTMiddleware { + // define jwt middleware + authMiddleware, errDefine := jwt.New(&jwt.GinJWTMiddleware{ + Realm: "nethserver", + Key: []byte(configuration.Config.SecretJWT), + Timeout: time.Hour * 24, // 1 day + MaxRefresh: time.Hour * 24, // 1 day + IdentityKey: identityKey, + Authenticator: func(c *gin.Context) (interface{}, error) { + // check login credentials exists + var loginVals login + if err := c.ShouldBind(&loginVals); err != nil { + return "", jwt.ErrMissingLoginValues + } + + // set login credentials + username := loginVals.Username + password := loginVals.Password + + // read user password hash + passwordHash := storage.GetPassword(username) + + // check password + valid := utils.CheckPasswordHash(password, passwordHash) + + if !valid { + // login fail action + logs.Logs.Println("[INFO][AUTH] authentication failed for user " + username) + + // return JWT error + return nil, jwt.ErrFailedAuthentication + } + + // login ok action + logs.Logs.Println("[INFO][AUTH] authentication success for user " + username) + + // return user auth model + return &models.UserAuthorizations{ + Username: username, + }, nil + + }, + PayloadFunc: func(data interface{}) jwt.MapClaims { + // read current user + if user, ok := data.(*models.UserAuthorizations); ok { + // define role + role := "user" + + // check if username is admin + isAdmin := storage.IsAdmin(user.Username) + if isAdmin { + role = "admin" + } + + // check if user require 2fa + status := storage.Is2FAEnabled(user.Username) + + // create claims map + return jwt.MapClaims{ + identityKey: user.Username, + "role": role, + "actions": []string{}, + "2fa": status, + } + } + + // return claims map + return jwt.MapClaims{} + }, + IdentityHandler: func(c *gin.Context) interface{} { + // handle identity and extract claims + claims := jwt.ExtractClaims(c) + + // create user object + user := &models.UserAuthorizations{ + Username: claims[identityKey].(string), + Role: "admin", + Actions: nil, + } + + // return user + return user + }, + Authorizator: func(data interface{}, c *gin.Context) bool { + // check token validation + claims, _ := InstanceJWT().GetClaimsFromJWT(c) + token, _ := InstanceJWT().ParseToken(c) + + // log request and body + reqMethod := c.Request.Method + reqURI := c.Request.RequestURI + + // check if token exists + if !methods.CheckTokenValidation(claims["id"].(string), token.Raw) { + // write logs + logs.Logs.Println("[INFO][AUTH] authorization failed for user " + claims["id"].(string) + ". " + reqMethod + " " + reqURI) + + // not authorized + return false + } + + // extract body + reqBody := "" + if reqMethod == "POST" || reqMethod == "PUT" { + // extract body + var buf bytes.Buffer + tee := io.TeeReader(c.Request.Body, &buf) + body, _ := io.ReadAll(tee) + c.Request.Body = io.NopCloser(&buf) + + // convert to map and flat it + var jsonDyn map[string]interface{} + json.Unmarshal(body, &jsonDyn) + in, _ := flat.Flatten(jsonDyn, nil) + + // search for sensitve data, in sensitive list + for k := range in { + for _, s := range configuration.Config.SensitiveList { + if strings.Contains(strings.ToLower(k), strings.ToLower(s)) { + in[k] = "XXX" + } + } + } + + // unflat the map + out, _ := flat.Unflatten(in, nil) + + // convert to json string + jsonOut, _ := json.Marshal(out) + + // compose string + reqBody = string(jsonOut) + } + + logs.Logs.Println("[INFO][AUTH] authorization success for user " + claims["id"].(string) + ". " + reqMethod + " " + reqURI + " " + reqBody) + + // authorized + return true + }, + LoginResponse: func(c *gin.Context, code int, token string, t time.Time) { + //get claims + tokenObj, _ := InstanceJWT().ParseTokenString(token) + claims := jwt.ExtractClaimsFromToken(tokenObj) + + // set token to valid, if not 2FA + if !claims["2fa"].(bool) { + methods.SetTokenValidation(claims["id"].(string), token) + } + + // write logs + logs.Logs.Println("[INFO][AUTH] login response success for user " + claims["id"].(string)) + + // return 200 OK + c.JSON(200, gin.H{"code": 200, "expire": t, "token": token}) + }, + RefreshResponse: func(c *gin.Context, code int, token string, t time.Time) { + //get claims + tokenObj, _ := InstanceJWT().ParseTokenString(token) + claims := jwt.ExtractClaimsFromToken(tokenObj) + + // set token to valid + methods.SetTokenValidation(claims["id"].(string), token) + + // write logs + logs.Logs.Println("[INFO][AUTH] refresh response success for user " + claims["id"].(string)) + + // return 200 OK + c.JSON(200, gin.H{"code": 200, "expire": t, "token": token}) + }, + LogoutResponse: func(c *gin.Context, code int) { + //get claims + tokenObj, _ := InstanceJWT().ParseToken(c) + claims := jwt.ExtractClaimsFromToken(tokenObj) + + // set token to invalid + methods.DelTokenValidation(claims["id"].(string), tokenObj.Raw) + + // write logs + logs.Logs.Println("[INFO][AUTH] logout response success for user " + claims["id"].(string)) + + // reutrn 200 OK + c.JSON(200, gin.H{"code": 200}) + }, + Unauthorized: func(c *gin.Context, code int, message string) { + // write logs + logs.Logs.Println("[INFO][AUTH] unauthorized request: " + message) + + // response not authorized + c.JSON(code, structs.Map(response.StatusUnauthorized{ + Code: code, + Message: message, + Data: nil, + })) + }, + SendCookie: true, + CookieName: cookieName, + SecureCookie: gin.Mode() != gin.DebugMode, + CookieHTTPOnly: true, + CookieSameSite: http.SameSiteLaxMode, + TokenLookup: "header: Authorization, token: jwt", + TokenHeadName: "Bearer", + TimeFunc: time.Now, + }) + + // check middleware errors + if errDefine != nil { + logs.Logs.Println("[ERR][AUTH] middleware definition error: " + errDefine.Error()) + } + + // init middleware + errInit := authMiddleware.MiddlewareInit() + + // check error on initialization + if errInit != nil { + logs.Logs.Println("[ERR][AUTH] middleware initialization error: " + errInit.Error()) + } + + // return object + return authMiddleware +} + +func BasicUnitAuth() gin.HandlerFunc { + return func(c *gin.Context) { + uuid, token, _ := c.Request.BasicAuth() + if uuid == "" || token == "" { + c.JSON(http.StatusBadRequest, structs.Map(response.StatusUnauthorized{ + Code: 400, + Message: "missing unit or token", + Data: nil, + })) + c.Abort() + return + } + + // validate registration token against configured one + if token != configuration.Config.RegistrationToken { + c.JSON(http.StatusUnauthorized, structs.Map(response.StatusBadRequest{ + Code: 401, + Message: "invalid registration token", + })) + c.Abort() + return + } + + // UnitId is invalid if there is no certificate issued for it + if _, err := os.Stat(configuration.Config.OpenVPNPKIDir + "/issued/" + uuid + ".crt"); err != nil { + c.JSON(http.StatusUnauthorized, structs.Map(response.StatusUnauthorized{ + Code: 401, + Message: "invalid unit id", + Data: nil, + })) + c.Abort() + return + } + + c.Set("UnitId", uuid) + c.Next() + } +} + +func BasicUserAuth() gin.HandlerFunc { + return func(c *gin.Context) { + // Try JWT authentication from cookie (used by Traefik ForwardAuth) + cookiePresent := false + tokenStr, err := c.Cookie(cookieName) + if err == nil && tokenStr != "" { + cookiePresent = true + token, err := InstanceJWT().ParseTokenString(tokenStr) + if err == nil && token.Valid { + claims := jwt.ExtractClaimsFromToken(token) + if id, ok := claims[identityKey].(string); ok && methods.CheckTokenValidation(id, token.Raw) { + // Enforce per-unit access check + unitID := c.Param("unit_id") + if unitID != "" && !methods.UserCanAccessUnit(id, unitID) { + c.JSON(http.StatusForbidden, structs.Map(response.StatusForbidden{ + Code: 403, + Message: "user does not have access to this unit", + Data: nil, + })) + logs.Logs.Println("[INFO][AUTH] user " + id + " does not have access to unit " + unitID) + c.Abort() + return + } + logs.Logs.Println("[INFO][AUTH] user " + id + " authenticated via JWT cookie") + c.Header("X-Auth-User", id) + c.Next() + return + } + } + // Cookie present but invalid/expired: clear it + c.SetCookie(cookieName, "", -1, "/", "", gin.Mode() != gin.DebugMode, true) + } + + // Fall back to Basic Auth + username, password, _ := c.Request.BasicAuth() + + if username == "" || password == "" { + if cookiePresent { + logs.Logs.Println("[INFO][AUTH] invalid or expired JWT cookie, cleared") + } + c.JSON(http.StatusUnauthorized, structs.Map(response.StatusUnauthorized{ + Code: 401, + Message: "missing or invalid credentials", + Data: nil, + })) + c.Abort() + return + } + + // read user password hash + passwordHash := storage.GetPassword(username) + + // check password and username + valid := utils.CheckPasswordHash(password, passwordHash) + + if !valid { + c.JSON(http.StatusUnauthorized, structs.Map(response.StatusUnauthorized{ + Code: 401, + Message: "invalid username or password", + Data: nil, + })) + logs.Logs.Println("[INFO][AUTH] user " + username + " authentication failed") + c.Abort() + return + } + + // Optionally load unit_id from query or header + unitID := c.Param("unit_id") + extra_log := "" + if unitID != "" { + if !methods.UserCanAccessUnit(username, unitID) { + c.JSON(http.StatusForbidden, structs.Map(response.StatusForbidden{ + Code: 403, + Message: "user does not have access to this unit", + Data: nil, + })) + logs.Logs.Println("[INFO][AUTH] user " + username + " does not have access to unit " + unitID) + c.Abort() + return + } + extra_log = " to unit " + unitID + } + // Just return success + logs.Logs.Println("[INFO][AUTH] user "+username+" authenticated successfully", extra_log) + c.Header("X-Auth-User", username) + c.Next() + } +} + +// BodyLimit rejects requests whose body exceeds maxBytes before any downstream +// binding reads it, so unauthenticated or semi-trusted routes cannot be used +// to exhaust memory with oversized payloads. +func BodyLimit(maxBytes int64) gin.HandlerFunc { + return func(c *gin.Context) { + c.Request.Body = http.MaxBytesReader(c.Writer, c.Request.Body, maxBytes) + c.Next() + } +} + +const rateLimiterStaleAfter = 3 * time.Minute +const rateLimiterCleanupInterval = time.Minute + +type rateLimiterVisitor struct { + limiter *rate.Limiter + lastSeen time.Time +} + +// RateLimiter throttles requests per client IP with a token-bucket limiter, +// so an unauthenticated route cannot be flooded with enough concurrent +// requests to exhaust memory before BodyLimit's per-request cap can help +// (BodyLimit bounds one request's body, not how many requests run at once). +// rps is the sustained rate and burst the number of requests allowed instantly. +func RateLimiter(rps rate.Limit, burst int) gin.HandlerFunc { + visitors := make(map[string]*rateLimiterVisitor) + var mu sync.Mutex + + go func() { + for { + time.Sleep(rateLimiterCleanupInterval) + mu.Lock() + for ip, v := range visitors { + if time.Since(v.lastSeen) > rateLimiterStaleAfter { + delete(visitors, ip) + } + } + mu.Unlock() + } + }() + + return func(c *gin.Context) { + ip := c.ClientIP() + + mu.Lock() + v, exists := visitors[ip] + if !exists { + v = &rateLimiterVisitor{limiter: rate.NewLimiter(rps, burst)} + visitors[ip] = v + } + v.lastSeen = time.Now() + limiter := v.limiter + mu.Unlock() + + if !limiter.Allow() { + logs.Logs.Println("[INFO][AUTH] rate limit exceeded for " + ip + " on " + c.Request.URL.Path) + c.JSON(http.StatusTooManyRequests, gin.H{"code": http.StatusTooManyRequests, "message": "too many requests", "data": nil}) + c.Abort() + return + } + + c.Next() + } +} diff --git a/controller/api/middleware/middleware_test.go b/controller/api/middleware/middleware_test.go new file mode 100644 index 00000000..66ea2f82 --- /dev/null +++ b/controller/api/middleware/middleware_test.go @@ -0,0 +1,247 @@ +/* + * Copyright (C) 2025 Nethesis S.r.l. + * http://www.nethesis.it - info@nethesis.it + * + * SPDX-License-Identifier: GPL-2.0-only + * + * author: Giacomo Sanchietti + */ + +package middleware + +import ( + "io" + "net/http" + "net/http/httptest" + "os" + "strings" + "testing" + "time" + + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/assert" + "golang.org/x/time/rate" + + "github.com/NethServer/nethsecurity-controller/api/models" + "github.com/NethServer/nethsecurity-controller/api/configuration" + "github.com/NethServer/nethsecurity-controller/api/logs" + "github.com/NethServer/nethsecurity-controller/api/methods" +) + +func TestMain(m *testing.M) { + gin.SetMode(gin.TestMode) + logs.Init("test") + os.Exit(m.Run()) +} + +func TestInitJWT(t *testing.T) { + // Set required config + configuration.Config.SecretJWT = "test_secret" + + mw := InitJWT() + assert.NotNil(t, mw) +} + +func TestInstanceJWT(t *testing.T) { + // Set required config + configuration.Config.SecretJWT = "test_secret" + + mw := InstanceJWT() + assert.NotNil(t, mw) +} + +func TestBasicUnitAuth(t *testing.T) { + // Set required config + configuration.Config.RegistrationToken = "test_token" + configuration.Config.OpenVPNPKIDir = "/tmp" // Mock path + + r := gin.New() + r.Use(BasicUnitAuth()) + r.GET("/test", func(c *gin.Context) { + c.JSON(200, gin.H{"message": "ok"}) + }) + + // Test missing auth + req, _ := http.NewRequest("GET", "/test", nil) + w := httptest.NewRecorder() + r.ServeHTTP(w, req) + assert.Equal(t, 400, w.Code) + + // Test invalid token + req.SetBasicAuth("unit1", "wrong_token") + w = httptest.NewRecorder() + r.ServeHTTP(w, req) + assert.Equal(t, 401, w.Code) + + // Test valid token but invalid unit (no cert file) + req.SetBasicAuth("unit1", "test_token") + w = httptest.NewRecorder() + r.ServeHTTP(w, req) + assert.Equal(t, 401, w.Code) +} + +func TestJWTLogin(t *testing.T) { + // Set required config + configuration.Config.SecretJWT = "test_secret" + + r := gin.New() + r.POST("/login", InstanceJWT().LoginHandler) + + // Test login with invalid credentials (will fail due to storage) + req, _ := http.NewRequest("POST", "/login", nil) + req.Header.Set("Content-Type", "application/json") + req.Body = io.NopCloser(strings.NewReader(`{"username":"admin","password":"pass"}`)) + w := httptest.NewRecorder() + r.ServeHTTP(w, req) + // Should return 401 since storage fails + assert.Equal(t, 401, w.Code) +} + +func generateTestToken(t *testing.T, username string) string { + t.Helper() + configuration.Config.SecretJWT = "test_secret" + mw := InstanceJWT() + + token, _, err := mw.TokenGenerator(&models.UserAuthorizations{Username: username}) + assert.NoError(t, err) + assert.NotEmpty(t, token) + + // Register the token as active + methods.SetTokenValidation(username, token) + + return token +} + +func TestBasicUserAuthCookie(t *testing.T) { + token := generateTestToken(t, "testuser") + + r := gin.New() + r.Use(BasicUserAuth()) + r.GET("/auth", func(c *gin.Context) { + c.JSON(200, gin.H{"message": "ok"}) + }) + + // Valid cookie returns 200 and sets X-Auth-User + req, _ := http.NewRequest("GET", "/auth", nil) + req.AddCookie(&http.Cookie{Name: cookieName, Value: token}) + w := httptest.NewRecorder() + r.ServeHTTP(w, req) + assert.Equal(t, 200, w.Code) + assert.Equal(t, "testuser", w.Header().Get("X-Auth-User")) + + // Invalid cookie → 401 and cookie is cleared (Set-Cookie with Max-Age=-1) + req, _ = http.NewRequest("GET", "/auth", nil) + req.AddCookie(&http.Cookie{Name: cookieName, Value: "invalid.jwt.token"}) + w = httptest.NewRecorder() + r.ServeHTTP(w, req) + assert.Equal(t, 401, w.Code) + assert.Contains(t, w.Header().Get("Set-Cookie"), cookieName+"=;") + + // No cookie and no Basic Auth → 401 + req, _ = http.NewRequest("GET", "/auth", nil) + w = httptest.NewRecorder() + r.ServeHTTP(w, req) + assert.Equal(t, 401, w.Code) + + // Clean up + methods.DelTokenValidation("testuser", token) +} + +func TestBasicUserAuthCookieUnitAccess(t *testing.T) { + token := generateTestToken(t, "limiteduser") + + r := gin.New() + r.Use(BasicUserAuth()) + r.GET("/auth/:unit_id", func(c *gin.Context) { + c.JSON(200, gin.H{"message": "ok"}) + }) + + // Non-admin user with cookie accessing a unit → 403 + // (limiteduser is not in adminUsers and has no unit assignments) + req, _ := http.NewRequest("GET", "/auth/unit-123", nil) + req.AddCookie(&http.Cookie{Name: cookieName, Value: token}) + w := httptest.NewRecorder() + r.ServeHTTP(w, req) + assert.Equal(t, 403, w.Code) + assert.Contains(t, w.Body.String(), "user does not have access to this unit") + + // Same user without unit_id param → 200 (no unit check needed) + r2 := gin.New() + r2.Use(BasicUserAuth()) + r2.GET("/auth", func(c *gin.Context) { + c.JSON(200, gin.H{"message": "ok"}) + }) + req, _ = http.NewRequest("GET", "/auth", nil) + req.AddCookie(&http.Cookie{Name: cookieName, Value: token}) + w = httptest.NewRecorder() + r2.ServeHTTP(w, req) + assert.Equal(t, 200, w.Code) + + // Clean up + methods.DelTokenValidation("limiteduser", token) +} + +func TestBodyLimit(t *testing.T) { + r := gin.New() + r.Use(BodyLimit(8)) + r.POST("/test", func(c *gin.Context) { + _, err := io.ReadAll(c.Request.Body) + if err != nil { + c.JSON(http.StatusRequestEntityTooLarge, gin.H{"error": err.Error()}) + return + } + c.JSON(200, gin.H{"message": "ok"}) + }) + + // Body within limit is accepted + req, _ := http.NewRequest("POST", "/test", strings.NewReader("1234567")) + w := httptest.NewRecorder() + r.ServeHTTP(w, req) + assert.Equal(t, 200, w.Code) + + // Body exceeding limit is rejected by the handler's read, not silently truncated + req, _ = http.NewRequest("POST", "/test", strings.NewReader("123456789")) + w = httptest.NewRecorder() + r.ServeHTTP(w, req) + assert.Equal(t, http.StatusRequestEntityTooLarge, w.Code) + + // Enforced on bytes actually read, independent of a spoofed Content-Length + req, _ = http.NewRequest("POST", "/test", strings.NewReader("123456789")) + req.ContentLength = 5 + w = httptest.NewRecorder() + r.ServeHTTP(w, req) + assert.Equal(t, http.StatusRequestEntityTooLarge, w.Code) +} + +func TestRateLimiter(t *testing.T) { + r := gin.New() + r.Use(RateLimiter(rate.Every(time.Minute), 2)) + r.GET("/test", func(c *gin.Context) { + c.JSON(200, gin.H{"message": "ok"}) + }) + + newReq := func(remoteAddr string) *http.Request { + req, _ := http.NewRequest("GET", "/test", nil) + req.RemoteAddr = remoteAddr + return req + } + + // Burst of 2 is allowed for the same client IP + w := httptest.NewRecorder() + r.ServeHTTP(w, newReq("1.2.3.4:1111")) + assert.Equal(t, 200, w.Code) + + w = httptest.NewRecorder() + r.ServeHTTP(w, newReq("1.2.3.4:2222")) + assert.Equal(t, 200, w.Code) + + // Third request from the same IP within the window is rejected + w = httptest.NewRecorder() + r.ServeHTTP(w, newReq("1.2.3.4:3333")) + assert.Equal(t, http.StatusTooManyRequests, w.Code) + + // A different client IP has its own independent bucket + w = httptest.NewRecorder() + r.ServeHTTP(w, newReq("5.6.7.8:1111")) + assert.Equal(t, 200, w.Code) +} diff --git a/controller/api/middleware_test.go b/controller/api/middleware_test.go new file mode 100644 index 00000000..6e01416a --- /dev/null +++ b/controller/api/middleware_test.go @@ -0,0 +1,318 @@ +/* + * Copyright (C) 2025 Nethesis S.r.l. + * http://www.nethesis.it - info@nethesis.it + * + * SPDX-License-Identifier: GPL-2.0-only + * + * author: Giacomo Sanchietti + */ + +package main + +import ( + "bytes" + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/assert" +) + +// TestBasicAuthUnit tests that BasicAuth works for unit authentication. +func TestBasicAuthUnit(t *testing.T) { + gin.SetMode(gin.TestMode) + router := setupRouter() + + // Test /prometheus/targets with valid basic auth + w := httptest.NewRecorder() + req, _ := http.NewRequest("GET", "/prometheus/targets", nil) + req.SetBasicAuth("prometheus", "prometheus") + router.ServeHTTP(w, req) + + // Should succeed with correct credentials + assert.Equal(t, http.StatusOK, w.Code, "BasicAuth should succeed with correct credentials") +} + +// TestBasicAuthInvalid tests that BasicAuth rejects invalid credentials. +func TestBasicAuthInvalid(t *testing.T) { + gin.SetMode(gin.TestMode) + router := setupRouter() + + // Test /prometheus/targets with invalid basic auth + w := httptest.NewRecorder() + req, _ := http.NewRequest("GET", "/prometheus/targets", nil) + req.SetBasicAuth("wrong", "credentials") + router.ServeHTTP(w, req) + + // Should fail with wrong credentials + assert.Equal(t, http.StatusUnauthorized, w.Code, "BasicAuth should fail with wrong credentials") +} + +// TestJWTExpiration tests that expired JWT tokens are rejected. +func TestJWTExpiration(t *testing.T) { + gin.SetMode(gin.TestMode) + router := setupRouter() + + // First, login to get a token + w := httptest.NewRecorder() + loginBody := `{"username": "admin", "password": "admin"}` + req, _ := http.NewRequest("POST", "/login", bytes.NewBuffer([]byte(loginBody))) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(w, req) + + var loginResp map[string]interface{} + json.NewDecoder(w.Body).Decode(&loginResp) + assert.Equal(t, float64(200), loginResp["code"], "Login should succeed") + + token := loginResp["token"].(string) + assert.NotEmpty(t, token, "Token should not be empty") + + // Test refresh endpoint (which checks token validity) + w = httptest.NewRecorder() + req, _ = http.NewRequest("GET", "/refresh", nil) + req.Header.Set("Authorization", "Bearer "+token) + router.ServeHTTP(w, req) + + // Should be successful with valid token + assert.Equal(t, http.StatusOK, w.Code, "Token should be valid immediately after login") +} + +// TestJWTInvalidSignature tests that malformed JWT tokens are rejected. +func TestJWTInvalidSignature(t *testing.T) { + gin.SetMode(gin.TestMode) + router := setupRouter() + + // Test /refresh with invalid token + w := httptest.NewRecorder() + req, _ := http.NewRequest("GET", "/refresh", nil) + req.Header.Set("Authorization", "Bearer invalid.token.signature") + router.ServeHTTP(w, req) + + // Should fail with invalid token + assert.Equal(t, http.StatusUnauthorized, w.Code, "Invalid JWT should be rejected") +} + +// TestJWTMissing tests that missing JWT tokens are rejected. +func TestJWTMissing(t *testing.T) { + gin.SetMode(gin.TestMode) + router := setupRouter() + + // Test /refresh without token + w := httptest.NewRecorder() + req, _ := http.NewRequest("GET", "/refresh", nil) + router.ServeHTTP(w, req) + + // Should fail without token + assert.Equal(t, http.StatusUnauthorized, w.Code, "Missing JWT should be rejected") +} + +// TestBasicAuthMissingCredentials tests that missing credentials are rejected. +func TestBasicAuthMissingCredentials(t *testing.T) { + gin.SetMode(gin.TestMode) + router := setupRouter() + + // Test /prometheus/targets without basic auth + w := httptest.NewRecorder() + req, _ := http.NewRequest("GET", "/prometheus/targets", nil) + router.ServeHTTP(w, req) + + // Should fail without credentials + assert.Equal(t, http.StatusUnauthorized, w.Code, "Missing BasicAuth should be rejected") +} + +// TestLoginFailure tests that login fails with wrong password. +func TestLoginFailure(t *testing.T) { + gin.SetMode(gin.TestMode) + router := setupRouter() + + w := httptest.NewRecorder() + loginBody := `{"username": "admin", "password": "wrongpassword"}` + req, _ := http.NewRequest("POST", "/login", bytes.NewBuffer([]byte(loginBody))) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(w, req) + + // Should fail with wrong password + assert.Equal(t, http.StatusUnauthorized, w.Code, "Login should fail with wrong password") +} + +// TestLogoutInvalidatesToken tests that logout invalidates the token. +func TestLogoutInvalidatesToken(t *testing.T) { + gin.SetMode(gin.TestMode) + router := setupRouter() + + // Login first + w := httptest.NewRecorder() + loginBody := `{"username": "admin", "password": "admin"}` + req, _ := http.NewRequest("POST", "/login", bytes.NewBuffer([]byte(loginBody))) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(w, req) + + var loginResp map[string]interface{} + json.NewDecoder(w.Body).Decode(&loginResp) + token := loginResp["token"].(string) + + // Verify token is valid before logout + w = httptest.NewRecorder() + req, _ = http.NewRequest("GET", "/refresh", nil) + req.Header.Set("Authorization", "Bearer "+token) + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusOK, w.Code, "Token should be valid before logout") + + // Logout + w = httptest.NewRecorder() + req, _ = http.NewRequest("POST", "/logout", nil) + req.Header.Set("Authorization", "Bearer "+token) + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusOK, w.Code, "Logout should succeed") + + // Try to use token after logout (should fail) + w = httptest.NewRecorder() + req, _ = http.NewRequest("GET", "/refresh", nil) + req.Header.Set("Authorization", "Bearer "+token) + router.ServeHTTP(w, req) + // The middleware returns 403 for failed authorization checks (CheckTokenValidation returns false) + assert.Equal(t, http.StatusForbidden, w.Code, "Token should be invalid after logout") +} + +// TestAuthHeaderFormats tests different JWT header formats. +func TestAuthHeaderFormats(t *testing.T) { + gin.SetMode(gin.TestMode) + router := setupRouter() + + // Login to get token + w := httptest.NewRecorder() + loginBody := `{"username": "admin", "password": "admin"}` + req, _ := http.NewRequest("POST", "/login", bytes.NewBuffer([]byte(loginBody))) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(w, req) + + var loginResp map[string]interface{} + json.NewDecoder(w.Body).Decode(&loginResp) + token := loginResp["token"].(string) + + // Test valid "Bearer" format + w = httptest.NewRecorder() + req, _ = http.NewRequest("GET", "/refresh", nil) + req.Header.Set("Authorization", "Bearer "+token) + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusOK, w.Code, "Standard Bearer format should work") + + // Test without "Bearer" prefix (should fail) + w = httptest.NewRecorder() + req, _ = http.NewRequest("GET", "/refresh", nil) + req.Header.Set("Authorization", token) // Missing "Bearer" prefix + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusUnauthorized, w.Code, "Missing Bearer prefix should fail") +} + +// TestAdminPrivilegeCheck tests that non-admin users cannot access admin endpoints. +func TestAdminPrivilegeCheck(t *testing.T) { + gin.SetMode(gin.TestMode) + router := setupRouter() + + // Create and login with limited user + w := httptest.NewRecorder() + adminLoginBody := `{"username": "admin", "password": "admin"}` + req, _ := http.NewRequest("POST", "/login", bytes.NewBuffer([]byte(adminLoginBody))) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(w, req) + + var adminLoginResp map[string]interface{} + json.NewDecoder(w.Body).Decode(&adminLoginResp) + adminToken := adminLoginResp["token"].(string) + + // Create limited user + w = httptest.NewRecorder() + addBody := `{"username": "limiteduser2", "password": "limited", "display_name": "Limited", "admin": false}` + req, _ = http.NewRequest("POST", "/accounts", bytes.NewBuffer([]byte(addBody))) + req.Header.Set("Authorization", "Bearer "+adminToken) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(w, req) + + // Login as limited user + w = httptest.NewRecorder() + limitedLoginBody := `{"username": "limiteduser2", "password": "limited"}` + req, _ = http.NewRequest("POST", "/login", bytes.NewBuffer([]byte(limitedLoginBody))) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(w, req) + + var limitedLoginResp map[string]interface{} + json.NewDecoder(w.Body).Decode(&limitedLoginResp) + limitedToken := limitedLoginResp["token"].(string) + + // Try to access /accounts endpoint (admin only) + w = httptest.NewRecorder() + req, _ = http.NewRequest("GET", "/accounts", nil) + req.Header.Set("Authorization", "Bearer "+limitedToken) + router.ServeHTTP(w, req) + + // Should be forbidden (403) + assert.Equal(t, http.StatusForbidden, w.Code, "Non-admin user should not access admin endpoint") +} + +// TestTokenValidationCaching tests token caching for performance. +func TestTokenValidationCaching(t *testing.T) { + gin.SetMode(gin.TestMode) + router := setupRouter() + + // Login to get token + w := httptest.NewRecorder() + loginBody := `{"username": "admin", "password": "admin"}` + req, _ := http.NewRequest("POST", "/login", bytes.NewBuffer([]byte(loginBody))) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(w, req) + + var loginResp map[string]interface{} + json.NewDecoder(w.Body).Decode(&loginResp) + token := loginResp["token"].(string) + + // Make multiple requests with the same token + for i := 0; i < 3; i++ { + w = httptest.NewRecorder() + req, _ = http.NewRequest("GET", "/refresh", nil) + req.Header.Set("Authorization", "Bearer "+token) + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusOK, w.Code, "Cached token should remain valid") + } +} + +// TestConcurrentRequests tests that middleware handles concurrent requests safely. +func TestConcurrentRequests(t *testing.T) { + gin.SetMode(gin.TestMode) + router := setupRouter() + + // Login to get token + w := httptest.NewRecorder() + loginBody := `{"username": "admin", "password": "admin"}` + req, _ := http.NewRequest("POST", "/login", bytes.NewBuffer([]byte(loginBody))) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(w, req) + + var loginResp map[string]interface{} + json.NewDecoder(w.Body).Decode(&loginResp) + token := loginResp["token"].(string) + + // Simulate concurrent requests (simplified test) + results := make(chan int, 5) + for i := 0; i < 5; i++ { + go func() { + w := httptest.NewRecorder() + req, _ := http.NewRequest("GET", "/refresh", nil) + req.Header.Set("Authorization", "Bearer "+token) + router.ServeHTTP(w, req) + results <- w.Code + }() + } + + // Verify all requests succeeded + successCount := 0 + for i := 0; i < 5; i++ { + code := <-results + if code == http.StatusOK { + successCount++ + } + } + assert.Equal(t, 5, successCount, "All concurrent requests should succeed") +} diff --git a/controller/api/models/account.go b/controller/api/models/account.go new file mode 100644 index 00000000..bb1e7f66 --- /dev/null +++ b/controller/api/models/account.go @@ -0,0 +1,39 @@ +/* + * Copyright (C) 2024 Nethesis S.r.l. + * http://www.nethesis.it - info@nethesis.it + * + * SPDX-License-Identifier: GPL-2.0-only + * + * author: Edoardo Spadoni + */ + +package models + +import ( + "time" +) + +type Account struct { + ID int `json:"id" structs:"id"` + Username string `json:"username" structs:"username" binding:"required,excludesall= "` + Password string `json:"password" structs:"password" db:"-" binding:"required"` + // Watch out: un/marshalling booleans is a pain, see https://github.com/gin-gonic/gin/issues/814 + Admin bool `json:"admin" structs:"admin"` + DisplayName string `json:"display_name" structs:"display_name"` + UnitGroups []int `json:"unit_groups" structs:"unit_groups"` + Created time.Time `json:"created" structs:"created_at"` + Updated time.Time `json:"updated" structs:"updated_at"` + TwoFA bool `json:"two_fa" structs:"two_fa"` +} + +type AccountUpdate struct { + Password string `json:"password" structs:"password"` + DisplayName string `json:"display_name" structs:"display_name"` + Admin bool `json:"admin" structs:"admin"` + UnitGroups []int `json:"unit_groups" structs:"unit_groups" binding:"required"` +} + +type PasswordChange struct { + OldPassword string `json:"old_password" binding:"required"` + NewPassword string `json:"new_password" binding:"required"` +} diff --git a/controller/api/models/auth.go b/controller/api/models/auth.go new file mode 100644 index 00000000..f8f6614b --- /dev/null +++ b/controller/api/models/auth.go @@ -0,0 +1,40 @@ +/* + * Copyright (C) 2024 Nethesis S.r.l. + * http://www.nethesis.it - info@nethesis.it + * + * SPDX-License-Identifier: GPL-2.0-only + * + * author: Edoardo Spadoni + */ + +package models + +type LoginRequest struct { + Username string `json:"username" binding:"required"` + Password string `json:"password" binding:"required"` + // controller user originating the request, ignored by older units + OnBehalfOf string `json:"on_behalf_of,omitempty"` +} + +type LoginResponse struct { + Code int `json:"code" binding:"required"` + Expire string `json:"expire" binding:"required"` + Token string `json:"token" binding:"required"` +} + +type SSHGenerate struct { + Passphrase string `json:"passphrase" binding:"required"` +} + +type OTPJson struct { + Username string `json:"username" structs:"username"` + Token string `json:"token" structs:"token"` + OTP string `json:"otp" structs:"otp"` +} + +type UserAuthorizations struct { + Username string `json:"username" structs:"username"` + Role string `json:"role" structs:"role"` + Actions []string `json:"actions" structs:"actions"` + SudoRequested bool `json:"sudo_requested" structs:"sudo_requested"` +} diff --git a/controller/api/models/platform.go b/controller/api/models/platform.go new file mode 100644 index 00000000..3d660d95 --- /dev/null +++ b/controller/api/models/platform.go @@ -0,0 +1,18 @@ +/* + * Copyright (C) 2025 Nethesis S.r.l. + * http://www.nethesis.it - info@nethesis.it + * + * SPDX-License-Identifier: GPL-2.0-only + * + * author: Giacomo Sanchietti + */ + +package models + +type PlatformInfo struct { + VpnPort string `json:"vpn_port" structs:"vpn_port"` + VpnNetwork string `json:"vpn_network" structs:"vpn_network"` + ControllerVersion string `json:"controller_version" structs:"controller_version"` + MetricsRetentionDays int `json:"metrics_retention_days" structs:"metrics_retention_days"` + LogsRetentionDays int `json:"logs_retention_days" structs:"logs_retention_days"` +} diff --git a/controller/api/models/report.go b/controller/api/models/report.go new file mode 100644 index 00000000..af47cd35 --- /dev/null +++ b/controller/api/models/report.go @@ -0,0 +1,111 @@ +/* + * Copyright (C) 2024 Nethesis S.r.l. + * http://www.nethesis.it - info@nethesis.it + * + * SPDX-License-Identifier: GPL-2.0-only + * + * author: Giacomo Sanchietti + */ + +package models + +type MwanEvent struct { + Timestamp int64 `json:"timestamp" binding:"required"` + Wan string `json:"wan" binding:"required"` + Interface string `json:"interface" binding:"required"` + Event string `json:"event" binding:"required"` +} + +type MwanEvents []MwanEvent + +type MwanEventRequest struct { + Data MwanEvents `json:"data" binding:"required"` +} + +type TsAttack struct { + Timestamp int64 `json:"timestamp" binding:"required"` + Ip string `json:"ip" binding:"required"` +} + +type TsAttacks []TsAttack + +type TsAttackRequest struct { + Data TsAttacks `json:"data" binding:"required"` +} + +type TsMalware struct { + Timestamp int64 `json:"timestamp" binding:"required"` + Src string `json:"src" binding:"required"` + Dst string `json:"dst" binding:"required"` + Category string `json:"category" binding:"required"` + Chain string `json:"chain" binding:"required"` +} + +type TsMalwares []TsMalware + +type TsMalwareRequest struct { + Data TsMalwares `json:"data" binding:"required"` +} + +type OvpnRwConnection struct { + Timestamp int64 `json:"timestamp" binding:"required"` + Instance string `json:"instance" binding:"required"` + CommonName string `json:"common_name" binding:"required"` + VirtualIpAddr string `json:"virtual_ip_addr" binding:"required"` + RemoteIpAddr string `json:"remote_ip_addr" binding:"required"` + StartTime int64 `json:"start_time" binding:"required"` + Duration int64 `json:"duration" binding:"required"` + BytesReceived int64 `json:"bytes_received" binding:"required"` + BytesSent int64 `json:"bytes_sent" binding:"required"` +} + +type OvpnRwConnections []OvpnRwConnection + +type OvpnRwConnectionsRequest struct { + Data OvpnRwConnections `json:"data" binding:"required"` +} + +type DpiStat struct { + Timestamp int64 `json:"timestamp" binding:"required"` + ClientAddress string `json:"client_address" binding:"required"` + ClientName string `json:"client_name" binding:"required"` + Protocol string `json:"protocol"` + Host string `json:"host"` + Application string `json:"application"` + Bytes int64 `json:"bytes" binding:"required"` +} + +type DpiStats []DpiStat + +type DpiStatsRequest struct { + Data DpiStats `json:"data" binding:"required"` +} + +type UnitNameRequest struct { + Name string `json:"name" binding:"required"` +} + +type OpenVPNConfiguration struct { + Instance string `json:"instance" binding:"required"` + Name string `json:"name" binding:"required"` + Device string `json:"device" binding:"required"` + Type string `json:"type"` // valid values are: rw (for roadwarrior), client (for tunnel client), server (for tunnel server) +} + +type OpenVPNConfigurations []OpenVPNConfiguration + +type UnitOpenVPNRWRequest struct { + Data OpenVPNConfigurations `json:"data" binding:"required"` +} + +type Wan struct { + Interface string `json:"interface" binding:"required"` + Device string `json:"device" binding:"required"` + Status string `json:"status" binding:"required"` +} + +type Wans []Wan + +type UnitWanRequest struct { + Data Wans `json:"data" binding:"required"` +} diff --git a/controller/api/models/ubus.go b/controller/api/models/ubus.go new file mode 100644 index 00000000..aaab0534 --- /dev/null +++ b/controller/api/models/ubus.go @@ -0,0 +1,24 @@ +/* + * Copyright (C) 2024 Nethesis S.r.l. + * http://www.nethesis.it - info@nethesis.it + * + * SPDX-License-Identifier: GPL-2.0-only + * + * author: Giacomo Sanchietti + */ + +package models + +type UbusCommand struct { + Path string `json:"path" binding:"required"` + Method string `json:"method" binding:"required"` + Payload map[string]interface{} `json:"payload" binding:"required"` +} + +// Response example: +// {"code":200,"data":{"depends": "on the call"},"message":"ubus call action success"} +type UbusResponse[T any] struct { + Code int `json:"code"` + Data T `json:"data"` + Message string `json:"message"` +} diff --git a/controller/api/models/unit.go b/controller/api/models/unit.go new file mode 100644 index 00000000..3143d47c --- /dev/null +++ b/controller/api/models/unit.go @@ -0,0 +1,67 @@ +/* + * Copyright (C) 2024 Nethesis S.r.l. + * http://www.nethesis.it - info@nethesis.it + * + * SPDX-License-Identifier: GPL-2.0-only + * + * author: Edoardo Spadoni + */ + +package models + +import ( + "time" +) + +type AddRequest struct { + UnitId string `json:"unit_id" binding:"required"` +} + +type RegisterRequest struct { + UnitId string `json:"unit_id" binding:"required"` + Username string `json:"username" binding:"required"` + Password string `json:"password" binding:"required"` + UnitName string `json:"unit_name" binding:"required"` + Version string `json:"version"` + SubscriptionType string `json:"subscription_type"` + SystemId string `json:"system_id"` +} + +type Unit struct { + ID string `json:"unit_id" structs:"id"` + Name string `json:"unit_name" structs:"name"` + Version string `json:"version" structs:"version"` + SubscriptionType string `json:"subscription_type" structs:"subscription_type"` + SystemID string `json:"system_id" structs:"system_id"` + Created time.Time `json:"created" structs:"created"` +} + +type UnitInfo struct { + UnitName string `json:"unit_name"` + Version string `json:"version"` + VersionUpdate string `json:"version_update"` + ScheduledUpdate int `json:"scheduled_update"` + SubscriptionType string `json:"subscription_type"` + SystemID string `json:"system_id"` + SSHPort int `json:"ssh_port"` + FQDN string `json:"fqdn"` + APIVersion string `json:"api_version"` // ns-api package + Description string `json:"description"` + UIVersion string `json:"ui_version"` // ns-ui package +} + +type CheckSystemUpdate struct { + LastVersion string `json:"lastVersion"` + ScheduledAt int `json:"scheduledAt"` + CurrentVersion string `json:"currentVersion"` +} + +type UnitGroup struct { + ID int `json:"id" structs:"id"` + Name string `json:"name" structs:"name"` + Description string `json:"description" structs:"description"` + Units []string `json:"units" structs:"units"` + CreatedAt time.Time `json:"created_at" structs:"created_at"` + UpdatedAt time.Time `json:"updated_at" structs:"updated_at"` + UsedBy []string `json:"used_by" structs:"used_by"` +} diff --git a/controller/api/report_test.go b/controller/api/report_test.go new file mode 100644 index 00000000..4b71b2fc --- /dev/null +++ b/controller/api/report_test.go @@ -0,0 +1,142 @@ +/* + * Copyright (C) 2025 Nethesis S.r.l. + * http://www.nethesis.it - info@nethesis.it + * + * SPDX-License-Identifier: GPL-2.0-only + * + * author: Giacomo Sanchietti + */ + +package main + +import ( + "bytes" + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + + "github.com/NethServer/nethsecurity-controller/api/models" + "github.com/stretchr/testify/assert" +) + +// These tests exercise the /ingest/* endpoints and require a database. +// Set REPORT_DB_URI environment variable before running tests. + +func TestIngestUnitName(t *testing.T) { + ginRouter := setupRouter() + + unitID := addUnit(t) + + payload := models.UnitNameRequest{Name: "test-unit-name"} + b, _ := json.Marshal(payload) + + req := httptest.NewRequest("POST", "/ingest/dump-nsplug-config", bytes.NewBuffer(b)) + req.SetBasicAuth(unitID, "1234") + req.Header.Set("Content-Type", "application/json") + + w := httptest.NewRecorder() + ginRouter.ServeHTTP(w, req) + assert.Equal(t, http.StatusOK, w.Code, "set unit name should return 200") +} + +func TestIngestWanConfig(t *testing.T) { + ginRouter := setupRouter() + unitID := addUnit(t) + + reqBody := `{"data":[{"interface":"wan","device":"eth0","status":"up"},{"interface":"wan2","device":"eth1","status":"down"}]}` + req := httptest.NewRequest("POST", "/ingest/dump-wan-config", bytes.NewBufferString(reqBody)) + req.SetBasicAuth(unitID, "1234") + req.Header.Set("Content-Type", "application/json") + + w := httptest.NewRecorder() + ginRouter.ServeHTTP(w, req) + assert.Equal(t, http.StatusOK, w.Code, "set wan config should return 200") +} + +func TestIngestOvpnConfig(t *testing.T) { + ginRouter := setupRouter() + unitID := addUnit(t) + + reqBody := `{"data":[{"instance":"server1","name":"srv","device":"tun0","type":"server"}]}` + req := httptest.NewRequest("POST", "/ingest/dump-ovpn-config", bytes.NewBufferString(reqBody)) + req.SetBasicAuth(unitID, "1234") + req.Header.Set("Content-Type", "application/json") + + w := httptest.NewRecorder() + ginRouter.ServeHTTP(w, req) + assert.Equal(t, http.StatusOK, w.Code, "set ovpn config should return 200") +} + +func TestIngestMwanEvents(t *testing.T) { + ginRouter := setupRouter() + unitID := addUnit(t) + + reqBody := `{"data":[{"timestamp":1600000000,"wan":"wan","event":"up","interface":"eth0"},{"timestamp":1600000001,"wan":"wan2","event":"down","interface":"eth1"}]}` + req := httptest.NewRequest("POST", "/ingest/dump-mwan-events", bytes.NewBufferString(reqBody)) + req.SetBasicAuth(unitID, "1234") + req.Header.Set("Content-Type", "application/json") + + w := httptest.NewRecorder() + ginRouter.ServeHTTP(w, req) + assert.Equal(t, http.StatusOK, w.Code, "update mwan events should return 200") +} + +func TestIngestTsAttacksAndDpiAndOvpnConnections(t *testing.T) { + ginRouter := setupRouter() + unitID := addUnit(t) + + // TS attacks + attacks := `{"data":[{"timestamp":1600000000,"ip":"8.8.8.8"}]}` + req := httptest.NewRequest("POST", "/ingest/dump-ts-attacks", bytes.NewBufferString(attacks)) + req.SetBasicAuth(unitID, "1234") + req.Header.Set("Content-Type", "application/json") + w := httptest.NewRecorder() + ginRouter.ServeHTTP(w, req) + assert.Equal(t, http.StatusOK, w.Code, "update ts attacks should return 200") + + // DPI stats + dpi := `{"data":[{"timestamp":1600000000,"client_address":"10.0.0.1","bytes":1234,"client_name":"host","protocol":"tcp","host":"example.com","application":"http"}]}` + req = httptest.NewRequest("POST", "/ingest/dump-dpi-stats", bytes.NewBufferString(dpi)) + req.SetBasicAuth(unitID, "1234") + req.Header.Set("Content-Type", "application/json") + w = httptest.NewRecorder() + ginRouter.ServeHTTP(w, req) + assert.Equal(t, http.StatusOK, w.Code, "update dpi stats should return 200") + + // OVPN connections + ovpn := `{"data":[{"timestamp":1600000000,"instance":"server1","common_name":"cn","virtual_ip_addr":"10.8.0.2","remote_ip_addr":"1.2.3.4","start_time":1600000000,"duration":60,"bytes_received":100,"bytes_sent":200}]}` + req = httptest.NewRequest("POST", "/ingest/dump-ovpn-connections", bytes.NewBufferString(ovpn)) + req.SetBasicAuth(unitID, "1234") + req.Header.Set("Content-Type", "application/json") + w = httptest.NewRecorder() + ginRouter.ServeHTTP(w, req) + assert.Equal(t, http.StatusOK, w.Code, "update ovpn connections should return 200") +} + +func TestIngestInvalidData(t *testing.T) { + // This test does not need DB; send malformed JSON or missing required fields to ensure handler returns 400 + ginRouter := setupRouter() + unitID := addUnit(t) + + // malformed JSON + req := httptest.NewRequest("POST", "/ingest/dump-dpi-stats", bytes.NewBufferString("{notjson")) + req.SetBasicAuth(unitID, "1234") + req.Header.Set("Content-Type", "application/json") + w := httptest.NewRecorder() + ginRouter.ServeHTTP(w, req) + assert.Equal(t, http.StatusBadRequest, w.Code, "malformed JSON should return 400") +} + +func TestIngestUnauthorized(t *testing.T) { + // wrong credentials should be unauthorized + ginRouter := setupRouter() + unitID := addUnit(t) + + req := httptest.NewRequest("POST", "/ingest/dump-wan-config", bytes.NewBufferString(`{"data":[]}`)) + req.SetBasicAuth(unitID, "wrongtoken") + req.Header.Set("Content-Type", "application/json") + w := httptest.NewRecorder() + ginRouter.ServeHTTP(w, req) + assert.Equal(t, http.StatusUnauthorized, w.Code, "wrong basic auth should return 401") +} diff --git a/controller/api/response/response.go b/controller/api/response/response.go new file mode 100644 index 00000000..3dc882ad --- /dev/null +++ b/controller/api/response/response.go @@ -0,0 +1,54 @@ +/* + * Copyright (C) 2024 Nethesis S.r.l. + * SPDX-License-Identifier: GPL-2.0-only + */ + +package response + +type StatusOK struct { + Code int `json:"code" example:"200" structs:"code"` + Message string `json:"message" example:"Success" structs:"message"` + Data interface{} `json:"data" structs:"data"` +} + +type StatusBadRequest struct { + Code int `json:"code" example:"400" structs:"code"` + Message string `json:"message" example:"Bad request" structs:"message"` + Data interface{} `json:"data" structs:"data"` +} + +type StatusNotFound struct { + Code int `json:"code" example:"404" structs:"code"` + Message string `json:"message" example:"Not found" structs:"message"` + Data interface{} `json:"data" structs:"data"` +} + +type StatusCreated struct { + Code int `json:"code" example:"201" structs:"code"` + Message string `json:"message" example:"Created" structs:"message"` + Data interface{} `json:"data" structs:"data"` +} + +type StatusUnauthorized struct { + Code int `json:"code" example:"401" structs:"code"` + Message string `json:"message" example:"Unauthorized" structs:"message"` + Data interface{} `json:"data" structs:"data"` +} + +type StatusForbidden struct { + Code int `json:"code" example:"403" structs:"code"` + Message string `json:"message" example:"Forbidden" structs:"message"` + Data interface{} `json:"data" structs:"data"` +} + +type StatusConflict struct { + Code int `json:"code" example:"409" structs:"code"` + Message string `json:"message" example:"Not found" structs:"message"` + Data interface{} `json:"data" structs:"data"` +} + +type StatusInternalServerError struct { + Code int `json:"code" example:"500" structs:"code"` + Message string `json:"message" example:"Internal server error" structs:"message"` + Data interface{} `json:"data" structs:"data"` +} diff --git a/controller/api/routines/routines.go b/controller/api/routines/routines.go new file mode 100644 index 00000000..39305531 --- /dev/null +++ b/controller/api/routines/routines.go @@ -0,0 +1,47 @@ +/* + * Copyright (C) 2024 Nethesis S.r.l. + * http://www.nethesis.it - info@nethesis.it + * + * SPDX-License-Identifier: GPL-2.0-only + * + * author: Edoardo Spadoni + */ + +package routines + +import ( + "github.com/NethServer/nethsecurity-controller/api/utils" + "time" + + "github.com/NethServer/nethsecurity-controller/api/logs" + "github.com/NethServer/nethsecurity-controller/api/methods" +) + +func RefreshRemoteInfoLoop() { + ticker := time.NewTicker(1 * time.Hour) + for range ticker.C { + // load all units info into cache + units, err := methods.ListConnectedUnits() + if err != nil { + return + } + + for _, unit := range units { + // no user originated this request + _, err := methods.GetRemoteInfo(unit, "") + if err != nil { + logs.Logs.Println("[ERR][ROUTINE] loop for remote info failed: " + err.Error()) + } + } + } +} + +func RefreshGeoIPDatabase() { + ticker := time.NewTicker(24 * time.Hour) + for range ticker.C { + err := utils.InitGeoIP() + if err != nil { + logs.Logs.Println("[ERR][ROUTINE] loop for geoip database failed: " + err.Error()) + } + } +} diff --git a/controller/api/socket/socket.go b/controller/api/socket/socket.go new file mode 100644 index 00000000..9c875ae6 --- /dev/null +++ b/controller/api/socket/socket.go @@ -0,0 +1,59 @@ +/* + * Copyright (C) 2023 Nethesis S.r.l. + * http://www.nethesis.it - info@nethesis.it + * + * SPDX-License-Identifier: GPL-2.0-only + * + * author: Edoardo Spadoni + */ + +package socket + +import ( + "net" + + "github.com/NethServer/nethsecurity-controller/api/configuration" + "github.com/NethServer/nethsecurity-controller/api/logs" +) + +var Socket net.Conn + +func Init() { + //establish connection + connection, err := net.Dial("unix", configuration.Config.OpenVPNMGMTSock) + + // check error + if err != nil { + logs.Logs.Println("[ERR][OPENVPN SOCKET] can't connect to openvpn socket: " + configuration.Config.OpenVPNMGMTSock) + } + + // assign object + Socket = connection +} + +func Write(message string) string { + // avoid panic if socket is not initialized + if Socket == nil { + return "" + } + + // compose message + _, err := Socket.Write([]byte(message + "\n")) + + // check write error + if err != nil { + logs.Logs.Println("[ERR][OPENVPN SOCKET] can't write to openvpn socket: " + configuration.Config.OpenVPNMGMTSock + ". error: " + err.Error()) + } + + // compose buffer + buffer := make([]byte, 4096) + bLen, err := Socket.Read(buffer) + + // check read error + if err != nil { + logs.Logs.Println("[ERR][OPENVPN SOCKET] can't read from openvpn socket: " + configuration.Config.OpenVPNMGMTSock + ". error: " + err.Error()) + } + + // return string + return string(buffer[:bLen]) +} diff --git a/controller/api/socket/socket_test.go b/controller/api/socket/socket_test.go new file mode 100644 index 00000000..04ce4b04 --- /dev/null +++ b/controller/api/socket/socket_test.go @@ -0,0 +1,40 @@ +/* + * Copyright (C) 2025 Nethesis S.r.l. + * http://www.nethesis.it - info@nethesis.it + * + * SPDX-License-Identifier: GPL-2.0-only + * + * author: Giacomo Sanchietti + */ + +package socket + +import ( + "testing" + + "github.com/NethServer/nethsecurity-controller/api/configuration" + "github.com/NethServer/nethsecurity-controller/api/logs" + "github.com/stretchr/testify/assert" +) + +func TestMain(m *testing.M) { + logs.Init("test") + // Set config + configuration.Config.OpenVPNMGMTSock = "/tmp/nonexistent.sock" + m.Run() +} + +func TestInit(t *testing.T) { + // Test that Init doesn't panic even if socket doesn't exist + assert.NotPanics(t, func() { + Init() + }) + // Socket should be nil since connection fails + assert.Nil(t, Socket) +} + +func TestWrite(t *testing.T) { + // Test Write when Socket is nil + result := Write("test") + assert.Equal(t, "", result) +} diff --git a/controller/api/storage/grafana_user.sql.tmpl b/controller/api/storage/grafana_user.sql.tmpl new file mode 100644 index 00000000..a6d42da2 --- /dev/null +++ b/controller/api/storage/grafana_user.sql.tmpl @@ -0,0 +1,13 @@ +DO +$do$ + BEGIN + IF NOT EXISTS (SELECT + FROM pg_catalog.pg_roles + WHERE rolname = 'grafana') THEN + CREATE USER grafana WITH PASSWORD '{{ .GrafanaPostgresPassword }}'; + END IF; + END +$do$; + +GRANT USAGE ON SCHEMA public TO grafana; +GRANT SELECT ON ALL TABLES IN SCHEMA public TO grafana; diff --git a/controller/api/storage/report_schema.sql.tmpl b/controller/api/storage/report_schema.sql.tmpl new file mode 100644 index 00000000..32d5f1de --- /dev/null +++ b/controller/api/storage/report_schema.sql.tmpl @@ -0,0 +1,587 @@ +/* + * Copyright (C) 2024 Nethesis S.r.l. + * http://www.nethesis.it - info@nethesis.it + * + * SPDX-License-Identifier: GPL-2.0-only + * + * author: Giacomo Sanchietti + */ + +--------------------------------------------------------------------------------------------- +-- CORE CONFIGURATION TABLES +-- These tables are part of the core, used to store the core configuration of the application +--------------------------------------------------------------------------------------------- + +-- This table contains all user accounts +CREATE TABLE IF NOT EXISTS accounts ( + id SERIAL PRIMARY KEY, + username TEXT NOT NULL UNIQUE, + password TEXT NOT NULL, + admin BOOLEAN NOT NULL DEFAULT FALSE, + display_name TEXT, + otp_secret TEXT, + otp_recovery_codes TEXT, + unit_groups int[], + created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP, + updated_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP +); + +-- This table contains all unit groups +CREATE TABLE IF NOT EXISTS unit_groups ( + id SERIAL PRIMARY KEY, + name TEXT NOT NULL UNIQUE, + description TEXT, + units uuid[], + created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP, + updated_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP +); + +-- This table contains all units +CREATE TABLE IF NOT EXISTS units ( + uuid UUID PRIMARY KEY, + name TEXT, + info JSONB, + vpn_address TEXT, + vpn_connected_since TIMESTAMP, + created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP, + updated_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP +); + +-- This table contains all unit credentials +CREATE TABLE IF NOT EXISTS unit_credentials ( + uuid UUID PRIMARY KEY, + username TEXT, + password TEXT +); + +------------------------------------------------------------------ +-- REPORT TABLES +-- All the following tables are used for reporting inside Grafana +------------------------------------------------------------------ + +-- This table contains the list of OpenVPN instances configured inside the units +CREATE TABLE IF NOT EXISTS openvpn_config ( + uuid UUID NOT NULL, + instance TEXT NOT NULL, + name TEXT, + device TEXT NOT NULL, + type TEXT NOT NULL, + created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP +); +CREATE INDEX IF NOT EXISTS idx_openvpn_config_uuid ON openvpn_config (uuid); + +-- This table contains the list of WAN interfaces configured inside the units +CREATE TABLE IF NOT EXISTS wan_config ( + uuid UUID NOT NULL, + interface TEXT NOT NULL, + device TEXT, + status TEXT, + created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP +); +CREATE INDEX IF NOT EXISTS idx_wan_config_uuid ON wan_config (uuid); + +-- General retention policies + +-- Keep raw data for 30 days +-- Keep downsampled data for 60 days + + +--------------- +-- Mwan events +--------------- + +CREATE TABLE IF NOT EXISTS mwan_events ( + time TIMESTAMPTZ NOT NULL, + uuid UUID NOT NULL, + wan TEXT NOT NULL, + event TEXT NOT NULL, + interface TEXT, + UNIQUE (time, uuid) +); +CREATE INDEX IF NOT EXISTS idx_mwan_events_uuid ON mwan_events (uuid); + +SELECT + create_hypertable('mwan_events', by_range('time'), if_not_exists => TRUE); + +-- Drop raw data after 30 days +SELECT remove_retention_policy('mwan_events', if_exists => TRUE); +SELECT add_retention_policy('mwan_events', drop_after => INTERVAL '{{ .RetentionDays }} days', if_not_exists => TRUE); + +-- Continuous aggregates and retention policy + +CREATE MATERIALIZED VIEW IF NOT EXISTS ca_mwan_events_hourly +WITH (timescaledb.continuous, timescaledb.materialized_only = false) AS +SELECT uuid, + time_bucket(INTERVAL '1 hour', time) AS bucket, + wan, + event, + interface +FROM mwan_events +GROUP BY uuid, bucket, wan, event, interface +WITH NO DATA; + +SELECT add_continuous_aggregate_policy('ca_mwan_events_hourly', + start_offset => NULL, + end_offset => INTERVAL '30 minutes', + schedule_interval => INTERVAL '15 minutes', + if_not_exists => TRUE +); + +-- Drop downsampled data after 60 days +SELECT remove_retention_policy('ca_mwan_events_hourly', if_exists => TRUE); +SELECT add_retention_policy('ca_mwan_events_hourly', drop_after => INTERVAL '{{ .RetentionDays }} days', if_not_exists => TRUE); + +--------------- +-- Malware +--------------- + +CREATE TABLE IF NOT EXISTS ts_malware ( + time TIMESTAMPTZ NOT NULL, + uuid UUID NOT NULL, + src TEXT NOT NULL, + dst TEXT NOT NULL, + category TEXT NOT NULL, + chain TEXT NOT NULL, + country VARCHAR(2), + UNIQUE (time, uuid) +); +CREATE INDEX IF NOT EXISTS idx_ts_malware_uuid ON ts_malware (uuid); + +SELECT + create_hypertable('ts_malware', by_range('time'), if_not_exists => TRUE); + +-- Drop raw data after 30 days +SELECT remove_retention_policy('ts_malware', if_exists => TRUE); +SELECT add_retention_policy('ts_malware', drop_after => INTERVAL '{{ .RetentionDays }} days', if_not_exists => TRUE); + +-- Continuous aggregates + +CREATE MATERIALIZED VIEW IF NOT EXISTS ca_ts_malware_hourly_direction +WITH (timescaledb.continuous, timescaledb.materialized_only = false) AS +SELECT uuid, + time_bucket(INTERVAL '1 hour', time) AS bucket, + src, + dst +FROM ts_malware +GROUP BY uuid, bucket, src, dst +WITH NO DATA; + +SELECT add_continuous_aggregate_policy('ca_ts_malware_hourly_direction', + start_offset => NULL, + end_offset => INTERVAL '30 minutes', + schedule_interval => INTERVAL '15 minutes', + if_not_exists => TRUE +); + +SELECT remove_retention_policy('ca_ts_malware_hourly_direction', if_exists => TRUE); +SELECT add_retention_policy('ca_ts_malware_hourly_direction', drop_after => INTERVAL '{{ .RetentionDays }} days', if_not_exists => TRUE); + +CREATE MATERIALIZED VIEW IF NOT EXISTS ca_ts_malware_hourly_category +WITH (timescaledb.continuous, timescaledb.materialized_only = false) AS +SELECT uuid, + time_bucket(INTERVAL '1 hour', time) AS bucket, + category, + count(category) as count +FROM ts_malware +GROUP BY uuid, bucket, category +WITH NO DATA; + +SELECT add_continuous_aggregate_policy('ca_ts_malware_hourly_category', + start_offset => NULL, + end_offset => INTERVAL '30 minutes', + schedule_interval => INTERVAL '15 minutes', + if_not_exists => TRUE +); + +SELECT remove_retention_policy('ca_ts_malware_hourly_category', if_exists => TRUE); +SELECT add_retention_policy('ca_ts_malware_hourly_category', drop_after => INTERVAL '{{ .RetentionDays }} days', if_not_exists => TRUE); + +CREATE MATERIALIZED VIEW IF NOT EXISTS ca_ts_malware_hourly_chain +WITH (timescaledb.continuous, timescaledb.materialized_only = false) AS +SELECT uuid, + time_bucket(INTERVAL '1 hour', time) AS bucket, + chain, + count(chain) as count +FROM ts_malware +GROUP BY uuid, bucket, chain +WITH NO DATA; + +SELECT add_continuous_aggregate_policy('ca_ts_malware_hourly_chain', + start_offset => NULL, + end_offset => INTERVAL '30 minutes', + schedule_interval => INTERVAL '15 minutes', + if_not_exists => TRUE +); + +SELECT remove_retention_policy('ca_ts_malware_hourly_chain', if_exists => TRUE); +SELECT add_retention_policy('ca_ts_malware_hourly_chain', drop_after => INTERVAL '{{ .RetentionDays }} days', if_not_exists => TRUE); + +CREATE MATERIALIZED VIEW IF NOT EXISTS ca_ts_malware_hourly_country +WITH (timescaledb.continuous, timescaledb.materialized_only = false) AS +SELECT uuid, + time_bucket(INTERVAL '1 hour', time) AS bucket, + country, + count(country) as count +FROM ts_malware +GROUP BY uuid, bucket, country +WITH NO DATA; + +SELECT add_continuous_aggregate_policy('ca_ts_malware_hourly_country', + start_offset => NULL, + end_offset => INTERVAL '30 minutes', + schedule_interval => INTERVAL '15 minutes', + if_not_exists => TRUE +); + +SELECT remove_retention_policy('ca_ts_malware_hourly_country', if_exists => TRUE); +SELECT add_retention_policy('ca_ts_malware_hourly_country', drop_after => INTERVAL '{{ .RetentionDays }} days', if_not_exists => TRUE); + +-- ---------------------- +-- -- OVPNRW connections +-- ---------------------- + +CREATE TABLE IF NOT EXISTS ovpnrw_connections ( + time TIMESTAMPTZ NOT NULL, + uuid UUID NOT NULL, + instance TEXT NOT NULL, + common_name TEXT NOT NULL, + virtual_ip_addr TEXT NOT NULL, + remote_ip_addr TEXT NOT NULL, + start_time BIGINT NOT NULL, + duration BIGINT, + bytes_received BIGINT, + bytes_sent BIGINT, + country VARCHAR(2), + UNIQUE (time, uuid, instance, common_name) +); +CREATE INDEX IF NOT EXISTS idx_ovpnrw_connections_uuid ON ovpnrw_connections (uuid); + +SELECT + create_hypertable('ovpnrw_connections', by_range('time'), if_not_exists => TRUE); + +-- Drop raw data after 30 days +SELECT remove_retention_policy('ovpnrw_connections', if_exists => TRUE); +SELECT add_retention_policy('ovpnrw_connections', drop_after => INTERVAL '{{ .RetentionDays }} days', if_not_exists => TRUE); + +-- Continuous aggregates + +CREATE MATERIALIZED VIEW IF NOT EXISTS ca_ovpnrw_connections_hourly_count +WITH (timescaledb.continuous, timescaledb.materialized_only = false) AS +SELECT uuid, + time_bucket(INTERVAL '1 hour', time) AS bucket, + common_name, + instance, + count(common_name) as count +FROM ovpnrw_connections +GROUP BY uuid, bucket, common_name, instance +WITH NO DATA; + +SELECT add_continuous_aggregate_policy('ca_ovpnrw_connections_hourly_count', + start_offset => NULL, + end_offset => INTERVAL '30 minutes', + schedule_interval => INTERVAL '15 minutes', + if_not_exists => TRUE +); + +SELECT remove_retention_policy('ca_ovpnrw_connections_hourly_count', if_exists => TRUE); +SELECT add_retention_policy('ca_ovpnrw_connections_hourly_count', drop_after => INTERVAL '{{ .RetentionDays }} days', if_not_exists => TRUE); + +CREATE MATERIALIZED VIEW IF NOT EXISTS ca_ovpnrw_connections_hourly_bytes +WITH (timescaledb.continuous, timescaledb.materialized_only = false) AS +SELECT uuid, + time_bucket(INTERVAL '1 hour', time) AS bucket, + common_name, + instance, + sum(bytes_received) as bytes_received, + sum(bytes_sent) as bytes_sent +FROM ovpnrw_connections +GROUP BY uuid, bucket, common_name, instance, bytes_received, bytes_sent +WITH NO DATA; + +SELECT add_continuous_aggregate_policy('ca_ovpnrw_connections_hourly_bytes', + start_offset => NULL, + end_offset => INTERVAL '30 minutes', + schedule_interval => INTERVAL '15 minutes', + if_not_exists => TRUE +); + +SELECT remove_retention_policy('ca_ovpnrw_connections_hourly_bytes', if_exists => TRUE); +SELECT add_retention_policy('ca_ovpnrw_connections_hourly_bytes', drop_after => INTERVAL '{{ .RetentionDays }} days', if_not_exists => TRUE); + +---------------------- +-- TS attacks +---------------------- + +CREATE TABLE IF NOT EXISTS ts_attacks ( + time TIMESTAMPTZ NOT NULL, + uuid UUID NOT NULL, + ip TEXT NOT NULL, + country VARCHAR(2), + UNIQUE (time, uuid, ip) +); +CREATE INDEX IF NOT EXISTS idx_ts_attacks_uuid ON ts_attacks (uuid); + +SELECT + create_hypertable('ts_attacks', by_range('time'), if_not_exists => TRUE); + +-- Drop raw data after 30 days +SELECT remove_retention_policy('ts_attacks', if_exists => TRUE); +SELECT add_retention_policy('ts_attacks', drop_after => INTERVAL '{{ .RetentionDays }} days', if_not_exists => TRUE); + +-- Continuous aggregates + +CREATE MATERIALIZED VIEW IF NOT EXISTS ca_ts_attacks_hourly +WITH (timescaledb.continuous, timescaledb.materialized_only = false) AS +SELECT uuid, + time_bucket(INTERVAL '1 hour', time) AS bucket, + count(ip) as count, + country +FROM ts_attacks +GROUP BY uuid, bucket, country +WITH NO DATA; + +SELECT add_continuous_aggregate_policy('ca_ts_attacks_hourly', + start_offset => NULL, + end_offset => INTERVAL '30 minutes', + schedule_interval => INTERVAL '15 minutes', + if_not_exists => TRUE +); + +SELECT remove_retention_policy('ca_ts_attacks_hourly', if_exists => TRUE); +SELECT add_retention_policy('ca_ts_attacks_hourly', drop_after => INTERVAL '{{ .RetentionDays }} days', if_not_exists => TRUE); + +---------------------- +-- DPI stats +---------------------- + +CREATE TABLE IF NOT EXISTS dpi_stats ( + time TIMESTAMPTZ NOT NULL, + uuid UUID NOT NULL, + client_address TEXT NOT NULL, + client_name TEXT, + protocol TEXT, + host TEXT, + application TEXT, + bytes BIGINT, + UNIQUE (time, uuid, client_address, protocol, host, application) +); +-- Create index on uuid for dpi_stats only if the table is empty (clean install) +-- A production table is too large to create an index on uuid if it already contains data +DO $$ +BEGIN + IF NOT EXISTS (SELECT uuid FROM dpi_stats LIMIT 1) THEN + CREATE INDEX IF NOT EXISTS idx_dpi_stats_uuid ON dpi_stats (uuid); + END IF; +END +$$; + +SELECT + create_hypertable('dpi_stats', by_range('time'), if_not_exists => TRUE); + +-- Drop raw data after 30 days +SELECT remove_retention_policy('dpi_stats', if_exists => TRUE); +SELECT add_retention_policy('dpi_stats', drop_after => INTERVAL '{{ .RetentionDays }} days', if_not_exists => TRUE); + +-- Continuous aggregates + +DO $$ +BEGIN + -- Check if the continuous aggregate exists and doesn't have the application filter + IF EXISTS ( + SELECT 1 + FROM timescaledb_information.continuous_aggregates + WHERE view_name = 'ca_dpi_stats_hourly_bytes' + ) AND NOT EXISTS ( + SELECT 1 + FROM timescaledb_information.continuous_aggregates + WHERE view_name = 'ca_dpi_stats_hourly_bytes' + AND view_definition LIKE '%WHERE application != ''''%' + ) THEN + -- Drop the existing continuous aggregate + DROP MATERIALIZED VIEW ca_dpi_stats_hourly_bytes CASCADE; + RAISE NOTICE 'Dropped ca_dpi_stats_hourly_bytes for recreation with application filter'; + END IF; +END +$$; + +CREATE MATERIALIZED VIEW IF NOT EXISTS ca_dpi_stats_hourly_bytes +WITH (timescaledb.continuous, timescaledb.materialized_only = false) AS +SELECT uuid, + time_bucket(INTERVAL '1 hour', time) AS bucket, + sum(bytes) as bytes +FROM dpi_stats +WHERE application != '' +GROUP BY uuid, bucket +WITH NO DATA; + +SELECT add_continuous_aggregate_policy('ca_dpi_stats_hourly_bytes', + start_offset => NULL, + end_offset => INTERVAL '30 minutes', + schedule_interval => INTERVAL '15 minutes', + if_not_exists => TRUE +); + +SELECT remove_retention_policy('ca_dpi_stats_hourly_bytes', if_exists => TRUE); +SELECT add_retention_policy('ca_dpi_stats_hourly_bytes', drop_after => INTERVAL '{{ .RetentionDays }} days', if_not_exists => TRUE); + +CREATE MATERIALIZED VIEW IF NOT EXISTS ca_dpi_stats_hourly_protocol +WITH (timescaledb.continuous, timescaledb.materialized_only = false) AS +SELECT uuid, + time_bucket(INTERVAL '1 hour', time) AS bucket, + protocol, + sum(bytes) as bytes +FROM dpi_stats +WHERE protocol != '' +GROUP BY uuid, bucket, protocol +WITH NO DATA; + +SELECT add_continuous_aggregate_policy('ca_dpi_stats_hourly_protocol', + start_offset => NULL, + end_offset => INTERVAL '30 minutes', + schedule_interval => INTERVAL '15 minutes', + if_not_exists => TRUE +); + +SELECT remove_retention_policy('ca_dpi_stats_hourly_protocol', if_exists => TRUE); +SELECT add_retention_policy('ca_dpi_stats_hourly_protocol', drop_after => INTERVAL '{{ .RetentionDays }} days', if_not_exists => TRUE); + +CREATE MATERIALIZED VIEW IF NOT EXISTS ca_dpi_stats_hourly_host +WITH (timescaledb.continuous, timescaledb.materialized_only = false) AS +SELECT uuid, + time_bucket(INTERVAL '1 hour', time) AS bucket, + host, + sum(bytes) as bytes +FROM dpi_stats +WHERE host != '' +GROUP BY uuid, bucket, host +WITH NO DATA; + +SELECT add_continuous_aggregate_policy('ca_dpi_stats_hourly_host', + start_offset => NULL, + end_offset => INTERVAL '30 minutes', + schedule_interval => INTERVAL '15 minutes', + if_not_exists => TRUE +); + +SELECT remove_retention_policy('ca_dpi_stats_hourly_host', if_exists => TRUE); +SELECT add_retention_policy('ca_dpi_stats_hourly_host', drop_after => INTERVAL '{{ .RetentionDays }} days', if_not_exists => TRUE); + +CREATE MATERIALIZED VIEW IF NOT EXISTS ca_dpi_stats_hourly_application +WITH (timescaledb.continuous, timescaledb.materialized_only = false) AS +SELECT uuid, + time_bucket(INTERVAL '1 hour', time) AS bucket, + application, + sum(bytes) as bytes +FROM dpi_stats +WHERE application != '' +GROUP BY uuid, bucket, application +WITH NO DATA; + +SELECT add_continuous_aggregate_policy('ca_dpi_stats_hourly_application', + start_offset => NULL, + end_offset => INTERVAL '30 minutes', + schedule_interval => INTERVAL '15 minutes', + if_not_exists => TRUE +); + +SELECT remove_retention_policy('ca_dpi_stats_hourly_application', if_exists => TRUE); +SELECT add_retention_policy('ca_dpi_stats_hourly_application', drop_after => INTERVAL '{{ .RetentionDays }} days', if_not_exists => TRUE); + +DO $$ +BEGIN + -- Check if the continuous aggregate exists and doesn't have the client_address filter + IF EXISTS ( + SELECT 1 + FROM timescaledb_information.continuous_aggregates + WHERE view_name = 'ca_dpi_stats_hourly_client' + ) AND NOT EXISTS ( + SELECT 1 + FROM timescaledb_information.continuous_aggregates + WHERE view_name = 'ca_dpi_stats_hourly_client' + AND view_definition LIKE '%WHERE client_address != ''''%' + ) THEN + -- Drop the existing continuous aggregate + DROP MATERIALIZED VIEW ca_dpi_stats_hourly_client CASCADE; + RAISE NOTICE 'Dropped ca_dpi_stats_hourly_client for recreation with client_address filter'; + END IF; +END +$$; + +CREATE MATERIALIZED VIEW IF NOT EXISTS ca_dpi_stats_hourly_client +WITH (timescaledb.continuous, timescaledb.materialized_only = false) AS +SELECT uuid, + time_bucket(INTERVAL '1 hour', time) AS bucket, + client_address, + client_name, + sum(bytes) as bytes +FROM dpi_stats +WHERE client_address != '' AND application != '' +GROUP BY uuid, bucket, client_address, client_name +WITH NO DATA; + +SELECT add_continuous_aggregate_policy('ca_dpi_stats_hourly_client', + start_offset => NULL, + end_offset => INTERVAL '30 minutes', + schedule_interval => INTERVAL '15 minutes', + if_not_exists => TRUE +); + +SELECT remove_retention_policy('ca_dpi_stats_hourly_client', if_exists => TRUE); +SELECT add_retention_policy('ca_dpi_stats_hourly_client', drop_after => INTERVAL '{{ .RetentionDays }} days', if_not_exists => TRUE); + +-- Stored procedures -- + +-- This function cleans up orphaned data in various tables that reference units +-- It deletes records in these tables where the uuid does not exist in the units table +CREATE OR REPLACE FUNCTION cleanup_orphaned_unit_data(job_id INT DEFAULT NULL, config JSONB DEFAULT NULL) +RETURNS void AS $$ +DECLARE + cnt INTEGER; +BEGIN + -- openvpn_config + DELETE FROM openvpn_config WHERE uuid NOT IN (SELECT uuid FROM units); + GET DIAGNOSTICS cnt = ROW_COUNT; + RAISE NOTICE 'openvpn_config: % rows deleted', cnt; + + -- wan_config + DELETE FROM wan_config WHERE uuid NOT IN (SELECT uuid FROM units); + GET DIAGNOSTICS cnt = ROW_COUNT; + RAISE NOTICE 'wan_config: % rows deleted', cnt; + + -- mwan_events + DELETE FROM mwan_events WHERE uuid NOT IN (SELECT uuid FROM units); + GET DIAGNOSTICS cnt = ROW_COUNT; + RAISE NOTICE 'mwan_events: % rows deleted', cnt; + + -- ts_malware + DELETE FROM ts_malware WHERE uuid NOT IN (SELECT uuid FROM units); + GET DIAGNOSTICS cnt = ROW_COUNT; + RAISE NOTICE 'ts_malware: % rows deleted', cnt; + + -- ovpnrw_connections + DELETE FROM ovpnrw_connections WHERE uuid NOT IN (SELECT uuid FROM units); + GET DIAGNOSTICS cnt = ROW_COUNT; + RAISE NOTICE 'ovpnrw_connections: % rows deleted', cnt; + + -- ts_attacks + DELETE FROM ts_attacks WHERE uuid NOT IN (SELECT uuid FROM units); + GET DIAGNOSTICS cnt = ROW_COUNT; + RAISE NOTICE 'ts_attacks: % rows deleted', cnt; + + -- dpi_stats + DELETE FROM dpi_stats WHERE uuid NOT IN (SELECT uuid FROM units); + GET DIAGNOSTICS cnt = ROW_COUNT; + RAISE NOTICE 'dpi_stats: % rows deleted', cnt; +END; +$$ LANGUAGE plpgsql; + +-- Add a job to run the cleanup function daily, if it does not already exist +DO $$ +BEGIN + IF NOT EXISTS ( + SELECT 1 FROM timescaledb_information.jobs WHERE proc_name = 'cleanup_orphaned_unit_data' + ) THEN + PERFORM add_job('cleanup_orphaned_unit_data', '1d'); + END IF; +END +$$; \ No newline at end of file diff --git a/controller/api/storage/storage.go b/controller/api/storage/storage.go new file mode 100644 index 00000000..0ef0c200 --- /dev/null +++ b/controller/api/storage/storage.go @@ -0,0 +1,1290 @@ +/* + * Copyright (C) 2024 Nethesis S.r.l. + * http://www.nethesis.it - info@nethesis.it + * + * SPDX-License-Identifier: GPL-2.0-only + * + * author: Edoardo Spadoni + */ + +package storage + +import ( + "bytes" + "context" + "database/sql" + _ "embed" + "encoding/json" + "fmt" + "html/template" + "os" + "strconv" + "strings" + "time" + + "github.com/NethServer/nethsecurity-controller/api/configuration" + "github.com/NethServer/nethsecurity-controller/api/logs" + "github.com/NethServer/nethsecurity-controller/api/models" + "github.com/NethServer/nethsecurity-controller/api/utils" + "github.com/jackc/pgx/v5/pgxpool" + + _ "github.com/mattn/go-sqlite3" +) + +var dbpool *pgxpool.Pool +var dbctx context.Context +var err error + +//go:embed report_schema.sql.tmpl +var reportSchemaSQL string + +//go:embed upgrade_schema.sql +var upgradeSchemaSQL string + +//go:embed grafana_user.sql.tmpl +var grafanaUserSQL string + +// userUnits is a map that holds the units for each user. +var userUnits = make(map[string][]string) + +// adminUsers is a list of user names that has the admin flag +var adminUsers = make([]string, 0) + +func Init() *pgxpool.Pool { + // Initialize PostgreSQL connection and schema + dbpool, dbctx = InitReportDb() + + toBeMigratedUnits := listCCDFiles() + + if len(toBeMigratedUnits) > 0 { + // Migrate unit info from file to Postgres if needed + migrateUnitInfoFromFileToPostgres(toBeMigratedUnits) + + // Migrate users from SQLite to Postgres if needed + migrateUsersFromSqliteToPostgres(toBeMigratedUnits) + + // Migrate unit credentials from file to Postgres + migrated := migrateUnitCredentialsFromFileToPostgres() + + // Safe guard to avoid data loss + if migrated > 0 { + // Remove all units that are not inside the migrated units + cleanupUnusedUnits(toBeMigratedUnits) + } + } else { + logs.Logs.Println("[INFO][MIGRATION] skipping migration: no units found in CCD directory") + } + + ReloadACLs() + + // Initialize PostgreSQL connection + dbctx = context.Background() + dbpool, err = pgxpool.New(dbctx, configuration.Config.ReportDbUri) + if err != nil { + logs.Logs.Println("[ERR][STORAGE] error in Postgres db connection:" + err.Error()) + os.Exit(1) + } + + err = dbpool.Ping(dbctx) + if err != nil { + logs.Logs.Println("[ERR][STORAGE] error in Postgres db ping:" + err.Error()) + os.Exit(1) + } + + // Check if admin user exists + var exists bool + err = dbpool.QueryRow(dbctx, "SELECT EXISTS (SELECT 1 FROM accounts WHERE username = $1)", configuration.Config.AdminUsername).Scan(&exists) + if err != nil { + logs.Logs.Println("[ERR][STORAGE] error checking admin user: " + err.Error()) + os.Exit(1) + } + if !exists { + admin := models.Account{ + Username: configuration.Config.AdminUsername, + Password: configuration.Config.AdminPassword, + Admin: true, + DisplayName: "Administrator", + Created: time.Now(), + } + _, _ = AddAccount(admin) + } + + return dbpool +} + +func listCCDFiles() []string { + ret := make([]string, 0) + // If the dir does not exists, just return + if _, err := os.Stat(configuration.Config.OpenVPNCCDDir); os.IsNotExist(err) { + fmt.Println("[INFO][MIGRATION] OpenVPN CCD directory does not exist") + return ret + } + files, err := os.ReadDir(configuration.Config.OpenVPNCCDDir) + if err != nil { + logs.Logs.Println("[WARNING][MIGRATION] error reading OpenVPN CCD directory: " + err.Error()) + return ret + } + for _, file := range files { + if !file.IsDir() { + ret = append(ret, file.Name()) + } + } + return ret +} + +func cleanupUnusedUnits(units []string) { + pgpool, pgctx := ReportInstance() + if len(units) == 0 { + return + } + + // Remove all units and credentials that are not in the provided list + if len(units) == 0 { + return + } + // Build the list of UUIDs as a comma-separated string for the SQL IN clause + quoted := make([]string, len(units)) + for i, v := range units { + quoted[i] = fmt.Sprintf("'%s'", v) + } + inClause := strings.Join(quoted, ",") + + // Delete from unit_credentials table for units not in the list + queryCreds := fmt.Sprintf("DELETE FROM unit_credentials WHERE uuid NOT IN (%s)", inClause) + res, err := pgpool.Exec(pgctx, queryCreds) + if err != nil { + logs.Logs.Println("[ERR][MIGRATION] error deleting unused unit credentials: " + err.Error()) + } else { + logs.Logs.Printf("[INFO][MIGRATION] cleaned up %d unused unit credentials not in the list of migrated units\n", res.RowsAffected()) + } + + // Delete from units table for units not in the list + query := fmt.Sprintf("DELETE FROM units WHERE uuid NOT IN (%s)", inClause) + res, err = pgpool.Exec(pgctx, query) + if err != nil { + logs.Logs.Println("[ERR][MIGRATION] error deleting unused units: " + err.Error()) + } else { + logs.Logs.Printf("[INFO][MIGRATION] cleaned up %d unused units not in the list of migrated units\n", res.RowsAffected()) + } +} + +// migrateUsersFromSqliteToPostgres migrates users from SQLite to PostgreSQL if needed +func migrateUsersFromSqliteToPostgres(units []string) { + // 1. Check if SQLite DB exists + sqlitePath := configuration.Config.DataDir + "/db.sqlite" + if _, err := os.Stat(sqlitePath); os.IsNotExist(err) { + return // No SQLite DB, nothing to migrate + } + + // 2. Open SQLite DB + sqliteDB, err := sql.Open("sqlite3", sqlitePath) + if err != nil { + logs.Logs.Println("[INFO][MIGRATION] cannot open SQLite DB: skipping user migration") + return + } + defer sqliteDB.Close() + + // 3. Check if SQLite exists and has users + var tableExists bool + err = sqliteDB.QueryRow("SELECT EXISTS (SELECT name FROM sqlite_master WHERE type='table' AND name='accounts')").Scan(&tableExists) + if err != nil && err != sql.ErrNoRows { + logs.Logs.Println("[ERR][MIGRATION] error checking accounts table in SQLite: " + err.Error()) + return + } + if !tableExists { + logs.Logs.Println("[INFO][MIGRATION] accounts table does not exist in SQLite: skipping user migration") + return + } + rows, err := sqliteDB.Query("SELECT id, username, password, display_name, created FROM accounts") + if err != nil { + logs.Logs.Println("[ERR][MIGRATION] cannot query SQLite accounts: " + err.Error()) + return + } + defer rows.Close() + var users []models.Account + for rows.Next() { + var acc models.Account + var createdStr string + if err := rows.Scan(&acc.ID, &acc.Username, &acc.Password, &acc.DisplayName, &createdStr); err != nil { + logs.Logs.Println("[ERR][MIGRATION] error scanning SQLite user: " + err.Error()) + continue + } + acc.Created, _ = time.Parse(time.RFC3339, createdStr) + if acc.ID == 1 { + acc.Admin = true + } else { + acc.Admin = false // Default to false for other users + } + users = append(users, acc) + } + if len(users) == 0 { + return // No users to migrate + } + + // 4. Check if admin user exists in Postgres + pgpool, pgctx := ReportInstance() + var adminExists bool + err = pgpool.QueryRow(pgctx, `SELECT EXISTS (SELECT 1 FROM accounts WHERE admin = true)`).Scan(&adminExists) + if err != nil { + logs.Logs.Println("[ERR][MIGRATION] error checking admin user in Postgres: " + err.Error()) + return + } + if adminExists { + return // Admin user exists, nothing to do + } + + // 5. Create a unit_group with all units + groupID := -1 + var groupErr error + if len(units) > 0 { + group := models.UnitGroup{ + Name: "Migrated", + Description: "All units migrated from old release", + Units: units, + } + groupID, groupErr = AddUnitGroup(group) + // Create a unit group for the migrated units + if groupErr != nil { + logs.Logs.Println("[ERR][MIGRATION] error creating unit group in Postgres: " + groupErr.Error()) + } + } + + // 6. Insert users into Postgres + for _, acc := range users { + // Read OTP secret + otp_secret := "" + otp_status := "0" + otp_status_f, otp_status_err := os.ReadFile(configuration.Config.SecretsDir + "/" + acc.Username + "/status") + if otp_status_err == nil { + otp_status = string(otp_status_f[:]) + } + // Only if status is 1, load the actual secret + if otp_status == "1" { + otp_secret_f, otp_secret_err := os.ReadFile(configuration.Config.SecretsDir + "/" + acc.Username + "/secret") + if otp_secret_err == nil { + otp_secret = string(otp_secret_f[:]) + } + } + + // Read recovery codes + recoveryCodes := "" + codesB, rerr := os.ReadFile(configuration.Config.SecretsDir + "/" + acc.Username + "/codes") + if rerr == nil { + recoveryCodes = strings.ReplaceAll(strings.TrimSpace(string(codesB[:])), "\n", "|") + } + // remove acc.Username directory + os.RemoveAll(configuration.Config.SecretsDir + "/" + acc.Username) + if groupID > 0 && !acc.Admin { + acc.UnitGroups = []int{groupID} // Set unit group for non admin users + } + // Insert user into Postgres + _, accountError := AddAccount(acc) // Use AddAccount to handle password hashing and other logic + if accountError != nil { + logs.Logs.Println("[ERR][MIGRATION] error migrating user to Postgres: " + accountError.Error()) + } + // Set the password directly using a raw query + _, rawPasswordError := pgpool.Exec(pgctx, "UPDATE accounts SET password = $1 WHERE username = $2", acc.Password, acc.Username) + if rawPasswordError != nil { + logs.Logs.Println("[ERR][MIGRATION] error setting raw password for user", acc.Username, ":", rawPasswordError.Error()) + } + + // Set OTP secret + if err := SetUserOtpSecret(acc.Username, otp_secret); err != nil { + logs.Logs.Println("[ERR][MIGRATION] error mirating OTP secret for user", acc.Username, ":", err.Error()) + } + if err := SetUserRecoveryCodes(acc.Username, strings.Split(recoveryCodes, "|")); err != nil { + logs.Logs.Println("[ERR][MIGRATION] error mirating recovery codes for user", acc.Username, ":", err.Error()) + } + } + logs.Logs.Println("[INFO][MIGRATION] migrated", len(users), "users from SQLite to Postgres") + + // 7. Rename SQLite DB to avoid future migrations + err = os.Rename(sqlitePath, sqlitePath+".bak") + if err != nil { + logs.Logs.Println("[ERR][MIGRATION] error renaming SQLite DB: " + err.Error()) + } +} + +// Refactored user functions to use PostgreSQL +func AddAccount(account models.Account) (int, error) { + pgpool, pgctx := ReportInstance() + var id int + err := pgpool.QueryRow(pgctx, + "INSERT INTO accounts (username, password, admin, display_name, unit_groups, created_at) VALUES ($1, $2, $3, $4, $5, $6) RETURNING id", + account.Username, + utils.HashPassword(account.Password), + account.Admin, + account.DisplayName, + account.UnitGroups, + account.Created, + ).Scan(&id) + if err != nil { + logs.Logs.Println("[ERR][STORAGE][ADD_ACCOUNT] error in insert accounts query: " + err.Error()) + } + + ReloadACLs() + + return id, err +} + +func UpdateAccount(accountID string, account models.AccountUpdate) error { + pgpool, pgctx := ReportInstance() + var err error + if len(account.Password) > 0 { + // Update password only if it is provided + _, err = pgpool.Exec(pgctx, + `UPDATE accounts + SET password = $1 + WHERE id = $2 + `, + utils.HashPassword(account.Password), + accountID, + ) + if err != nil { + logs.Logs.Println("[ERR][STORAGE][UPDATE_ACCOUNT] error in update accounts password query: " + err.Error()) + return err + } + } + // Set unit_groups array + unitGroupsStrs := make([]string, len(account.UnitGroups)) + for i, v := range account.UnitGroups { + unitGroupsStrs[i] = strconv.Itoa(v) + } + unitGroupsArray := "{" + strings.Join(unitGroupsStrs, ",") + "}" + _, err = pgpool.Exec(pgctx, + `UPDATE accounts + SET unit_groups = $1::int[], + display_name = $2, + admin = $3, + updated_at = NOW() + WHERE id = $4`, + unitGroupsArray, + account.DisplayName, + account.Admin, + accountID, + ) + if err != nil { + logs.Logs.Println("[ERR][STORAGE][UPDATE_ACCOUNT] error in update accounts query: " + err.Error()) + return err + } + + ReloadACLs() + + return err +} + +func IsAdmin(accountUsername string) bool { + for _, admin := range adminUsers { + if admin == accountUsername { + return true + } + } + return false +} + +func GetAccounts() ([]models.Account, error) { + pgpool, pgctx := ReportInstance() + rows, err := pgpool.Query(pgctx, "SELECT id, username, display_name, admin, unit_groups, created_at, updated_at FROM accounts ORDER BY id ASC") + if err != nil { + logs.Logs.Println("[ERR][STORAGE][GET_ACCOUNTS] error in query execution:" + err.Error()) + } + defer rows.Close() + var results []models.Account + for rows.Next() { + var accountRow models.Account + if err := rows.Scan(&accountRow.ID, &accountRow.Username, &accountRow.DisplayName, &accountRow.Admin, &accountRow.UnitGroups, &accountRow.Created, &accountRow.Updated); err != nil { + logs.Logs.Println("[ERR][STORAGE][GET_ACCOUNTS] error in query row extraction" + err.Error()) + } + accountRow.TwoFA = Is2FAEnabled(accountRow.Username) + results = append(results, accountRow) + } + return results, err +} + +func GetAccount(accountID string) ([]models.Account, error) { + pgpool, pgctx := ReportInstance() + rows, err := pgpool.Query(pgctx, "SELECT id, username, display_name, admin, unit_groups, created_at, updated_at FROM accounts where id = $1", accountID) + if err != nil { + logs.Logs.Println("[ERR][STORAGE][GET_ACCOUNT] error in query execution:" + err.Error()) + } + defer rows.Close() + var results []models.Account + for rows.Next() { + var accountRow models.Account + if err := rows.Scan(&accountRow.ID, &accountRow.Username, &accountRow.DisplayName, &accountRow.Admin, &accountRow.UnitGroups, &accountRow.Created, &accountRow.Updated); err != nil { + logs.Logs.Println("[ERR][STORAGE][GET_ACCOUNT] error in query row extraction" + err.Error()) + } + results = append(results, accountRow) + } + return results, err +} + +func GetPassword(accountUsername string) string { + pgpool, pgctx := ReportInstance() + var password string + err := pgpool.QueryRow(pgctx, "SELECT password FROM accounts where username = $1 LIMIT 1", accountUsername).Scan(&password) + if err != nil { + logs.Logs.Println("[ERR][STORAGE][GET_PASSWORD] error in query execution:" + err.Error()) + } + return password +} + +func DeleteAccount(accountID string) error { + pgpool, pgctx := ReportInstance() + _, err := pgpool.Exec(pgctx, "DELETE FROM accounts where id = $1", accountID) + if err != nil { + logs.Logs.Println("[ERR][STORAGE][DELETE_ACCOUNT] error in query execution:" + err.Error()) + } + + ReloadACLs() + + return err +} + +func UpdatePassword(accountUsername string, newPassword string) error { + pgpool, pgctx := ReportInstance() + _, err := pgpool.Exec(pgctx, + "UPDATE accounts set password = $1 WHERE username = $2", + utils.HashPassword(newPassword), + accountUsername, + ) + if err == nil { + // Update the updated_at timestamp + _, err = pgpool.Exec(pgctx, "UPDATE accounts SET updated_at = NOW() WHERE username = $1", accountUsername) + if err != nil { + logs.Logs.Println("[ERR][STORAGE][UPDATE_PASSWORD] error in updating updated_at timestamp: " + err.Error()) + } + } else { + logs.Logs.Println("[ERR][STORAGE][UPDATE_PASSWORD] error during update password: " + err.Error()) + } + return err +} + +func loadReportSchema(*pgxpool.Pool, context.Context) bool { + // execute create tables + logs.Logs.Println("[INFO][STORAGE] creating report tables") + reportTemplate, _ := template.New("report_schema").Parse(reportSchemaSQL) + var executedReportTemplate bytes.Buffer + errExecute := reportTemplate.Execute(&executedReportTemplate, configuration.Config) + if errExecute != nil { + logs.Logs.Println("[ERR][STORAGE] error in storage file schema init:" + errExecute.Error()) + return false + } + _, errExecute = dbpool.Exec(dbctx, executedReportTemplate.String()) + if errExecute != nil { + logs.Logs.Println("[ERR][STORAGE] error in storage file schema init:" + errExecute.Error()) + return false + } + + logs.Logs.Println("[INFO][STORAGE] creating grafana user") + grafanaUserTemplate, _ := template.New("grafana_user").Parse(grafanaUserSQL) + var executedGrafanaUserReport bytes.Buffer + errExecute = grafanaUserTemplate.Execute(&executedGrafanaUserReport, configuration.Config) + if errExecute != nil { + logs.Logs.Println("[ERR][STORAGE] error in storage file schema init:" + errExecute.Error()) + return false + } + _, errExecute = dbpool.Exec(dbctx, executedGrafanaUserReport.String()) + if errExecute != nil { + logs.Logs.Println("[ERR][STORAGE] error in storage file schema init:" + errExecute.Error()) + return false + } + + // execute upgrade schema + logs.Logs.Println("[INFO][STORAGE] upgrading report schema") + _, errExecute = dbpool.Exec(dbctx, upgradeSchemaSQL) + if errExecute != nil { + logs.Logs.Println("[ERR][STORAGE] error in storage file schema upgrade:" + errExecute.Error()) + return false + } + return true +} + +func InitReportDb() (*pgxpool.Pool, context.Context) { + dbctx = context.Background() + dbpool, err = pgxpool.New(dbctx, configuration.Config.ReportDbUri) + if err != nil { + logs.Logs.Println("[WARN][DB] error in db connection:" + err.Error()) + } + + err = dbpool.Ping(dbctx) + if err != nil { + logs.Logs.Println("[WARN][DB] error in db connection:" + err.Error()) + } + + loadReportSchema(dbpool, dbctx) + + return dbpool, dbctx +} + +func ReportInstance() (*pgxpool.Pool, context.Context) { + if dbpool == nil { + dbpool, dbctx = InitReportDb() + + } + + return dbpool, dbctx +} + +func GetUserOtpSecret(username string) string { + pgpool, pgctx := ReportInstance() + var otp_secret string + err := pgpool.QueryRow(pgctx, "SELECT otp_secret FROM accounts where username = $1 LIMIT 1", username).Scan(&otp_secret) + if err != nil { + logs.Logs.Println("[ERR][STORAGE][GET_USER_SECRET] error in query execution:" + err.Error()) + return "" + } + decrypted, err := utils.DecryptAESGCMFromString(otp_secret, []byte(configuration.Config.EncryptionKey)) + if err != nil { + logs.Logs.Println("[ERR][STORAGE][DECRYPT_USER_SECRET] error in decryption:" + err.Error()) + return "" + } + return string(decrypted) +} + +func GetRecoveryCodes(username string) []string { + pgpool, pgctx := ReportInstance() + var otp_recovery_codes string + err := pgpool.QueryRow(pgctx, "SELECT otp_recovery_codes FROM accounts where username = $1 LIMIT 1", username).Scan(&otp_recovery_codes) + if err != nil { + logs.Logs.Println("[ERR][STORAGE][GET_RECOVERY_CODES] error in query execution:" + err.Error()) + return []string{} + } + return strings.Split(otp_recovery_codes, "|") +} + +func Is2FAEnabled(username string) bool { + pgpool, pgctx := ReportInstance() + var status sql.NullString + err := pgpool.QueryRow(pgctx, "SELECT otp_secret FROM accounts where username = $1 LIMIT 1", username).Scan(&status) + if err != nil { + logs.Logs.Println("[ERR][STORAGE][GET_2FA_STATUS] error in query execution:" + err.Error()) + } + return status.Valid && status.String != "" +} + +func SetUserOtpSecret(username string, secret string) error { + pgpool, pgctx := ReportInstance() + var otp_secret string + if len(secret) > 0 { + otp_secret, _ = utils.EncryptAESGCMToString([]byte(secret), []byte(configuration.Config.EncryptionKey)) + } else { + otp_secret = "" + } + _, err := pgpool.Exec(pgctx, "UPDATE accounts set otp_secret = $1 WHERE username = $2", otp_secret, username) + if err != nil { + logs.Logs.Println("[ERR][STORAGE][SET_USER_OTP_SECRET] error in query execution:" + err.Error()) + } + return err +} + +func SetUserRecoveryCodes(username string, codes []string) error { + pgpool, pgctx := ReportInstance() + _, err := pgpool.Exec(pgctx, "UPDATE accounts set otp_recovery_codes = $1 WHERE username = $2", strings.Join(codes, "|"), username) + if err != nil { + logs.Logs.Println("[ERR][STORAGE][SET_USER_RECOVERY_CODES] error in query execution:" + err.Error()) + } + return err +} + +func AddUnit(uuid string, ipaddr string) error { + pgpool, pgctx := ReportInstance() + // Try to insert the unit; if it already exists, return an error + _, err := pgpool.Exec(pgctx, ` + INSERT INTO units (uuid, vpn_address, created_at, updated_at) + VALUES ($1, $2, NOW(), NOW()) + ON CONFLICT (uuid) DO NOTHING + `, uuid, ipaddr) + if err != nil { + logs.Logs.Println("[ERR][STORAGE][ADD_UNIT] error in query execution:" + err.Error()) + } + return err +} + +func SetUnitAddress(uuid string, ipaddr string) error { + pgpool, pgctx := ReportInstance() + // Try to update the unit; if no rows are affected, return an error + res, err := pgpool.Exec(pgctx, ` + UPDATE units SET vpn_address = $2, updated_at = NOW() + WHERE uuid = $1 + `, uuid, ipaddr) + if err != nil { + logs.Logs.Println("[ERR][STORAGE][SET_UNIT_ADDRESS] error in query execution:" + err.Error()) + return err + } + if res.RowsAffected() == 0 { + logs.Logs.Println("[WARN][STORAGE][SET_UNIT_ADDRESS] unit with uuid " + uuid + " does not exist") + } + return err +} + +func SetUnitInfo(uuid string, info models.UnitInfo) error { + pgpool, pgctx := ReportInstance() + // Try to update the unit; if no rows are affected, return an error + res, err := pgpool.Exec(pgctx, ` + UPDATE units SET name = $2, info = $3::jsonb, updated_at = NOW() + WHERE uuid = $1 + `, uuid, info.UnitName, info) + if err != nil { + logs.Logs.Println("[ERR][STORAGE][SET_UNIT_INFO] error in query execution:" + err.Error()) + return err + } + if res.RowsAffected() == 0 { + logs.Logs.Println("[WARN][STORAGE][SET_UNIT_INFO] unit with uuid " + uuid + " does not exist") + } + return err +} + +func GetUnitInfo(uuid string) map[string]interface{} { + pgpool, pgctx := ReportInstance() + var infoStr string + err := pgpool.QueryRow(pgctx, ` + SELECT info::text FROM units WHERE uuid = $1 + `, uuid).Scan(&infoStr) + if err != nil { + // No info found for this unit, return nil + return nil + } + var info map[string]interface{} + if err := json.Unmarshal([]byte(infoStr), &info); err != nil { + logs.Logs.Println("[ERR][STORAGE][GET_UNIT_INFO] error unmarshalling info:" + err.Error()) + return nil + } + return info +} + +func loadUnitIP(unitId string) string { + unitFile, err := os.ReadFile(configuration.Config.OpenVPNCCDDir + "/" + unitId) + if err != nil { + return "" + } + + // parse ccd dir file content + parts := strings.Split(string(unitFile), "\n") + parts = strings.Split(parts[0], " ") + + return parts[1] +} + +func migrateUnitInfoFromFileToPostgres(toBeMigratedUnits []string) { + migrated := 0 + for _, uuid := range toBeMigratedUnits { + softErrors := 0 + infoFile := configuration.Config.OpenVPNStatusDir + "/" + uuid + ".info" + ccdFile := configuration.Config.OpenVPNCCDDir + "/" + uuid + proxyFile := configuration.Config.OpenVPNProxyDir + "/" + uuid + ".yaml" + // If the proxy file does not exist, skip migration for this unit: this means that the unit is in a dirty state + if _, err := os.Stat(proxyFile); os.IsNotExist(err) { + logs.Logs.Println("[INFO][MIGRATION] proxy file missing for unit", uuid, "- skipping migration (dirty state)") + if err := os.Remove(ccdFile); err != nil { + logs.Logs.Println("[INFO][MIGRATION] could not remove CCD file for unit", uuid, ":", err.Error()) + } else { + logs.Logs.Println("[INFO][MIGRATION] removed CCD file for unit", uuid) + } + continue + } + + // uuid.vpn file is not migrated because it does not persist: vpn status is restored as soon as the client re-connects + ipaddr := loadUnitIP(uuid) + + // Check if the unit already exists in Postgres + exists, err := UnitExists(uuid) + if err != nil { + logs.Logs.Println("[WARNING][MIGRATION] error checking if unit exists in Postgres:", err.Error()) + continue + } + if exists { + setUnitError := SetUnitAddress(uuid, ipaddr) + if setUnitError != nil { + logs.Logs.Println("[WARNING][MIGRATION] error setting unit address in Postgres:", uuid, setUnitError.Error()) + softErrors++ + } + } else { + addUnitErr := AddUnit(uuid, ipaddr) + if addUnitErr != nil { + logs.Logs.Println("[WARNING][MIGRATION] error adding unit to Postgres:", uuid, addUnitErr.Error()) + continue + } + } + if _, err := os.Stat(infoFile); err == nil { + // read file, parse as JSON and then set unit info + data, err := os.ReadFile(infoFile) + if err != nil { + logs.Logs.Println("[WARNING][MIGRATION] error reading file:", infoFile, err.Error()) + softErrors++ + continue + } + var info models.UnitInfo + if err := json.Unmarshal(data, &info); err != nil { + logs.Logs.Println("[WARNING][MIGRATION] error parsing JSON in file:", infoFile, err.Error()) + softErrors++ + continue + } + if err := SetUnitInfo(uuid, info); err != nil { + logs.Logs.Println("[WARNING][MIGRATION] error setting unit info for", uuid, ":", err.Error()) + softErrors++ + } + // remove the info file + if err := os.Remove(infoFile); err != nil { + logs.Logs.Println("[WARNING][MIGRATION] error removing file:", infoFile, err.Error()) + softErrors++ + } else { + logs.Logs.Println("[INFO][MIGRATION] removed file:", infoFile) + } + } + // remove ccd file + if err := os.Remove(ccdFile); err != nil { + logs.Logs.Println("[WARNING][MIGRATION] error removing file:", ccdFile, err.Error()) + } else { + logs.Logs.Println("[INFO][MIGRATION] removed file:", ccdFile) + } + softErrors = 0 + migrated++ + logs.Logs.Println("[INFO][MIGRATION] migrated unit", uuid, "with IP address", ipaddr, "Error count:", softErrors) + } + logs.Logs.Println("[INFO][MIGRATION] migrated", migrated, "units from file to Postgres") +} + +func UnitExists(uuid string) (bool, error) { + pgpool, pgctx := ReportInstance() + var exists bool + err := pgpool.QueryRow(pgctx, "SELECT EXISTS (SELECT 1 FROM units WHERE uuid = $1)", uuid).Scan(&exists) + if err != nil { + logs.Logs.Println("[ERR][STORAGE][UNIT_EXISTS] error in query execution:" + err.Error()) + return false, err + } + return exists, nil +} + +// AddUnitGroup adds a new unit group. Only admin can execute. +func AddUnitGroup(group models.UnitGroup) (int, error) { + pgpool, pgctx := ReportInstance() + var id int + unitArray := "{" + strings.Join(group.Units, ",") + "}" + err := pgpool.QueryRow(pgctx, + `INSERT INTO unit_groups (name, description, units, created_at, updated_at) VALUES ($1, $2, $3::uuid[], NOW(), NOW()) RETURNING id`, + group.Name, group.Description, unitArray).Scan(&id) + if err != nil { + logs.Logs.Println("[ERR][STORAGE][ADD_UNIT_GROUP] error in insert unit_groups query: " + err.Error()) + } + + ReloadACLs() + + return id, err +} + +func UpdateUnitGroup(groupId int, group models.UnitGroup) error { + pgpool, pgctx := ReportInstance() + unitArray := "{" + strings.Join(group.Units, ",") + "}" + res, err := pgpool.Exec(pgctx, + `UPDATE unit_groups SET name = $1, description = $2, units = $3::uuid[], updated_at = NOW() WHERE id = $4`, + group.Name, group.Description, unitArray, groupId) + if err != nil { + logs.Logs.Println("[ERR][STORAGE][EDIT_UNIT_GROUP] error in update unit_groups query: " + err.Error()) + return err + } + if res.RowsAffected() == 0 { + logs.Logs.Println("[WARN][STORAGE][EDIT_UNIT_GROUP] no unit group updated with id " + strconv.Itoa(groupId)) + return fmt.Errorf("no unit group updated with id %d", groupId) + } + + ReloadACLs() + + return nil +} + +func DeleteUnitGroup(groupID int) error { + pgpool, pgctx := ReportInstance() + res, err := pgpool.Exec(pgctx, `DELETE FROM unit_groups WHERE id = $1`, groupID) + if err != nil { + logs.Logs.Println("[ERR][STORAGE][DELETE_UNIT_GROUP] error in delete unit_groups query: " + err.Error()) + return err + } + if res.RowsAffected() == 0 { + logs.Logs.Println("[WARN][STORAGE][DELETE_UNIT_GROUP] no unit group deleted with id " + strconv.Itoa(groupID)) + return fmt.Errorf("no unit group deleted with id %d", groupID) + } + + ReloadACLs() + + return nil +} + +func ListUnitGroups() ([]models.UnitGroup, error) { + pgpool, pgctx := ReportInstance() + rows, err := pgpool.Query(pgctx, ` + SELECT + ug.id, + ug.name, + ug.description, + ug.units, + ug.created_at, + ug.updated_at, + COALESCE(array_agg(a.username) FILTER (WHERE a.username IS NOT NULL), '{}') AS accounts + FROM unit_groups ug + LEFT JOIN accounts a ON ug.id = ANY(a.unit_groups) + GROUP BY ug.id, ug.name, ug.description, ug.units, ug.created_at, ug.updated_at + ORDER BY ug.id ASC + `) + if err != nil { + logs.Logs.Println("[ERR][STORAGE][GET_UNIT_GROUPS] error in query execution:" + err.Error()) + return nil, err + } + defer rows.Close() + + var groups []models.UnitGroup + for rows.Next() { + var group models.UnitGroup + var unitsArray []string + var accountsArray []string + if err := rows.Scan(&group.ID, &group.Name, &group.Description, &unitsArray, &group.CreatedAt, &group.UpdatedAt, &accountsArray); err != nil { + logs.Logs.Println("[ERR][STORAGE][GET_UNIT_GROUPS] error in query row extraction" + err.Error()) + continue + } + group.Units = unitsArray + group.UsedBy = accountsArray + groups = append(groups, group) + } + return groups, nil +} + +func GetUnitGroup(groupID int) (models.UnitGroup, error) { + pgpool, pgctx := ReportInstance() + row := pgpool.QueryRow(pgctx, `SELECT id, name, description, units, created_at, updated_at FROM unit_groups WHERE id = $1`, groupID) + + var group models.UnitGroup + var unitsArray []string + if err := row.Scan(&group.ID, &group.Name, &group.Description, &unitsArray, &group.CreatedAt, &group.UpdatedAt); err != nil { + logs.Logs.Println("[ERR][STORAGE][GET_UNIT_GROUP] error in query row extraction" + err.Error()) + return group, err + } + group.Units = unitsArray + return group, nil +} + +func IsUnitGroupUsed(groupID int) (bool, error) { + pgpool, pgctx := ReportInstance() + var exists bool + query := ` + SELECT EXISTS ( + SELECT 1 FROM accounts + WHERE $1 = ANY(unit_groups) + ) + ` + err := pgpool.QueryRow(pgctx, query, groupID).Scan(&exists) + if err != nil { + logs.Logs.Println("[ERR][STORAGE][IS_UNIT_GROUP_USED] error in query execution:" + err.Error()) + return false, err + } + return exists, nil +} + +func UnitGroupExists(groupID int) (bool, error) { + pgpool, pgctx := ReportInstance() + var exists bool + err := pgpool.QueryRow(pgctx, "SELECT EXISTS (SELECT 1 FROM unit_groups WHERE id = $1)", groupID).Scan(&exists) + if err != nil { + logs.Logs.Println("[ERR][STORAGE][UNIT_GROUP_EXISTS] error in query execution:" + err.Error()) + return false, err + } + return exists, nil +} + +func GetUserUnits() map[string][]string { + // If userUnits is already loaded, return it + if len(userUnits) > 0 { + return userUnits + } + + // Load user units from the database + userUnitsMap, err := LoadUserUnitsMap() + if err != nil { + logs.Logs.Println("[ERR][STORAGE][GET_USER_UNITS] error loading user units: " + err.Error()) + return nil + } + + userUnits = userUnitsMap + return userUnits +} + +func LoadUserUnitsMap() (map[string][]string, error) { + pgpool, pgctx := ReportInstance() + UserUnits := make(map[string][]string) + + // Use a join to get username and units in a single query + rows, err := pgpool.Query(pgctx, ` + SELECT a.username, COALESCE(u.units, '{}') AS units + FROM accounts a + LEFT JOIN LATERAL ( + SELECT array_agg(DISTINCT group_id) AS group_ids + FROM accounts acc, unnest(acc.unit_groups) AS group_id + WHERE acc.username = a.username AND acc.unit_groups IS NOT NULL AND array_length(acc.unit_groups, 1) > 0 + ) ag ON true + LEFT JOIN LATERAL ( + SELECT array_agg(DISTINCT unit_id) AS units + FROM unit_groups ug, unnest(ug.units) AS unit_id + WHERE ag.group_ids IS NOT NULL AND ug.id = ANY(ag.group_ids) + ) u ON true + `) + if err != nil { + logs.Logs.Println("[ERR][STORAGE][GET_USER_UNITS_MAP] error in query execution:" + err.Error()) + return nil, err + } + defer rows.Close() + + for rows.Next() { + var username string + var unitsArr []string + if err := rows.Scan(&username, &unitsArr); err != nil { + logs.Logs.Println("[ERR][STORAGE][GET_USER_UNITS_MAP] error in row scan: " + err.Error()) + continue + } + // If unitsArr is nil, assign empty slice + if unitsArr == nil { + unitsArr = []string{} + } + UserUnits[username] = unitsArr + } + return UserUnits, nil +} + +func LoadAdminUsersList() ([]string, error) { + pgpool, pgctx := ReportInstance() + rows, err := pgpool.Query(pgctx, "SELECT username FROM accounts WHERE admin = true") + if err != nil { + logs.Logs.Println("[ERR][STORAGE][LOAD_ADMIN_USERS] error in query execution:" + err.Error()) + return nil, err + } + defer rows.Close() + + admins := make([]string, 0) + for rows.Next() { + var username string + if err := rows.Scan(&username); err != nil { + logs.Logs.Println("[ERR][STORAGE][LOAD_ADMIN_USERS] error in row scan: " + err.Error()) + continue + } + admins = append(admins, username) + } + return admins, nil +} + +func ReloadACLs() { + // Reload user units from the database + userUnitsMap, err := LoadUserUnitsMap() + if err != nil { + logs.Logs.Println("[ERR][STORAGE][RELOAD_USER_UNITS] error loading user units: " + err.Error()) + return + } + userUnits = userUnitsMap + logs.Logs.Println("[INFO][STORAGE][RELOAD_USER_UNITS] user units reloaded successfully") + + // Reload admin users + admins, err := LoadAdminUsersList() + if err != nil { + logs.Logs.Println("[ERR][STORAGE][RELOAD_ADMIN_USERS] error loading admin users: " + err.Error()) + return + } + adminUsers = admins + logs.Logs.Println("[INFO][STORAGE][RELOAD_ADMIN_USERS] admin users reloaded successfully") +} + +func GetFreeIP() string { + // get all ips + IPs, _ := utils.ListIPs(configuration.Config.OpenVPNNetwork, configuration.Config.OpenVPNNetmask) + // remove first ip used for tun + IPs = IPs[1:] + + pgpool, pgctx := ReportInstance() + rows, err := pgpool.Query(pgctx, "SELECT vpn_address FROM units WHERE vpn_address IS NOT NULL") + if err != nil { + logs.Logs.Println("[ERR][STORAGE][GET_FREE_IP] error in query execution:" + err.Error()) + return "" + } + defer rows.Close() + + usedIPs := make([]string, 0) + for rows.Next() { + var ip string + if err := rows.Scan(&ip); err != nil { + logs.Logs.Println("[ERR][STORAGE][GET_FREE_IP] error in row scan: " + err.Error()) + continue + } + usedIPs = append(usedIPs, ip) + } + // usedIPs now contains all used IP addresses from the units table + // loop all IPs + for _, ip := range IPs { + if !utils.Contains(ip, usedIPs) { + return ip + } + } + return "" +} + +func ListUnits() ([]map[string]interface{}, error) { + + pgpool, pgctx := ReportInstance() + rows, err := pgpool.Query(pgctx, ` + SELECT + u.uuid, + u.name, + u.vpn_address, + u.info::text, + u.vpn_connected_since, + COALESCE(array_agg(g.name) FILTER (WHERE g.id IS NOT NULL), '{}') AS groups + FROM units u + LEFT JOIN unit_groups g ON u.uuid = ANY(g.units) + GROUP BY u.uuid, u.vpn_address, u.info, u.vpn_connected_since + ORDER BY u.created_at ASC + `) + if err != nil { + logs.Logs.Println("[ERR][STORAGE][LIST_UNITS] error in query execution:" + err.Error()) + return nil, err + } + defer rows.Close() + + units := make([]map[string]interface{}, 0) + for rows.Next() { + var uuid sql.NullString + var name sql.NullString + var ipaddress sql.NullString + var infoStr sql.NullString + var connectedSince sql.NullTime + var info map[string]interface{} + var groups []string + vpn_info := make(map[string]interface{}) + unit := make(map[string]interface{}) + + if err := rows.Scan(&uuid, &name, &ipaddress, &infoStr, &connectedSince, &groups); err != nil { + logs.Logs.Println("[ERR][STORAGE][LIST_UNITS] error in row scan: " + err.Error()) + continue + } + + unit["id"] = uuid.String + unit["ipaddress"] = ipaddress.String + unit["netmask"] = configuration.Config.OpenVPNNetmask + unit["vpn"] = vpn_info + unit["groups"] = groups + + if infoStr.Valid && infoStr.String != "" { + if err := json.Unmarshal([]byte(infoStr.String), &info); err == nil { + unit["info"] = info + } else { + unit["info"] = map[string]interface{}{ + "unit_name": name.String, + } + } + } else { + unit["info"] = map[string]interface{}{ + "unit_name": name.String, + } + } + + unit["join_code"] = utils.GetJoinCode(uuid.String) + if connectedSince.Valid { + vpn_info["connected_since"] = connectedSince.Time.Unix() + } + + units = append(units, unit) + } + return units, nil +} + +func ListConnectedUnits() ([]string, error) { + pgpool, pgctx := ReportInstance() + rows, err := pgpool.Query(pgctx, "SELECT uuid FROM units WHERE vpn_connected_since IS NOT NULL") + if err != nil { + logs.Logs.Println("[INFO][STORAGE][LIST_CONNECTED_UNITS] error in query execution:" + err.Error()) + return nil, err + } + defer rows.Close() + + var uuids []string + for rows.Next() { + var uuid string + if err := rows.Scan(&uuid); err != nil { + logs.Logs.Println("[ERR][STORAGE][LIST_CONNECTED_UNITS] error in row scan: " + err.Error()) + continue + } + uuids = append(uuids, uuid) + } + return uuids, nil +} + +func UpdateUnitVpnStatus(uuid string, connectedSince int) error { + pgpool, pgctx := ReportInstance() + // Convert connectedSince (seconds since epoch) to time.Time + connectedTime := time.Unix(int64(connectedSince), 0) + _, err := pgpool.Exec(pgctx, ` + UPDATE units + SET vpn_connected_since = $1, updated_at = NOW() + WHERE uuid = $2 + `, connectedTime, uuid) + if err != nil { + logs.Logs.Println("[ERR][STORAGE][UPDATE_UNIT_VPN_STATUS] error in query execution:" + err.Error()) + return err + } + return nil +} + +func GetUnit(uuid string) (map[string]interface{}, error) { + pgpool, pgctx := ReportInstance() + row := pgpool.QueryRow(pgctx, "SELECT uuid, vpn_address, info::text, vpn_connected_since FROM units WHERE uuid = $1", uuid) + + var unit map[string]interface{} + var ipaddress sql.NullString + var infoStr sql.NullString + var connectedSince sql.NullTime + + if err := row.Scan(&uuid, &ipaddress, &infoStr, &connectedSince); err != nil { + logs.Logs.Println("[ERR][STORAGE][GET_UNIT] error in query execution:" + err.Error()) + return nil, err + } + + unit = make(map[string]interface{}) + unit["id"] = uuid + unit["ipaddress"] = ipaddress.String + unit["netmask"] = configuration.Config.OpenVPNNetmask + vpn_info := make(map[string]interface{}) + vpn_info["connected_since"] = 0 + + if infoStr.Valid && infoStr.String != "" { + var info map[string]interface{} + if err := json.Unmarshal([]byte(infoStr.String), &info); err == nil { + unit["info"] = info + } else { + unit["info"] = map[string]interface{}{} + } + } else { + unit["info"] = map[string]interface{}{} + } + + if connectedSince.Valid { + vpn_info["connected_since"] = connectedSince.Time.Unix() + } + + unit["vpn"] = vpn_info + unit["join_code"] = utils.GetJoinCode(uuid) + + return unit, nil +} + +func DeleteUnit(uuid string) error { + pgpool, pgctx := ReportInstance() + + // Delete the unit from the database + res, err := pgpool.Exec(pgctx, "DELETE FROM units WHERE uuid = $1", uuid) + if err != nil { + logs.Logs.Println("[ERR][STORAGE][DELETE_UNIT] error in query execution:" + err.Error()) + return err + } + if res.RowsAffected() == 0 { + logs.Logs.Println("[WARN][STORAGE][DELETE_UNIT] no unit deleted with uuid " + uuid) + return fmt.Errorf("no unit deleted with uuid %s", uuid) + } + + // Also delete credentials for this unit + _, err = pgpool.Exec(pgctx, "DELETE FROM unit_credentials WHERE uuid = $1", uuid) + if err != nil { + logs.Logs.Println("[ERR][STORAGE][DELETE_UNIT] error deleting unit credentials:" + err.Error()) + return err + } + + // Remove the unit from all unit_groups arrays + _, err = pgpool.Exec(pgctx, ` + UPDATE unit_groups + SET units = array_remove(units, $1) + WHERE $1 = ANY(units) + `, uuid) + if err != nil { + logs.Logs.Println("[ERR][STORAGE][DELETE_UNIT] error removing unit from unit_groups:" + err.Error()) + return err + } + ReloadACLs() + + // Delete of report data is not required: data are cleaned up by a database job + return nil +} + +func GetUnitCredentials(uuid string) (string, string, error) { + pgpool, pgctx := ReportInstance() + var username, password string + err := pgpool.QueryRow(pgctx, "SELECT username, password FROM unit_credentials WHERE uuid = $1::uuid", uuid).Scan(&username, &password) + if err != nil { + logs.Logs.Println("[ERR][STORAGE][GET_UNIT_CREDENTIALS] error in query execution:" + err.Error()) + return "", "", err + } + + decrypted, err := utils.DecryptAESGCMFromString(password, []byte(configuration.Config.EncryptionKey)) + if err != nil { + logs.Logs.Println("[ERR][STORAGE][DECRYPT_UNIT_CREDENTIALS] error in decryption:" + err.Error()) + return "", "", err + } + + return username, string(decrypted), nil +} + +func SetUnitCredentials(uuid string, username string, password string) error { + encrypted, err := utils.EncryptAESGCMToString([]byte(password), []byte(configuration.Config.EncryptionKey)) + if err != nil { + logs.Logs.Println("[ERR][STORAGE][ENCRYPT_UNIT_CREDENTIALS] error in encryption:" + err.Error()) + return err + } + + pgpool, pgctx := ReportInstance() + // Update the account with the new credentials + _, err = pgpool.Exec(pgctx, ` + INSERT INTO unit_credentials (uuid, username, password) + VALUES ($1, $2, $3) + ON CONFLICT (uuid) DO UPDATE + SET username = EXCLUDED.username, password = EXCLUDED.password + `, uuid, username, encrypted) + if err != nil { + logs.Logs.Println("[ERR][STORAGE][SET_UNIT_CREDENTIALS] error in query execution:" + err.Error()) + return err + } + return nil +} + +func migrateUnitCredentialsFromFileToPostgres() int { + migrated := 0 + files, err := os.ReadDir(configuration.Config.CredentialsDir) + if err != nil { + logs.Logs.Println("[INFO][MIGRATION] credentials directory does not exists. Skipping migration.") + return migrated + } + for _, file := range files { + if file.IsDir() { + continue + } + unitID := file.Name() + // Check if the unit already exists in Postgres + exists, _ := UnitExists(unitID) + if !exists { + logs.Logs.Println("[INFO][MIGRATION] skipping unit credentials for", unitID, "because the unit does not exist") + continue + } + credPath := configuration.Config.CredentialsDir + "/" + unitID + jsonString, errRead := os.ReadFile(credPath) + if errRead != nil { + logs.Logs.Println("[WARNING][MIGRATION] error reading credentials file:", credPath, errRead.Error()) + continue + } + var credentials models.LoginRequest + if err := json.Unmarshal(jsonString, &credentials); err != nil { + logs.Logs.Println("[WARNING][MIGRATION] error parsing credentials JSON in file:", credPath, err.Error()) + continue + } + if err := SetUnitCredentials(unitID, credentials.Username, credentials.Password); err != nil { + logs.Logs.Println("[WARNING][MIGRATION] error setting credentials for unit", unitID, ":", err.Error()) + continue + } + if err := os.Remove(credPath); err != nil { + logs.Logs.Println("[WARNING][MIGRATION] error removing credentials file:", credPath, err.Error()) + } + migrated++ + } + logs.Logs.Printf("[INFO][MIGRATION] migrated %d unit credentials from file to Postgres\n", migrated) + return migrated +} diff --git a/controller/api/storage/storage_test.go b/controller/api/storage/storage_test.go new file mode 100644 index 00000000..be7208e6 --- /dev/null +++ b/controller/api/storage/storage_test.go @@ -0,0 +1,213 @@ +/* + * Copyright (C) 2025 Nethesis S.r.l. + * http://www.nethesis.it - info@nethesis.it + * + * SPDX-License-Identifier: GPL-2.0-only + * + * author: Giacomo Sanchietti + */ + +package storage + +import ( + "fmt" + "os" + "testing" + + "github.com/NethServer/nethsecurity-controller/api/configuration" + "github.com/NethServer/nethsecurity-controller/api/logs" + "github.com/NethServer/nethsecurity-controller/api/models" + "github.com/stretchr/testify/assert" +) + +func TestMain(m *testing.M) { + logs.Init("test") + // Set required env vars for config + os.Setenv("ENCRYPTION_KEY", "12345678901234567890123456789012") + os.Setenv("GRAFANA_POSTGRES_PASSWORD", "grafana_pass") + os.Setenv("REPORT_DB_URI", "postgres://report:password@localhost:5432/report") + os.Setenv("GRAFANA_PATH", "/grafana") + os.Setenv("WEBSSH_PATH", "/webssh") + os.Setenv("PROMETHEUS_PATH", "/prometheus") + os.Setenv("PROMTAIL_PORT", "3100") + os.Setenv("PROMTAIL_ADDRESS", "localhost") + os.Setenv("ISSUER_2FA", "issuer") + os.Setenv("DATA_DIR", "/tmp/data") + os.Setenv("CREDENTIALS_DIR", "/tmp/creds") + os.Setenv("REGISTRATION_TOKEN", "token") + os.Setenv("SECRET_JWT", "secret") + os.Setenv("ADMIN_PASSWORD", "password") + os.Setenv("ADMIN_USERNAME", "admin") + os.Setenv("LISTEN_ADDRESS", "127.0.0.1:5000") + os.Setenv("SENSITIVE_LIST", "password,secret,token,passphrase,private,key") + os.Setenv("OVPN_DIR", "/etc/openvpn") + os.Setenv("SECRETS_DIR", "/tmp/secrets") + os.Setenv("FQDN", "example.com") + os.Setenv("VALID_SUBSCRIPTION", "false") + os.Setenv("PROMETHEUS_AUTH_PASSWORD", "prometheus") + os.Setenv("PROMETHEUS_AUTH_USERNAME", "prometheus") + os.Setenv("RETENTION_DAYS", "60") + os.Setenv("CACHE_TTL", "7200") + os.Setenv("OVPN_UDP_PORT", "1194") + os.Setenv("OVPN_NETMASK", "255.255.0.0") + os.Setenv("OVPN_NETWORK", "172.21.0.0") + configuration.Init() + // Assume DB is running + os.Exit(m.Run()) +} + +func TestAddUnit(t *testing.T) { + // Test adding a unit + err := AddUnit("550e8400-e29b-41d4-a716-446655440000", "192.168.1.10") + assert.NoError(t, err) + + // Verify it exists + unit, err := GetUnit("550e8400-e29b-41d4-a716-446655440000") + assert.NoError(t, err) + assert.Equal(t, "550e8400-e29b-41d4-a716-446655440000", unit["id"]) + assert.Equal(t, "192.168.1.10", unit["ipaddress"]) + + // Clean up + DeleteUnit("550e8400-e29b-41d4-a716-446655440000") +} + +func TestGetUnit(t *testing.T) { + // Add a unit first + AddUnit("550e8400-e29b-41d4-a716-446655440001", "192.168.1.11") + + // Get it + unit, err := GetUnit("550e8400-e29b-41d4-a716-446655440001") + assert.NoError(t, err) + assert.Equal(t, "550e8400-e29b-41d4-a716-446655440001", unit["id"]) + + // Clean up + DeleteUnit("550e8400-e29b-41d4-a716-446655440001") +} + +func TestGetFreeIP(t *testing.T) { + // This depends on network config, but test that it returns something + ip := GetFreeIP() + assert.NotEmpty(t, ip) +} + +func TestGetUnitCredentials(t *testing.T) { + // Set credentials + err := SetUnitCredentials("550e8400-e29b-41d4-a716-446655440002", "user", "pass") + assert.NoError(t, err) + + // Get them + user, pass, err := GetUnitCredentials("550e8400-e29b-41d4-a716-446655440002") + assert.NoError(t, err) + assert.Equal(t, "user", user) + assert.Equal(t, "pass", pass) + + // Clean up + DeleteUnit("550e8400-e29b-41d4-a716-446655440002") +} + +func TestUnitGroupExists(t *testing.T) { + // Test non-existing group + exists, err := UnitGroupExists(999) + assert.NoError(t, err) + assert.False(t, exists) +} + +func TestAddAccount(t *testing.T) { + account := models.Account{ + Username: "testuser", + Password: "testpass", + } + id, err := AddAccount(account) + assert.NoError(t, err) + assert.Greater(t, id, 0) + + // Clean up + DeleteAccount(fmt.Sprintf("%d", id)) +} + +func TestGetPassword(t *testing.T) { + // Add account + account := models.Account{ + Username: "testuser2", + Password: "testpass", + } + id, _ := AddAccount(account) + + // Get password + pass := GetPassword("testuser2") + assert.NotEmpty(t, pass) + + // Clean up + DeleteAccount(fmt.Sprintf("%d", id)) +} + +func TestIsAdmin(t *testing.T) { + // Test with non-admin + assert.False(t, IsAdmin("testuser")) + // Admin might not be loaded, so skip asserting true +} + +func TestGetAccounts(t *testing.T) { + accounts, err := GetAccounts() + assert.NoError(t, err) + // DB might be empty, so >=0 + assert.GreaterOrEqual(t, len(accounts), 0) +} + +func TestUpdatePassword(t *testing.T) { + // Add account + account := models.Account{ + Username: "testuser3", + Password: "testpass", + } + id, _ := AddAccount(account) + + // Update password + err := UpdatePassword("testuser3", "newpass") + assert.NoError(t, err) + + // Verify + pass := GetPassword("testuser3") + assert.NotEmpty(t, pass) + + // Clean up + DeleteAccount(fmt.Sprintf("%d", id)) +} + +func TestListUnits(t *testing.T) { + // Add a unit + AddUnit("550e8400-e29b-41d4-a716-446655440003", "192.168.1.12") + + // List units + units, err := ListUnits() + assert.NoError(t, err) + assert.Greater(t, len(units), 0) + + // Clean up + DeleteUnit("550e8400-e29b-41d4-a716-446655440003") +} + +func TestAddUnitGroup(t *testing.T) { + // Add unit group + group := models.UnitGroup{Name: "testgroup", Description: "", Units: []string{}} + id, err := AddUnitGroup(group) + assert.NoError(t, err) + assert.Greater(t, id, 0) + + // Clean up + DeleteUnitGroup(id) +} + +func TestGetUnitGroup(t *testing.T) { + // Add group + group := models.UnitGroup{Name: "testgroup2", Description: "", Units: []string{}} + id, _ := AddUnitGroup(group) + + // Get it + group2, err := GetUnitGroup(id) + assert.NoError(t, err) + assert.Equal(t, "testgroup2", group2.Name) + + // Clean up + DeleteUnitGroup(id) +} diff --git a/controller/api/storage/upgrade_schema.sql b/controller/api/storage/upgrade_schema.sql new file mode 100644 index 00000000..fb40c5c6 --- /dev/null +++ b/controller/api/storage/upgrade_schema.sql @@ -0,0 +1,36 @@ +/* + * Copyright (C) 2025 Nethesis S.r.l. + * http://www.nethesis.it - info@nethesis.it + * + * SPDX-License-Identifier: GPL-2.0-only + * + * author: Giacomo Sanchietti + */ + + +ALTER TABLE units ADD COLUMN IF NOT EXISTS info JSONB; +ALTER TABLE units ADD COLUMN IF NOT EXISTS updated_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP; +ALTER TABLE units ADD COLUMN IF NOT EXISTS vpn_address TEXT; +ALTER TABLE units ADD COLUMN IF NOT EXISTS vpn_connected_since TIMESTAMP NULL; +/* Remove foreign key constraints: the cascade trigger causes very slow deletes */ +ALTER TABLE openvpn_config + DROP CONSTRAINT IF EXISTS openvpn_config_uuid_fkey, + DROP CONSTRAINT IF EXISTS fk_unit; +ALTER TABLE wan_config + DROP CONSTRAINT IF EXISTS wan_config_uuid_fkey, + DROP CONSTRAINT IF EXISTS fk_unit; +ALTER TABLE mwan_events + DROP CONSTRAINT IF EXISTS mwan_events_uuid_fkey, + DROP CONSTRAINT IF EXISTS fk_unit; +ALTER TABLE ts_malware + DROP CONSTRAINT IF EXISTS ts_malware_uuid_fkey, + DROP CONSTRAINT IF EXISTS fk_unit; +ALTER TABLE ovpnrw_connections + DROP CONSTRAINT IF EXISTS ovpnrw_connections_uuid_fkey, + DROP CONSTRAINT IF EXISTS fk_unit; +ALTER TABLE ts_attacks + DROP CONSTRAINT IF EXISTS ts_attacks_uuid_fkey, + DROP CONSTRAINT IF EXISTS fk_unit; +ALTER TABLE dpi_stats + DROP CONSTRAINT IF EXISTS dpi_stats_uuid_fkey, + DROP CONSTRAINT IF EXISTS fk_unit; diff --git a/controller/api/storage_test.go b/controller/api/storage_test.go new file mode 100644 index 00000000..b02b6672 --- /dev/null +++ b/controller/api/storage_test.go @@ -0,0 +1,142 @@ +/* + * Copyright (C) 2025 Nethesis S.r.l. + * http://www.nethesis.it - info@nethesis.it + * + * SPDX-License-Identifier: GPL-2.0-only + * + * author: Giacomo Sanchietti + */ + +package main + +import ( + "testing" + + "github.com/NethServer/nethsecurity-controller/api/storage" + "github.com/google/uuid" + "github.com/stretchr/testify/assert" +) + +// TestGetFreeIP tests that GetFreeIP returns a valid available IP address. +func TestGetFreeIP(t *testing.T) { + + ip := storage.GetFreeIP() + assert.NotEmpty(t, ip, "GetFreeIP should return a non-empty IP address") + + // Verify it looks like an IP address + octets := 0 + for i := 0; i < len(ip); i++ { + if ip[i] == '.' { + octets++ + } + } + assert.Equal(t, 3, octets, "GetFreeIP should return a valid IPv4 address") +} + +// TestUnitExists tests checking if a unit exists in the database. +func TestUnitExists(t *testing.T) { + + // Create a test unit + unitID := uuid.New().String() + ip := storage.GetFreeIP() + storage.AddUnit(unitID, ip) + + // Verify unit exists + exists, err := storage.UnitExists(unitID) + assert.NoError(t, err, "UnitExists should not return an error") + assert.True(t, exists, "UnitExists should return true for existing unit") + + // Verify non-existent unit + nonExistentID := uuid.New().String() + exists, err = storage.UnitExists(nonExistentID) + assert.NoError(t, err, "UnitExists should not return an error for non-existent unit") + assert.False(t, exists, "UnitExists should return false for non-existent unit") +} + +// TestUnitCredentialsCRUD tests Create, Read, Update operations on unit credentials. +func TestUnitCredentialsCRUD(t *testing.T) { + + unitID := uuid.New().String() + ip := storage.GetFreeIP() + storage.AddUnit(unitID, ip) + + // Test SetUnitCredentials (Create) + username := "testuser" + password := "testpass123" + err := storage.SetUnitCredentials(unitID, username, password) + assert.NoError(t, err, "SetUnitCredentials should not return an error") + + // Test GetUnitCredentials (Read) + retrievedUser, retrievedPass, err := storage.GetUnitCredentials(unitID) + assert.NoError(t, err, "GetUnitCredentials should not return an error") + assert.Equal(t, username, retrievedUser, "Retrieved username should match") + assert.Equal(t, password, retrievedPass, "Retrieved password should match (after decryption)") + + // Test SetUnitCredentials (Update) + newPassword := "newpass456" + err = storage.SetUnitCredentials(unitID, username, newPassword) + assert.NoError(t, err, "SetUnitCredentials should not return an error for update") + + // Verify update + retrievedUser, retrievedPass, err = storage.GetUnitCredentials(unitID) + assert.NoError(t, err, "GetUnitCredentials should not return an error after update") + assert.Equal(t, username, retrievedUser, "Retrieved username should still match") + assert.Equal(t, newPassword, retrievedPass, "Retrieved password should be updated") +} + +// TestReloadACLs tests that ReloadACLs doesn't crash and completes successfully. +func TestReloadACLs(t *testing.T) { + + // ReloadACLs should not return an error and should complete without panic + assert.NotPanics(t, func() { + storage.ReloadACLs() + }, "ReloadACLs should not panic") +} + +// TestDatabaseConnectivity tests basic database connectivity. +func TestDatabaseConnectivity(t *testing.T) { + + // Try to get free IP which requires database connectivity + ip := storage.GetFreeIP() + + // If we get here without panic and have an IP (or empty if all used), DB is connected + assert.True(t, ip != "" || ip == "", "Database should be accessible") +} + +// TestGetFreeIPConsistency tests that GetFreeIP doesn't return duplicate IPs. +func TestGetFreeIPConsistency(t *testing.T) { + + // Get first free IP + ip1 := storage.GetFreeIP() + if ip1 == "" { + t.Skip("No free IPs available in network") + } + + // Add a unit with this IP + unitID1 := uuid.New().String() + storage.AddUnit(unitID1, ip1) + + // Get next free IP + ip2 := storage.GetFreeIP() + + // They should be different + assert.NotEqual(t, ip1, ip2, "GetFreeIP should return different IPs on successive calls") +} + +// TestUnitCredentialsEncryption tests that credentials are properly encrypted/decrypted. +func TestUnitCredentialsEncryption(t *testing.T) { + + unitID := uuid.New().String() + ip := storage.GetFreeIP() + storage.AddUnit(unitID, ip) + + // Test with sensitive password + sensitivePassword := "P@ssw0rd!#$%^&*()" + err := storage.SetUnitCredentials(unitID, "admin", sensitivePassword) + assert.NoError(t, err, "SetUnitCredentials should handle special characters") + + // Verify the password is correctly decrypted + _, retrievedPass, err := storage.GetUnitCredentials(unitID) + assert.NoError(t, err, "GetUnitCredentials should decrypt properly") + assert.Equal(t, sensitivePassword, retrievedPass, "Special characters should be preserved") +} diff --git a/controller/api/utils/geoip.go b/controller/api/utils/geoip.go new file mode 100644 index 00000000..af817a63 --- /dev/null +++ b/controller/api/utils/geoip.go @@ -0,0 +1,100 @@ +/* + * Copyright (C) 2024 Nethesis S.r.l. + * http://www.nethesis.it - info@nethesis.it + * + * SPDX-License-Identifier: GPL-2.0-only + * + * author: Giacomo Sanchietti + */ + +package utils + +import ( + "net" + "os" + "os/exec" + "strings" + "time" + + "github.com/NethServer/nethsecurity-controller/api/configuration" + "github.com/NethServer/nethsecurity-controller/api/logs" + "github.com/oschwald/geoip2-golang" +) + +var db *geoip2.Reader + +func InitGeoIP() error { + // try to download the GeoLite2-Country.mmdb file + err := DownloadGeoIpDatabase() + if err != nil { + logs.Logs.Println("[ERR][GEOIP] error downloading geoip db file") + return err + } + // open geoip db, path from config, name is always the same: GeoLite2-Country.mmdb + db, err = geoip2.Open(configuration.Config.GeoIPDbDir + "/GeoLite2-Country.mmdb") + + if err != nil { + logs.Logs.Println("[ERR][GEOIP] error reading geoip db file :" + err.Error()) + return err + } else { + logs.Logs.Println("[INFO][GEOIP] geoip db file loaded") + } + + return nil +} + +func GetCountryShort(ip string) string { + if ip == "" || db == nil { + return "" + } + + // Parse IP and get country record from GeoLite2-Country database + record, err := db.Country(net.ParseIP(ip)) + if err != nil { + logs.Logs.Println("[ERR][GEOIP] error looking up IP " + ip + ": " + err.Error()) + return "" + } + + return record.Country.IsoCode +} + +func DownloadGeoIpDatabase() error { + databaseFile, err := os.Stat(configuration.Config.GeoIPDbDir + "/GeoLite2-Country.mmdb") + if err == nil && time.Since(databaseFile.ModTime()).Hours() < 72 { + logs.Logs.Println("[INFO][GEOIP] geoip db file is up to date") + return nil + } + cmd := exec.Command( + "curl", + "-L", + "--fail", + "--silent", + "--show-error", + "--retry", "5", + "--retry-max-time", "120", + "https://download.maxmind.com/app/geoip_download?edition_id=GeoLite2-Country&license_key="+configuration.Config.MaxmindLicense+"&suffix=tar.gz", + "-o", configuration.Config.GeoIPDbDir+"/GeoLite2-Country.tar.gz", + ) + var out strings.Builder + cmd.Stderr = &out + err = cmd.Run() + if err != nil { + logs.Logs.Println("[ERR][GEOIP] error downloading geoip db file: " + out.String()) + return err + } + cmd = exec.Command( + "tar", + "xzf", + configuration.Config.GeoIPDbDir+"/GeoLite2-Country.tar.gz", + "--strip-components=1", + ) + cmd.Dir = configuration.Config.GeoIPDbDir + cmd.Stderr = &out + err = cmd.Run() + if err != nil { + logs.Logs.Println("[ERR][GEOIP] error extracting geoip db file: " + out.String()) + return err + } + + return nil +} diff --git a/controller/api/utils/utils.go b/controller/api/utils/utils.go new file mode 100644 index 00000000..b225eb67 --- /dev/null +++ b/controller/api/utils/utils.go @@ -0,0 +1,215 @@ +/* + * Copyright (C) 2024 Nethesis S.r.l. + * http://www.nethesis.it - info@nethesis.it + * + * SPDX-License-Identifier: GPL-2.0-only + * + * author: Edoardo Spadoni + */ + +package utils + +import ( + "encoding/base64" + "encoding/json" + "fmt" + "net" + "strconv" + + "crypto/aes" + "crypto/cipher" + "crypto/rand" + "io" + + "github.com/NethServer/nethsecurity-controller/api/configuration" + "github.com/gin-gonic/gin" + "golang.org/x/crypto/bcrypt" +) + +func inc(ip net.IP) { + for j := len(ip) - 1; j >= 0; j-- { + ip[j]++ + if ip[j] > 0 { + break + } + } +} + +func Contains(a string, values []string) bool { + for _, b := range values { + if b == a { + return true + } + } + return false +} + +func ListIPs(ipArg string, netmaskArg string) ([]string, error) { + // convert netmask to prefix + prefixMask, _ := net.IPMask(net.ParseIP(netmaskArg).To4()).Size() + + // create network + ip, ipnet, err := net.ParseCIDR(ipArg + "/" + strconv.Itoa(prefixMask)) + if err != nil { + return nil, err + } + + // loop all ips in network + var ips []string + for ip := ip.Mask(ipnet.Mask); ipnet.Contains(ip); inc(ip) { + ips = append(ips, ip.String()) + } + + // remove network address and broadcast address + lenIPs := len(ips) + switch { + case lenIPs < 2: + return ips, nil + + default: + return ips[1 : len(ips)-1], nil + } +} + +func HashPassword(password string) string { + bytes, _ := bcrypt.GenerateFromPassword([]byte(password), bcrypt.DefaultCost) + return string(bytes) +} + +func CheckPasswordHash(password, hash string) bool { + err := bcrypt.CompareHashAndPassword([]byte(hash), []byte(password)) + return err == nil +} + +// generate join code +// the join code is a JSON encoded in base64 with the following fields: +// - unit_id +// - registration token +// - fqdn +func GetJoinCode(unitId string) string { + // compose join code + joinCode := gin.H{ + "unit_id": unitId, + "token": configuration.Config.RegistrationToken, + "fqdn": configuration.Config.FQDN, + } + + // encode in base64 + joinCodeString, _ := json.Marshal(joinCode) + return base64.StdEncoding.EncodeToString([]byte(joinCodeString)) +} + +func Remove(a string, values []string) []string { + for i, v := range values { + if v == a { + return append(values[:i], values[i+1:]...) + } + } + return values +} + +// EncryptAESGCM encrypts plaintext using AES-GCM with the provided key. +func EncryptAESGCM(plaintext, key []byte) ([]byte, error) { + block, err := aes.NewCipher(key) + if err != nil { + return nil, err + } + gcm, err := cipher.NewGCM(block) + if err != nil { + return nil, err + } + nonce := make([]byte, gcm.NonceSize()) + if _, err := io.ReadFull(rand.Reader, nonce); err != nil { + return nil, err + } + ciphertext := gcm.Seal(nonce, nonce, plaintext, nil) + return ciphertext, nil +} + +// EncryptAESGCMToString encrypts plaintext and returns a base64 string for DB storage. +func EncryptAESGCMToString(plaintext, key []byte) (string, error) { + ciphertext, err := EncryptAESGCM(plaintext, key) + if err != nil { + return "", err + } + return base64.StdEncoding.EncodeToString(ciphertext), nil +} + +// DecryptAESGCMFromString decodes base64 string and decrypts using AES-GCM. +func DecryptAESGCMFromString(ciphertextB64 string, key []byte) ([]byte, error) { + ciphertext, err := base64.StdEncoding.DecodeString(ciphertextB64) + if err != nil { + return nil, err + } + plaintext, err := DecryptAESGCM(ciphertext, key) + if err == nil { + return plaintext, nil + } + return []byte(""), err +} + +// DecryptAESGCM decrypts ciphertext using AES-GCM with the provided key. +func DecryptAESGCM(ciphertext, key []byte) ([]byte, error) { + block, err := aes.NewCipher(key) + if err != nil { + return nil, err + } + gcm, err := cipher.NewGCM(block) + if err != nil { + return nil, err + } + if len(ciphertext) < gcm.NonceSize() { + return nil, io.ErrUnexpectedEOF + } + nonce := ciphertext[:gcm.NonceSize()] + ciphertext = ciphertext[gcm.NonceSize():] + plaintext, err := gcm.Open(nil, nonce, ciphertext, nil) + if err != nil { + return nil, err + } + return plaintext, nil +} + +// ToCIDR converts an IPv4 address and a netmask string into CIDR notation. +// It returns an empty string if the input IP or mask is invalid. +// For example, given "192.168.0.1" and "255.255.255.0", it returns "192.168.0.1/24". +func ToCIDR(ipStr, maskStr string) string { + // Parse the IP address string. + ip := net.ParseIP(ipStr) + if ip == nil { + return "" + } + // The function works with IPv4 addresses, so we ensure it's a 4-byte representation. + ipv4 := ip.To4() + if ipv4 == nil { + return "" + } + // Parse the netmask string as an IP address. + maskIP := net.ParseIP(maskStr) + if maskIP == nil { + return "" + } + // Convert the parsed netmask IP to a 4-byte representation. + maskIPv4 := maskIP.To4() + if maskIPv4 == nil { + return "" + } + // Create an IPMask type from the 4-byte mask. + mask := net.IPMask(maskIPv4) + // Get the prefix size (the number of leading '1's in the mask). + // The second return value is the total number of bits, which is always 32 for IPv4. + prefixSize, _ := mask.Size() + return fmt.Sprintf("%s/%d", ipStr, prefixSize) +} + +// ToIpMask takes an IP address in CIDR notation (e.g., "192.168.100.2/24") +// and returns the IP address and its corresponding netmask string (e.g., "255.255.255.0"). +func ToIpMask(cidr string) (string, string) { + ip, ipnet, err := net.ParseCIDR(cidr) + if err != nil { + return "", "" + } + mask := ipnet.Mask + netmask := fmt.Sprintf("%d.%d.%d.%d", mask[0], mask[1], mask[2], mask[3]) + return ip.String(), netmask +} diff --git a/controller/api/utils/utils_test.go b/controller/api/utils/utils_test.go new file mode 100644 index 00000000..0a856709 --- /dev/null +++ b/controller/api/utils/utils_test.go @@ -0,0 +1,419 @@ +/* + * Copyright (C) 2025 Nethesis S.r.l. + * http://www.nethesis.it - info@nethesis.it + * + * SPDX-License-Identifier: GPL-2.0-only + * + * author: Giacomo Sanchietti + */ + +package utils + +import ( + "bytes" + "crypto/rand" + "encoding/base64" + "encoding/json" + "os" + "testing" + + "github.com/NethServer/nethsecurity-controller/api/configuration" + "github.com/NethServer/nethsecurity-controller/api/logs" + "github.com/stretchr/testify/assert" +) // TestAESGCMEncryption tests AES-GCM encryption and decryption. +func TestAESGCMEncryption(t *testing.T) { + key := make([]byte, 32) + _, err := rand.Read(key) + assert.NoError(t, err) + + plaintext := []byte("Hello, AES-GCM encryption!") + ciphertext, err := EncryptAESGCM(plaintext, key) + assert.NoError(t, err) + assert.NotNil(t, ciphertext) + assert.False(t, bytes.Equal(ciphertext, plaintext)) + + decrypted, err := DecryptAESGCM(ciphertext, key) + assert.NoError(t, err) + assert.Equal(t, plaintext, decrypted) + + wrongKey := make([]byte, 32) + _, err = rand.Read(wrongKey) + assert.NoError(t, err) + _, err = DecryptAESGCM(ciphertext, wrongKey) + assert.Error(t, err) +} + +// TestAESGCMToString tests base64 string conversion helpers. +func TestAESGCMToString(t *testing.T) { + key := make([]byte, 32) + _, err := rand.Read(key) + assert.NoError(t, err) + + plaintext := []byte("Store this in DB as base64!") + ciphertextB64, err := EncryptAESGCMToString(plaintext, key) + assert.NoError(t, err) + assert.NotEmpty(t, ciphertextB64) + + decrypted, err := DecryptAESGCMFromString(ciphertextB64, key) + assert.NoError(t, err) + assert.Equal(t, plaintext, decrypted) + + wrongKey := []byte("abcdefghabcdefghabcdefghabcdefgh") + _, err = DecryptAESGCMFromString(ciphertextB64, wrongKey) + assert.Error(t, err) + + _, err = DecryptAESGCMFromString("not-valid-base64!!!", key) + assert.Error(t, err) +} + +// TestPasswordHashing tests password hashing security. +func TestPasswordHashing(t *testing.T) { + password := "TestPassword123!" + hash1 := HashPassword(password) + hash2 := HashPassword(password) + + assert.NotEqual(t, hash1, hash2) + assert.True(t, CheckPasswordHash(password, hash1)) + assert.True(t, CheckPasswordHash(password, hash2)) + + wrongPassword := "WrongPassword123!" + assert.False(t, CheckPasswordHash(wrongPassword, hash1)) + assert.False(t, CheckPasswordHash(wrongPassword, hash2)) +} + +// TestToCIDR tests IP/netmask to CIDR conversion. +func TestToCIDR(t *testing.T) { + testCases := []struct { + name string + ip string + mask string + want string + isValid bool + }{ + {"C/24", "192.168.1.10", "255.255.255.0", "192.168.1.10/24", true}, + {"B/16", "172.16.5.4", "255.255.0.0", "172.16.5.4/16", true}, + {"Single/32", "192.168.1.10", "255.255.255.255", "192.168.1.10/32", true}, + {"BadMask", "192.168.1.10", "255.255.0", "", false}, + {"BadIP", "notanip", "255.255.255.0", "", false}, + } + + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + got := ToCIDR(tc.ip, tc.mask) + if tc.isValid { + assert.Equal(t, tc.want, got) + } else { + assert.Equal(t, "", got) + } + }) + } +} + +// TestToIpMask tests CIDR to IP/netmask conversion. +func TestToIpMask(t *testing.T) { + testCases := []struct { + name string + cidr string + wantIP string + wantNet string + isValid bool + }{ + {"C/24", "192.168.1.10/24", "192.168.1.10", "255.255.255.0", true}, + {"B/16", "172.16.5.4/16", "172.16.5.4", "255.255.0.0", true}, + {"Single/32", "10.0.0.1/32", "10.0.0.1", "255.255.255.255", true}, + {"BadPrefix", "192.168.1.10/33", "", "", false}, + {"BadIP", "notanip/24", "", "", false}, + } + + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + ip, mask := ToIpMask(tc.cidr) + if tc.isValid { + assert.Equal(t, tc.wantIP, ip) + assert.Equal(t, tc.wantNet, mask) + } else { + assert.Equal(t, "", ip) + assert.Equal(t, "", mask) + } + }) + } +} + +// TestListIPs tests ListIPs with various network sizes. +func TestListIPs(t *testing.T) { + testCases := []struct { + name string + ip string + netmask string + minCount int + }{ + {"30", "192.168.1.0", "255.255.255.252", 2}, + {"28", "192.168.1.0", "255.255.255.240", 14}, + {"29", "10.0.0.0", "255.255.255.248", 6}, // Smaller than /24 to avoid large allocations + } + + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + ips, err := ListIPs(tc.ip, tc.netmask) + assert.NoError(t, err) + assert.GreaterOrEqual(t, len(ips), tc.minCount) + + for _, ip := range ips { + assert.NotEmpty(t, ip) + octets := 0 + for i := 0; i < len(ip); i++ { + if ip[i] == '.' { + octets++ + } + } + assert.Equal(t, 3, octets) + } + }) + } +} + +// TestListIPsEdgeCases tests edge cases for ListIPs. +func TestListIPsEdgeCases(t *testing.T) { + testCases := []struct { + name string + ip string + netmask string + want int + }{ + {"32", "192.168.1.1", "255.255.255.255", 1}, + {"31", "192.168.1.0", "255.255.255.254", 0}, + {"30", "192.168.1.0", "255.255.255.252", 2}, + } + + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + ips, err := ListIPs(tc.ip, tc.netmask) + assert.NoError(t, err) + assert.Equal(t, tc.want, len(ips)) + }) + } +} + +// TestContains tests the Contains helper. +func TestContains(t *testing.T) { + testCases := []struct { + name string + val string + slice []string + expect bool + }{ + {"Start", "apple", []string{"apple", "banana", "cherry"}, true}, + {"Middle", "banana", []string{"apple", "banana", "cherry"}, true}, + {"End", "cherry", []string{"apple", "banana", "cherry"}, true}, + {"NotFound", "grape", []string{"apple", "banana", "cherry"}, false}, + {"Empty", "apple", []string{}, false}, + {"EmptyStr", "", []string{"apple", "banana"}, false}, + {"PartialMatch", "app", []string{"apple"}, false}, + } + + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + got := Contains(tc.val, tc.slice) + assert.Equal(t, tc.expect, got) + }) + } +} + +// TestRemove tests the Remove helper. +func TestRemove(t *testing.T) { + testCases := []struct { + name string + val string + slice []string + wantLen int + }{ + {"Middle", "banana", []string{"apple", "banana", "cherry"}, 2}, + {"Start", "apple", []string{"apple", "banana", "cherry"}, 2}, + {"End", "cherry", []string{"apple", "banana", "cherry"}, 2}, + {"NotFound", "grape", []string{"apple", "banana", "cherry"}, 3}, + {"Single", "apple", []string{"apple"}, 0}, + {"Empty", "apple", []string{}, 0}, + } + + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + got := Remove(tc.val, tc.slice) + assert.Equal(t, tc.wantLen, len(got)) + }) + } +} + +// TestIncViaListIPs tests the inc() function behavior indirectly via ListIPs. +func TestIncViaListIPs(t *testing.T) { + ips, err := ListIPs("192.168.1.0", "255.255.255.252") + assert.NoError(t, err) + assert.Len(t, ips, 2) + assert.Equal(t, "192.168.1.1", ips[0]) + assert.Equal(t, "192.168.1.2", ips[1]) +} + +// TestGetJoinCode tests GetJoinCode function. +func TestGetJoinCode(t *testing.T) { + // Set up config + originalToken := configuration.Config.RegistrationToken + originalFQDN := configuration.Config.FQDN + defer func() { + configuration.Config.RegistrationToken = originalToken + configuration.Config.FQDN = originalFQDN + }() + + configuration.Config.RegistrationToken = "test-token" + configuration.Config.FQDN = "test.example.com" + + unitId := "unit-123" + joinCode := GetJoinCode(unitId) + + assert.NotEmpty(t, joinCode) + + // Decode and verify + decoded, err := base64.StdEncoding.DecodeString(joinCode) + assert.NoError(t, err) + + var data map[string]interface{} + err = json.Unmarshal(decoded, &data) + assert.NoError(t, err) + + assert.Equal(t, unitId, data["unit_id"]) + assert.Equal(t, "test-token", data["token"]) + assert.Equal(t, "test.example.com", data["fqdn"]) +} + +// TestGeoIPDownloadAndLookup tests the GeoIP2 database download and country lookup functionality. +func TestGeoIPDownloadAndLookup(t *testing.T) { + // Check if MaxMind license is set in environment + maxmindLicense := os.Getenv("MAXMIND_LICENSE") + if maxmindLicense == "" { + t.Skip("Skipping test - MAXMIND_LICENSE environment variable not set") + } + + // Create a temporary directory for the test + tmpDir, err := os.MkdirTemp("", "geoip-test-*") + assert.NoError(t, err) + defer os.RemoveAll(tmpDir) + + // Save original config and restore after test + originalGeoIPDbDir := configuration.Config.GeoIPDbDir + originalMaxmindLicense := configuration.Config.MaxmindLicense + defer func() { + configuration.Config.GeoIPDbDir = originalGeoIPDbDir + configuration.Config.MaxmindLicense = originalMaxmindLicense + }() + + // Set test configuration with the license from environment + configuration.Config.GeoIPDbDir = tmpDir + configuration.Config.MaxmindLicense = maxmindLicense + + // Initialize logs if not already done + if logs.Logs == nil { + logs.Init("test") + } + + // Test downloading the GeoIP database + err = DownloadGeoIpDatabase() + if err != nil { + t.Logf("Download failed with error: %v", err) + t.Skip("Skipping test - MaxMind download may require valid license or network access") + } + + // Verify the database file was created + dbPath := tmpDir + "/GeoLite2-Country.mmdb" + info, err := os.Stat(dbPath) + if err != nil { + t.Logf("Database file not found at %s: %v", dbPath, err) + t.Skip("Skipping test - database file was not created") + } + t.Logf("Database file created: %s (size: %d bytes)", dbPath, info.Size()) + + // Test InitGeoIP + err = InitGeoIP() + if err != nil { + t.Logf("InitGeoIP failed: %v", err) + t.Skip("Skipping test - failed to initialize GeoIP reader") + } + + // Test country lookup for known IPs + testCases := []struct { + name string + ip string + shouldResolve bool + description string + }{ + {"Google DNS", "8.8.8.8", true, "Public DNS should have country"}, + {"OpenDNS", "208.67.222.222", true, "OpenDNS should have country"}, + {"Empty IP", "", false, "Empty IP should not resolve"}, + {"Invalid IP", "999.999.999.999", false, "Invalid IP should not resolve"}, + } + + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + country := GetCountryShort(tc.ip) + if tc.shouldResolve { + assert.NotEmpty(t, country, "IP %s should resolve to a country code (%s)", tc.ip, tc.description) + assert.Len(t, country, 2, "Country code should be 2 characters (ISO 3166-1 alpha-2)") + t.Logf("IP %s resolved to country: %s", tc.ip, country) + } else { + assert.Empty(t, country, "IP %s should not resolve (%s)", tc.ip, tc.description) + } + }) + } +} + +// TestGeoIPDatabaseCaching tests that the database is not re-downloaded if it's recent. +func TestGeoIPDatabaseCaching(t *testing.T) { + // Check if MaxMind license is set in environment + maxmindLicense := os.Getenv("MAXMIND_LICENSE") + if maxmindLicense == "" { + t.Skip("Skipping test - MAXMIND_LICENSE environment variable not set") + } + + // Create a temporary directory for the test + tmpDir, err := os.MkdirTemp("", "geoip-cache-test-*") + assert.NoError(t, err) + defer os.RemoveAll(tmpDir) + + // Save original config and restore after test + originalGeoIPDbDir := configuration.Config.GeoIPDbDir + originalMaxmindLicense := configuration.Config.MaxmindLicense + defer func() { + configuration.Config.GeoIPDbDir = originalGeoIPDbDir + configuration.Config.MaxmindLicense = originalMaxmindLicense + }() + + // Set test configuration with the license from environment + configuration.Config.GeoIPDbDir = tmpDir + configuration.Config.MaxmindLicense = maxmindLicense + + // Initialize logs if not already done + if logs.Logs == nil { + logs.Init("test") + } + + // First download + err = DownloadGeoIpDatabase() + if err != nil { + t.Logf("First download failed: %v", err) + t.Skip("Skipping test - network or license issue") + } + + dbPath := tmpDir + "/GeoLite2-Country.mmdb" + firstStat, err := os.Stat(dbPath) + if err != nil { + t.Logf("Database file not found: %v", err) + t.Skip("Skipping test - database file not created") + } + + // Second download should skip because file is recent (less than 72 hours old) + err = DownloadGeoIpDatabase() + assert.NoError(t, err, "Second download should succeed (but skip actual download)") + + secondStat, err := os.Stat(dbPath) + assert.NoError(t, err) + + // Modification time should be the same (file wasn't re-downloaded) + assert.Equal(t, firstStat.ModTime(), secondStat.ModTime(), "Database file should not be re-downloaded") +} diff --git a/controller/controller.te b/controller/controller.te new file mode 100644 index 00000000..37e132c8 --- /dev/null +++ b/controller/controller.te @@ -0,0 +1,25 @@ + +module controller 1.0; + +require { + type user_tmp_t; + type pasta_t; + type container_runtime_t; + type tun_tap_device_t; + type container_t; + type unconfined_t; + class dir read; + class chr_file { read write }; + class fifo_file setattr; + class tun_socket relabelfrom; +} + +#============= container_t ============== +allow container_t container_runtime_t:fifo_file setattr; + +#!!!! This avc is allowed in the current policy +allow container_t tun_tap_device_t:chr_file { read write }; +allow container_t unconfined_t:tun_socket relabelfrom; + +#============= pasta_t ============== +allow pasta_t user_tmp_t:dir read; diff --git a/controller/dev.sh b/controller/dev.sh new file mode 100755 index 00000000..7fa6c1a9 --- /dev/null +++ b/controller/dev.sh @@ -0,0 +1,109 @@ +#!/bin/bash + +# This script manages a Podman pod for the NethSecurity project. +# It can start or stop a pod with multiple containers (VPN, API, UI, Proxy, and TimescaleDB). +# Optionally, it mounts a local directory as the UI's document root if provided: this option is useful for development purposes. + +POD="nethsecurity-pod" +branch_name=$(git -C "$(dirname "$0")" rev-parse --abbrev-ref HEAD 2>/dev/null || echo "latest") +image_tag=${IMAGE_TAG:-$branch_name} + +start_pod() { + # Check if network device tunsec exists, if not fail + if ! ip link show dev tunsec > /dev/null 2>&1; then + echo "Network device tunsec does not exist, create it using root privileges:" + echo + echo " ip tuntap add dev tunsec mod tun" + echo " ip addr add 172.21.0.1/16 dev tunsec" + echo " ip link set dev tunsec up" + exit 1 + fi + + # Stop the pod if it is already running + if podman pod exists $POD; then + echo "Pod $POD already exists" + exit 0 + fi + echo "Starting pod $POD with image tag $image_tag" + podman pod create --replace --name $POD + + # Helper function to determine image pull policy + get_image() { + local image_name="$1" + # Check if image exists locally + if podman image exists "$image_name"; then + echo "$image_name" + else + # If not local, podman run will try to pull it + echo "$image_name" + fi + } + + # The --pull option, sets pull policy to never if image exists locally to avoid registry lookup errors (required for GitHub CI) + podman run --rm --detach --pull="missing" --network=host --privileged --cap-add=NET_ADMIN --device /dev/net/tun -v ovpn-data:/etc/openvpn/:z --pod $POD --name $POD-vpn $(get_image ghcr.io/nethserver/nethsecurity-vpn:$image_tag) + # renovate: datasource=docker depName=docker.io/timescale/timescaledb + podman run --rm --detach --pull="missing" --network=host --name $POD-db --pod $POD -e POSTGRES_PASSWORD=password -e POSTGRES_USER=report docker.io/timescale/timescaledb:2.23.1-pg16 + # Wait for Postgres to be ready + echo -n "Waiting for Postgres to start..." + for i in {1..30}; do + if podman exec $POD-db pg_isready -U report > /dev/null 2>&1; then + break + fi + sleep 1 + done + # wait for db, pg_isready is not enough + sleep 5 + echo "OK" + cat > api.env < /config.yaml +entryPoints: + web: + address: "$ip:$port" + forwardedHeaders: + trustedIPs: + - "127.0.0.1/32" + +accessLog: {} + +providers: + file: + directory: $CONFIG_DIR + watch: true + +serversTransport: + insecureSkipVerify: true + +EOF + +cat << EOF > "${CONFIG_DIR}api.yaml" +http: + routers: +$(output_public_routers) + routerapi: + entryPoints: + - web + priority: 50 + middlewares: +$(output_middlewares_list stripprefix-api) + service: service-api + rule: PathPrefix(\`/api\`) + + services: + service-api: + loadBalancer: + servers: + - url: http://127.0.0.1:${api_port}/ + passHostHeader: true + + middlewares: + stripprefix-api: + stripPrefix: + prefixes: + - "/api" +$(output_whitelist_middleware) +EOF + +# separating single stripprefix into -api and -ui, shared name across files collides on merge +# ipallowlist stays duplicated in both files, gets merged anyway +cat << EOF > "${CONFIG_DIR}ui.yaml" +http: + routers: + routerui: + entryPoints: + - web + priority: 50 + middlewares: +$(output_middlewares_list stripprefix-ui) + service: service-ui + rule: PathPrefix(\`/ui\`) + routerui-root: + entryPoints: + - web + # fallback to keep backward compatibility + priority: 1 +$(output_ui_middlewares_list) + service: service-ui + rule: PathPrefix(\`/\`) + + services: + service-ui: + loadBalancer: + servers: + - url: http://127.0.0.1:${ui_port}/ + passHostHeader: true + + middlewares: + stripprefix-ui: + stripPrefix: + prefixes: + - "/ui" +$(output_whitelist_middleware) +EOF + +exec "$@" diff --git a/controller/test/smoke.sh b/controller/test/smoke.sh new file mode 100755 index 00000000..a70cb6d6 --- /dev/null +++ b/controller/test/smoke.sh @@ -0,0 +1,178 @@ +#!/bin/bash + +set -e + +# Smoke test script for NethSecurity Controller +# This script builds all containers, starts the stack, and verifies all services are running + +BASE_URL="http://localhost:5000" +POD="nethsecurity-pod" +SCRIPT_DIR=$(cd "$(dirname "$0")/.." && pwd) + +echo "=== NethSecurity Controller Smoke Test ===" +echo + +# Clean up function +cleanup() { + echo "Cleaning up..." + cd "$SCRIPT_DIR" + ./dev.sh stop 2>/dev/null || true +} + +trap cleanup EXIT + +# Check prerequisites +echo "Checking prerequisites..." +for cmd in podman buildah curl jq; do + if ! command -v $cmd &> /dev/null; then + echo "Error: $cmd is not installed" + exit 1 + fi +done + +# Check tunsec device +if ! ip link show dev tunsec > /dev/null 2>&1; then + echo "Error: tunsec network device does not exist" + echo "Create it with:" + echo " sudo ip tuntap add dev tunsec mod tun" + echo " sudo ip addr add 172.21.0.1/16 dev tunsec" + echo " sudo ip link set dev tunsec up" + exit 1 +fi + +echo "Prerequisites OK" +echo + +# build-images.sh tags local images as latest +IMAGE_TAG=latest + +# Build containers +echo "Building containers..." +cd "$SCRIPT_DIR/.." +./build-images.sh +cd "$SCRIPT_DIR" + +echo +echo "Containers built successfully" +podman images +echo + +# Stop any existing pod +echo "Stopping any existing pod..." +./dev.sh stop 2>/dev/null || true +sleep 2 + +# Start the stack with the locally built images +echo "Starting the stack with IMAGE_TAG=$IMAGE_TAG..." +IMAGE_TAG="$IMAGE_TAG" ./dev.sh start + +# Wait for API to be ready +echo "Waiting for API server to be ready..." +for i in {1..60}; do + if curl -s -f "$BASE_URL/health" > /dev/null 2>&1; then + echo "API server is ready" + break + fi + if [ $i -eq 60 ]; then + echo "Error: API server did not become ready in time" + podman logs nethsecurity-pod-api + exit 1 + fi + sleep 1 +done + +# Check all containers are running +echo +echo "Checking all containers are running..." +for container in nethsecurity-pod-vpn nethsecurity-pod-db nethsecurity-pod-api nethsecurity-pod-ui nethsecurity-pod-proxy; do + if ! podman ps --filter name=$container --format "{{.Status}}" | grep -q "Up"; then + echo "Error: Container $container is not running" + podman ps -a + exit 1 + fi + echo " ✓ $container is running" +done + +# Test login +echo +echo "Testing login..." +LOGIN_RESPONSE=$(curl -s -X POST "$BASE_URL/login" \ + -H "Content-Type: application/json" \ + -d '{"username":"admin","password":"admin"}') + +TOKEN=$(echo "$LOGIN_RESPONSE" | jq -r '.token') +if [ "$TOKEN" = "null" ] || [ -z "$TOKEN" ]; then + echo "Error: Login failed" + echo "Response: $LOGIN_RESPONSE" + exit 1 +fi +echo " ✓ Login successful" + +# Test adding a unit +echo +echo "Testing unit creation..." +UNIT_ID=$(uuidgen) +ADD_RESPONSE=$(curl -s -X POST "$BASE_URL/units" \ + -H "Authorization: Bearer $TOKEN" \ + -H "Content-Type: application/json" \ + -d "{\"unit_id\":\"$UNIT_ID\",\"subscription\":\"active\"}") + +if ! echo "$ADD_RESPONSE" | jq -e '.code == 200' > /dev/null; then + echo "Error: Failed to add unit" + echo "Response: $ADD_RESPONSE" + exit 1 +fi +echo " ✓ Unit created successfully" + +# Test retrieving the unit +echo +echo "Testing unit retrieval..." +UNIT_RESPONSE=$(curl -s -X GET "$BASE_URL/units/$UNIT_ID" \ + -H "Authorization: Bearer $TOKEN") + +if ! echo "$UNIT_RESPONSE" | jq -e '.code == 200' > /dev/null; then + echo "Error: Failed to retrieve unit" + echo "Response: $UNIT_RESPONSE" + exit 1 +fi +echo " ✓ Unit retrieved successfully" + +# Test health endpoint +echo +echo "Testing health endpoint..." +HEALTH_RESPONSE=$(curl -s "$BASE_URL/health") +if ! echo "$HEALTH_RESPONSE" | jq -e '.status == "ok"' > /dev/null 2>&1; then + echo "Error: Health check failed" + echo "Response: $HEALTH_RESPONSE" + exit 1 +fi +echo " ✓ Health check passed" + +# Check VPN certificates directory +echo +echo "Checking VPN certificates..." +if ! podman exec nethsecurity-pod-vpn test -d /etc/openvpn/pki; then + echo "Error: VPN PKI directory not found" + exit 1 +fi +echo " ✓ VPN PKI directory exists" + +# Check database connectivity +echo +echo "Testing database connectivity..." +if ! podman exec nethsecurity-pod-db pg_isready -U report > /dev/null 2>&1; then + echo "Error: Database is not ready" + exit 1 +fi +echo " ✓ Database is ready" + +echo +echo "=== All smoke tests passed! ===" +echo +echo "Stack is running and healthy:" +echo " - VPN: Running on port 1194 (UDP)" +echo " - API: http://localhost:5000" +echo " - UI: http://localhost:3000" +echo " - Proxy: http://localhost:8080" +echo +echo "Run './dev.sh stop' to stop the stack" diff --git a/controller/ui/Containerfile b/controller/ui/Containerfile new file mode 100644 index 00000000..ec95ffd0 --- /dev/null +++ b/controller/ui/Containerfile @@ -0,0 +1,16 @@ +FROM node:22.23.3 AS build +WORKDIR /build +# renovate: datasource=github-releases depName=NethServer/nethsecurity-ui +ARG UI_VERSION=2.25.1 +RUN git init . \ + && git remote add origin https://github.com/NethServer/nethsecurity-ui \ + && git fetch --depth=1 origin "$UI_VERSION" \ + && git checkout FETCH_HEAD \ + && npm ci \ + && VITE_UI_MODE=controller npm run build + +FROM docker.io/alpine:3.22.6 AS dist +RUN apk add --no-cache lighttpd +COPY entrypoint.sh /entrypoint.sh +ENTRYPOINT ["/entrypoint.sh"] +COPY --from=build /build/dist /var/www/localhost/htdocs diff --git a/controller/ui/entrypoint.sh b/controller/ui/entrypoint.sh new file mode 100755 index 00000000..b5c70c25 --- /dev/null +++ b/controller/ui/entrypoint.sh @@ -0,0 +1,30 @@ +#!/bin/sh + +ui_port=${UI_PORT:-3000} +ui_bind=${UI_BIND_IP:-0.0.0.0} + +exec 3>&1 +chown lighttpd:lighttpd /dev/fd/3 + +cat < /etc/lighttpd/lighttpd.conf +var.basedir = "/var/www/localhost" +var.logdir = "/var/log/lighttpd" +var.statedir = "/var/lib/lighttpd" +server.modules = ( + "mod_access", + "mod_accesslog" +) +server.bind = "${ui_bind}" +server.port = ${ui_port} +server.username = "lighttpd" +server.groupname = "lighttpd" +server.document-root = var.basedir + "/htdocs" +server.pid-file = "/run/lighttpd.pid" +server.indexfiles = ("index.php", "index.html", + "index.htm", "default.htm") +server.follow-symlink = "enable" +static-file.exclude-extensions = (".php", ".pl", ".cgi", ".fcgi") +url.access-deny = ("~", ".inc") +EOF + +lighttpd -D -f /etc/lighttpd/lighttpd.conf diff --git a/controller/vpn/Containerfile b/controller/vpn/Containerfile new file mode 100644 index 00000000..fd439979 --- /dev/null +++ b/controller/vpn/Containerfile @@ -0,0 +1,13 @@ +FROM docker.io/alpine:3.22.6 AS dist +RUN apk add --no-cache \ + openvpn \ + easy-rsa \ + postgresql-client +COPY ip /sbin/ip +COPY controller-auth /usr/local/bin/controller-auth +COPY handle-connection /usr/local/bin/handle-connection +COPY handle-disconnection /usr/local/bin/handle-disconnection +COPY renew-certs /usr/local/bin/renew-certs +COPY entrypoint.sh /entrypoint.sh +ENTRYPOINT ["/entrypoint.sh"] +CMD ["/usr/sbin/openvpn", "/etc/openvpn/server.conf"] diff --git a/controller/vpn/controller-auth b/controller/vpn/controller-auth new file mode 100755 index 00000000..0d7d544c --- /dev/null +++ b/controller/vpn/controller-auth @@ -0,0 +1,7 @@ +#!/bin/sh + +if [ -f /etc/openvpn/ccd/$username ]; then + exit 0 +else + exit 1 +fi diff --git a/controller/vpn/entrypoint.sh b/controller/vpn/entrypoint.sh new file mode 100755 index 00000000..1c6682d5 --- /dev/null +++ b/controller/vpn/entrypoint.sh @@ -0,0 +1,87 @@ +#!/bin/sh + +set -e + +ovpn_network=${OVPN_NETWORK:-172.21.0.0} +ovpn_netmask=${OVPN_NETMASK:-255.255.0.0} +cn=${OVPN_CN:-nethsec} +ovpn_port=${OVPN_UDP_PORT:-1194} +tun=${OVPN_TUN:-tunsec} +tun_mtu=${OVPN_TUN_MTU:-1500} +mssfix=${OVPN_MSSFIX:-1450} + +if [ ! -f /etc/openvpn/pki/ca.crt ]; then + cd /etc/openvpn + EASYRSA_BATCH=1 /usr/share/easy-rsa/easyrsa init-pki + EASYRSA_BATCH=1 EASYRSA_REQ_CN=$cn /usr/share/easy-rsa/easyrsa build-ca nopass + openssl dhparam -dsaparam -out pki/dh.pem 2048 + EASYRSA_BATCH=1 EASYRSA_CERT_EXPIRE=3650 EASYRSA_REQ_CN=$cn /usr/share/easy-rsa/easyrsa build-server-full server nopass + EASYRSA_BATCH=1 EASYRSA_CRL_DAYS=3650 EASYRSA_REQ_CN=$cn /usr/share/easy-rsa/easyrsa gen-crl + cd - +fi + +if [ ! -d /etc/openvpn/ccd ]; then + mkdir -p /etc/openvpn/ccd +fi + +if [ ! -d /etc/openvpn/run ]; then + mkdir -p /etc/openvpn/run +fi + +if [ ! -d /etc/openvpn/proxy ]; then + mkdir -p /etc/openvpn/proxy +fi + +if [ ! -d /etc/openvpn/status ]; then + mkdir -p /etc/openvpn/status +else + find /etc/openvpn/status -name "*.vpn" -delete 2>/dev/null || true +fi + +cat << EOF > /etc/openvpn/server.conf +dev $tun +dev-type tun +server $ovpn_network $ovpn_netmask +push "route $ovpn_network $ovpn_netmask" + +topology subnet +client-config-dir /etc/openvpn/ccd + +ifconfig-pool-persist host-to-net.pool 0 + +port $ovpn_port +script-security 3 +float +multihome + +tun-mtu $tun_mtu +mssfix $mssfix + +tls-server +remote-cert-tls server +dh /etc/openvpn/pki/dh.pem +ca /etc/openvpn/pki/ca.crt +cert /etc/openvpn/pki/issued/server.crt +key /etc/openvpn/pki/private/server.key +crl-verify /etc/openvpn/pki/crl.pem + +client-connect /usr/local/bin/handle-connection +client-disconnect /usr/local/bin/handle-disconnection + +# configuration for old easy RSA certs +remote-cert-ku e0 80 +remote-cert-eku "TLS Web Client Authentication" + +management /etc/openvpn/run/mgmt.sock unix + +errors-to-stderr +keepalive 20 120 +persist-key +persist-tun +verb 3 +EOF + +# renew expiring certificates before starting the server +/usr/local/bin/renew-certs || echo "[entrypoint] certificate renewal failed" + +exec "$@" diff --git a/controller/vpn/handle-connection b/controller/vpn/handle-connection new file mode 100755 index 00000000..68535b87 --- /dev/null +++ b/controller/vpn/handle-connection @@ -0,0 +1,61 @@ +#!/bin/sh + +source /etc/openvpn/conf.env + +# Send output to stdout to avoid flooding the logs +/usr/bin/psql $REPORT_DB_URI -c "UPDATE units SET vpn_connected_since = NOW() WHERE uuid = '$common_name';" > /dev/null + +# Dynamically assign VPN IP address +tmp_config=$1 +if [ -z "$tmp_config" ]; then + exit 0 +fi + +vpn_address=$(/usr/bin/psql "$REPORT_DB_URI" -t -A -c "SELECT vpn_address FROM units WHERE uuid = '$common_name';") +if [ -n "$vpn_address" ]; then + echo "ifconfig-push $vpn_address $OVPN_NETMASK" >> $tmp_config +else + vpn_address=$ifconfig_pool_remote_ip +fi + +# Add route to traefik +# using tmp to make the load atomic, moving the file at the end of the script +cat < /etc/openvpn/proxy/$common_name.yaml.tmp +http: + # Add the router + routers: + router$common_name: + entryPoints: + - web + # must outrank the catch-all serving the controller UI + priority: 100 + # addslash must come before stripprefix + middlewares: + - m$common_name-addslash + - m$common_name-stripprefix + service: service-$common_name + rule: PathPrefix(\`/$common_name\`) + + # Add the service + services: + service-$common_name: + loadBalancer: + servers: + - url: https://$vpn_address:9090/ + passHostHeader: true + + # Add middleware + middlewares: + # without the trailing slash the unit's relative asset URLs resolve against the origin root + # and hit the controller's bundle + m$common_name-addslash: + redirectRegex: + regex: "^https?://[^/]+/$common_name(\\\\?.*)?\$" + replacement: "/$common_name/\${1}" + permanent: false + m$common_name-stripprefix: + stripPrefix: + prefixes: + - "/$common_name" +EOF +mv /etc/openvpn/proxy/$common_name.yaml.tmp /etc/openvpn/proxy/$common_name.yaml \ No newline at end of file diff --git a/controller/vpn/handle-disconnection b/controller/vpn/handle-disconnection new file mode 100755 index 00000000..b55b850e --- /dev/null +++ b/controller/vpn/handle-disconnection @@ -0,0 +1,7 @@ +#!/bin/sh + +source /etc/openvpn/conf.env +# Mark the unit as disconnected +/usr/bin/psql $REPORT_DB_URI -c "UPDATE units SET vpn_connected_since = NULL WHERE uuid = '$common_name';" >/dev/null + +exit 0 diff --git a/controller/vpn/ip b/controller/vpn/ip new file mode 100755 index 00000000..296ef781 --- /dev/null +++ b/controller/vpn/ip @@ -0,0 +1,3 @@ +#!/bin/sh + +true diff --git a/controller/vpn/renew-certs b/controller/vpn/renew-certs new file mode 100755 index 00000000..d1978474 --- /dev/null +++ b/controller/vpn/renew-certs @@ -0,0 +1,68 @@ +#!/bin/sh + +# +# Renew the OpenVPN PKI before it expires. Runs at container startup. +# Certificates last 10 years and are renewed with 6 months left. +# + +set -e + +easyrsa=${EASYRSA_BIN:-/usr/share/easy-rsa/easyrsa} +pki=${EASYRSA_PKI:-/etc/openvpn/pki} +cert_expire=3650 # 10 years +renew_window=15552000 # 6 months + +export EASYRSA_BATCH=1 +export EASYRSA_PKI="$pki" + +[ -f "$pki/ca.crt" ] || exit 0 + +is_fresh() { + openssl x509 -in "$1" -checkend "$renew_window" -noout >/dev/null 2>&1 +} + +# openssl has no -checkend for CRLs. Any read or parse failure means "not +# fresh", so a broken CRL is rebuilt rather than left to lapse. +crl_is_fresh() { + [ -f "$pki/crl.pem" ] || return 1 + next=$(openssl crl -in "$pki/crl.pem" -noout -nextupdate 2>/dev/null | cut -d= -f2) + [ -n "$next" ] || return 1 + next_epoch=$(date -u -D '%b %d %H:%M:%S %Y' -d "${next% GMT}" +%s 2>/dev/null) || return 1 + [ "$next_epoch" -gt $(( $(date -u +%s) + renew_window )) ] +} + +# Drop the superseded copy instead of revoking it: easy-rsa will not renew +# while it exists, and revoking would disconnect units that have not +# re-registered yet. +renew_cert() { + rm -f "$pki/renewed/issued/$1.crt" + EASYRSA_CERT_EXPIRE="$cert_expire" "$easyrsa" renew "$1" nopass +} + +# A renewed CA keeps its key but gets a new serial, which every issued +# certificate pins. Renewing it means re-issuing all of them. +if is_fresh "$pki/ca.crt"; then + renew_ca=0 +else + renew_ca=1 + echo "[renew-certs] CA is about to expire, renewing it and the whole PKI" + EASYRSA_CA_EXPIRE="$cert_expire" "$easyrsa" renew-ca nopass +fi + +for crt in "$pki"/issued/*.crt; do + [ -e "$crt" ] || continue + cn=$(basename "$crt" .crt) + + if [ "$renew_ca" -eq 0 ] && is_fresh "$crt"; then + continue + fi + + echo "[renew-certs] renewing certificate for $cn" + renew_cert "$cn" +done + +# The CRL expires too, and OpenVPN then rejects all new connections. +if [ "$renew_ca" -eq 1 ] || ! crl_is_fresh; then + echo "[renew-certs] refreshing the certificate revocation list" + EASYRSA_CRL_DAYS="$cert_expire" "$easyrsa" gen-crl +fi diff --git a/renovate.json b/renovate.json index 33eb188d..b9771b72 100644 --- a/renovate.json +++ b/renovate.json @@ -1,19 +1,30 @@ { "$schema": "https://docs.renovatebot.com/renovate-schema.json", "extends": [ - "local>NethServer/.github:ns8-automerge" + "local>NethServer/.github:ns8-automerge", + "customManagers:dockerfileVersions" ], "customManagers": [ { "customType": "regex", + "description": "Bump the first image tag following a `renovate: datasource=... depName=...` comment", "fileMatch": [ - "build-images.sh" + "^controller/dev\\.sh$", + "^controller/README\\.md$", + "^\\.github/workflows/controller-tests\\.yml$" ], "matchStrings": [ - "controller_version=\"(?[-0-9a-zA-Z_.]+)\"" + "renovate: datasource=(?[a-z-]+) depName=(?\\S+)[\\s\\S]*?\\s[^\\s:]+/[^\\s:]+:(?[-0-9a-zA-Z_.]+)" + ] + } + ], + "packageRules": [ + { + "matchDepNames": [ + "NethServer/nethsecurity-ui" ], - "depNameTemplate": "NethServer/nethsecurity-controller", - "datasourceTemplate": "github-releases" + "matchUpdateTypes": "minor", + "automerge": false } ] }