From 7721446a309a09b03e950e00c3bcfc2bfc7de2f5 Mon Sep 17 00:00:00 2001 From: 0x7fffff92 <0x7fffff92@example.com> Date: Thu, 24 Sep 2026 21:48:01 +0800 Subject: [PATCH] feat: add shared policy API --- API.md | 338 ++++++++++++++++++++++++ go.mod | 1 + go.sum | 3 + main.go | 8 + policy.go | 678 ++++++++++++++++++++++++++++++++++++++++++++++++ service_auth.go | 173 ++++++++++++ 6 files changed, 1201 insertions(+) create mode 100644 API.md create mode 100644 policy.go create mode 100644 service_auth.go diff --git a/API.md b/API.md new file mode 100644 index 0000000..846a152 --- /dev/null +++ b/API.md @@ -0,0 +1,338 @@ +# Headscale API Wrapper Specification + +This document describes the APIs exposed by `headscale-api-wrapper` to other Olares components. Here, "external" means callers outside the wrapper process. These APIs must remain internal to the cluster through Kubernetes Services and NetworkPolicies and must not be exposed directly to the public internet. + +## 1. API categories + +| Category | Caller | Port | Path prefix | Credential | +| --- | --- | --- | --- | --- | +| Auth key | Vault/LarePass call chain | `9000` | `/headscale` | Olares AccessToken | +| User device management | Settings and user-service | `8000` | `/headscale` | Olares AccessToken | +| Platform policy management | app-service | `8000` | `/internal/policy` | Kubernetes ServiceAccount token | + +Recommended in-cluster addresses: + +- Auth key: `http://headscale-authkey-svc.os-network:9000` +- User device and platform policy APIs: `http://headscale-server-svc.os-network:8000` + +## 2. Common response format + +All wrapper APIs use the following response envelope: + +```json +{ + "code": 0, + "message": "", + "data": {} +} +``` + +- `code = 0`: success. +- `code = 1001`: request, authentication, or Headscale call failure. +- Callers must check both the HTTP status and `code`. + +Common HTTP status codes: + +- `200`: success. +- `400`: invalid request body or port format. +- `401`: missing or invalid AccessToken or ServiceAccount token. +- `403`: the node does not belong to the current user, or the ServiceAccount is not authorized. +- `409`: the current policy structure cannot be modified safely. +- `500`: the wrapper failed to read or parse state. +- `502`: a Headscale API call failed. + +## 3. Olares user authentication + +Auth key and user device management requests must include: + +```http +X-Authorization: Bearer +``` + +The wrapper sends the AccessToken to the LLDAP token verification endpoint and uses the returned `username` as the Headscale username. + +`X-BFL-USER` is used only for diagnostics. It does not select the Headscale user and cannot override the username in the AccessToken. + +## 4. Auth key API + +### 4.1 Get a pre-auth key + +```http +GET /headscale/preauthkey +``` + +Port: `9000` + +Behavior: + +1. Verify the Olares AccessToken. +2. Find the Headscale user with the same name, creating it if it does not exist. +3. Create a reusable, non-ephemeral pre-auth key that is valid for 24 hours. + +Request body: none. + +On success, `data` retains the response structure from the Headscale create pre-auth key API. For example: + +```json +{ + "code": 0, + "message": "", + "data": { + "preAuthKey": { + "id": "21", + "key": "hskey-auth-...", + "reusable": true, + "ephemeral": false, + "expiration": "2026-09-25T10:00:00Z", + "user": { + "name": "alice" + } + } + } +} +``` + +## 5. User device management APIs + +These APIs run on port `8000` and may operate only on nodes owned by the Headscale user identified by the AccessToken. + +### 5.1 List nodes owned by the current user + +```http +POST /headscale/node +Content-Type: application/json +``` + +Request body: + +```json +{} +``` + +The wrapper fetches all nodes from Headscale and returns only nodes whose `node.user.name` matches the authenticated username. + +### 5.2 Delete a node owned by the current user + +```http +POST /headscale/node +Content-Type: application/json +``` + +Request body: + +```json +{ + "id": "12" +} +``` + +When `id` is present, this endpoint deletes the node instead of listing nodes. The wrapper verifies ownership first and returns `403` for a node owned by another user. + +### 5.3 Rename a node owned by the current user + +```http +POST /headscale/node/rename +Content-Type: application/json +``` + +Request body: + +```json +{ + "id": "12", + "name": "macbook-pro" +} +``` + +Both `id` and `name` are required. The wrapper verifies ownership before applying the change. + +### 5.4 Approve routes for a node owned by the current user + +```http +POST /headscale/node/approve_routes +Content-Type: application/json +``` + +Request body: + +```json +{ + "id": "12", + "routes": ["192.168.1.0/24"] +} +``` + +- `id` is required. +- `routes` is the complete list of approved routes. +- `routes: []` clears all approved routes for the node. +- The wrapper verifies ownership before applying the change. + +The user-facing APIs do not support transferring nodes between users or assigning arbitrary tags. This prevents users from bypassing shared-Headscale ACL isolation by changing node ownership or tags. + +## 6. Platform policy management APIs + +These APIs are intended only for app-service and run on port `8000`. + +Requests must include the app-service Kubernetes ServiceAccount token: + +```http +Authorization: Bearer +``` + +The wrapper verifies the token through Kubernetes TokenReview and requires this identity by default: + +```text +system:serviceaccount:os-framework:os-internal +``` + +The caller must use a projected ServiceAccount token with the dedicated `headscale-policy` audience. The wrapper sends the same audience in the TokenReview request, so the normal Kubernetes API token is rejected. + +The expected namespace and ServiceAccount can be overridden with: + +- `POLICY_CLIENT_NAMESPACE` +- `POLICY_CLIENT_SERVICE_ACCOUNT` +- `POLICY_TOKEN_AUDIENCE` + +### 6.1 Read the application port policy + +```http +GET /internal/policy/application-ports +``` + +Example response: + +```json +{ + "code": 0, + "message": "", + "data": { + "defaultPorts": { + "tcp": ["53", "80", "443", "18088"], + "udp": ["53"] + }, + "applicationPorts": [ + { + "user": "alice", + "tcp": ["445", "5000"], + "udp": ["5353"] + } + ], + "effectivePorts": [ + { + "user": "alice", + "tcp": ["53", "80", "443", "5000", "18088"], + "udp": ["53", "5353"] + } + ], + "revision": "sha256:...", + "updatedAt": "2026-09-24T10:00:00Z", + "inSync": true + } +} +``` + +Field definitions: + +- `defaultPorts`: platform ports available to every Headscale member. +- `applicationPorts`: dynamic ports declared by Applications and aggregated by Olares user. +- `effectivePorts`: the union of default and dynamic ports for users with dynamic entries. Users not present in this array still receive `defaultPorts`. +- `revision`: a digest of normalized `applicationPorts`. +- `updatedAt`: the last Headscale policy update time. +- `inSync`: whether the default and dynamic rules in the database match the wrapper's canonical structure. + +### 6.2 Replace the application port policy + +```http +PUT /internal/policy/application-ports +Content-Type: application/json +``` + +Request body: + +```json +{ + "applicationPorts": [ + { + "user": "alice", + "tcp": ["445", "5000", "6000-6010"], + "udp": ["5353"] + }, + { + "user": "bob", + "tcp": ["8080"], + "udp": [] + } + ] +} +``` + +Update semantics: + +- This is a full replacement, not an incremental patch. +- Dynamic application ports for users omitted from the request are removed. +- `applicationPorts: []` removes all dynamic application ports while preserving platform default ports. +- Duplicate entries for the same user are merged and deduplicated. +- Platform default ports are removed from dynamic entries even if the request includes them. +- The wrapper ensures that a matching Headscale user exists before writing the policy. +- If the content is unchanged and the existing rules are canonical, the wrapper does not write to Headscale again. + +Supported port formats: + +- Single port: `"443"` +- Inclusive range: `"6000-6010"` +- Every numeric port must be within `1..65535`. + +The success response has the same structure as GET and also includes: + +```json +{ + "changed": true +} +``` + +`changed` indicates whether this request actually updated the Headscale policy. + +## 7. Platform default ports + +The wrapper currently manages these default ports: + +```text +TCP: 53, 80, 443, 18088 +UDP: 53 +``` + +They correspond to this Headscale policy rule: + +```json +{ + "action": "accept", + "src": ["autogroup:member"], + "proto": "tcp", + "dst": ["tag:olares:53,80,443,18088"] +} +``` + +Dynamic application rules are generated per user. For example: + +```json +{ + "action": "accept", + "src": ["alice@"], + "proto": "tcp", + "dst": ["tag:olares:445,5000"] +} +``` + +During a read-modify-write operation, the wrapper preserves unrelated ACLs, groups, tag owners, auto-approvers, and policy fields it does not recognize. + +## 8. Removed and unsupported APIs + +The shared Headscale architecture no longer exposes these legacy capabilities: + +- `/inner/*` forwarding endpoints. +- Reading the Headscale control URL. +- Registering nodes through the management API. +- Transferring a node to another Headscale user. +- Assigning arbitrary tags to a node through a user-facing API. + +New callers must not depend on these legacy APIs. diff --git a/go.mod b/go.mod index 7a5e968..0a89bb5 100644 --- a/go.mod +++ b/go.mod @@ -10,6 +10,7 @@ require ( github.com/pkg/errors v0.9.1 github.com/sirupsen/logrus v1.9.3 github.com/spf13/pflag v1.0.5 + github.com/tailscale/hujson v0.0.0-20221223112325-20486734a56a golang.org/x/crypto v0.9.0 gopkg.in/yaml.v2 v2.4.0 ) diff --git a/go.sum b/go.sum index 6ca3107..e531ab7 100644 --- a/go.sum +++ b/go.sum @@ -30,6 +30,7 @@ github.com/goccy/go-json v0.10.2/go.mod h1:6MelG93GURQebXPDq3khkgXZkazVtN9CRI+MG github.com/golang/protobuf v1.5.0/go.mod h1:FsONVRAS9T7sI+LIUmWTfcYkHO4aIWwzhcaSAoJOfIk= github.com/google/go-cmp v0.5.5 h1:Khx7svrCpmxxtHBq5j2mp/xVjsi8hQMfNLvJFAlrGgU= github.com/google/go-cmp v0.5.5/go.mod h1:v8dTdLbMG2kIc/vJvl+f65V22dbkXbowE6jgT/gNBxE= +github.com/google/go-cmp v0.5.8 h1:e6P7q2lk1O+qJJb4BtCQXlK8vWEO8V1ZeuEdJNOqZyg= github.com/google/gofuzz v1.0.0/go.mod h1:dBl0BpW6vV/+mYPU4Po3pmUjxk6FQPldtuIdl/M65Eg= github.com/json-iterator/go v1.1.12 h1:PV8peI4a0ysnczrg+LtxykD8LfKY9ML6u2jnxaEnrnM= github.com/json-iterator/go v1.1.12/go.mod h1:e30LSqwooZae/UwlEbR2852Gd8hjQvJoHmT4TnhNGBo= @@ -74,6 +75,8 @@ github.com/stretchr/testify v1.8.1/go.mod h1:w2LPCIKwWwSfY2zedu0+kehJoqGctiVI29o github.com/stretchr/testify v1.8.2/go.mod h1:w2LPCIKwWwSfY2zedu0+kehJoqGctiVI29o6fzry7u4= github.com/stretchr/testify v1.8.3 h1:RP3t2pwF7cMEbC1dqtB6poj3niw/9gnV4Cjg5oW5gtY= github.com/stretchr/testify v1.8.3/go.mod h1:sz/lmYIOXD/1dqDmKjjqLyZ2RngseejIcXlSw2iwfAo= +github.com/tailscale/hujson v0.0.0-20221223112325-20486734a56a h1:SJy1Pu0eH1C29XwJucQo73FrleVK6t4kYz4NVhp34Yw= +github.com/tailscale/hujson v0.0.0-20221223112325-20486734a56a/go.mod h1:DFSS3NAGHthKo1gTlmEcSBiZrRJXi28rLNd/1udP1c8= 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.2.11 h1:BMaWp1Bb6fHwEtbplGBGJ498wD+LKlNSl25MjdZY4dU= diff --git a/main.go b/main.go index 008f52e..ca4ccef 100644 --- a/main.go +++ b/main.go @@ -8,6 +8,7 @@ import ( "os" "strconv" "strings" + "sync" "time" "github.com/gin-gonic/gin" @@ -66,6 +67,8 @@ var config string var headers map[string]string var proxyPrefix string = "/headscale" +var policyUpdateMu sync.Mutex + const authenticatedUserContextKey = "authenticated-user" func init() { @@ -176,6 +179,11 @@ func main() { router := gin.Default() router.SetTrustedProxies(nil) + internal := router.Group("/internal") + internal.Use(requireServiceAccount()) + internal.GET("/policy/application-ports", getApplicationPorts) + internal.PUT("/policy/application-ports", putApplicationPorts) + rgProxy := router.Group(proxyPrefix) rgProxy.Use(requireAuthenticatedUser()) rgProxy.POST(getMachineStr, func(c *gin.Context) { diff --git a/policy.go b/policy.go new file mode 100644 index 0000000..8d0aa20 --- /dev/null +++ b/policy.go @@ -0,0 +1,678 @@ +package main + +import ( + "crypto/sha256" + "encoding/hex" + "encoding/json" + "errors" + "fmt" + "net/http" + "sort" + "strconv" + "strings" + "time" + "unicode" + + "github.com/gin-gonic/gin" + resty "github.com/go-resty/resty/v2" + "github.com/tailscale/hujson" +) + +var defaultApplicationServicePorts = protocolPorts{ + TCP: []string{"53", "80", "443", "18088"}, + UDP: []string{"53"}, +} + +type protocolPorts struct { + TCP []string `json:"tcp"` + UDP []string `json:"udp"` +} + +type userProtocolPorts struct { + User string `json:"user"` + protocolPorts +} + +type applicationPortsRequest struct { + ApplicationPorts *[]userProtocolPorts `json:"applicationPorts"` +} + +type applicationPortsState struct { + DefaultPorts protocolPorts `json:"defaultPorts"` + ApplicationPorts []userProtocolPorts `json:"applicationPorts"` + EffectivePorts []userProtocolPorts `json:"effectivePorts"` + Revision string `json:"revision"` + UpdatedAt string `json:"updatedAt,omitempty"` + InSync bool `json:"inSync"` + Changed bool `json:"changed,omitempty"` +} + +type headscalePolicyResponse struct { + Policy string `json:"policy"` + UpdatedAt string `json:"updatedAt"` +} + +// policyACL retains unknown fields so upgrading Headscale cannot cause a +// read-modify-write to silently strip fields owned by a newer policy schema. +type policyACL struct { + fields map[string]json.RawMessage + Action string + Src []string + Proto string + Dst []string +} + +func (acl *policyACL) UnmarshalJSON(data []byte) error { + fields := make(map[string]json.RawMessage) + if err := json.Unmarshal(data, &fields); err != nil { + return err + } + type knownACL struct { + Action string `json:"action"` + Src []string `json:"src"` + Proto string `json:"proto,omitempty"` + Dst []string `json:"dst"` + } + var known knownACL + if err := json.Unmarshal(data, &known); err != nil { + return err + } + acl.fields = fields + acl.Action = known.Action + acl.Src = known.Src + acl.Proto = known.Proto + acl.Dst = known.Dst + return nil +} + +func (acl policyACL) MarshalJSON() ([]byte, error) { + fields := make(map[string]json.RawMessage, len(acl.fields)) + for key, value := range acl.fields { + fields[key] = value + } + set := func(key string, value interface{}) error { + encoded, err := json.Marshal(value) + if err != nil { + return err + } + fields[key] = encoded + return nil + } + if err := set("action", acl.Action); err != nil { + return nil, err + } + if err := set("src", acl.Src); err != nil { + return nil, err + } + if acl.Proto == "" { + delete(fields, "proto") + } else if err := set("proto", acl.Proto); err != nil { + return nil, err + } + if err := set("dst", acl.Dst); err != nil { + return nil, err + } + return json.Marshal(fields) +} + +type policyDocument struct { + fields map[string]json.RawMessage + ACLs []policyACL +} + +func getApplicationPorts(c *gin.Context) { + policyUpdateMu.Lock() + defer policyUpdateMu.Unlock() + + state, err := readApplicationPorts() + if err != nil { + c.JSON(http.StatusInternalServerError, response{Code: requestHeadscaleError, Message: err.Error()}) + return + } + c.JSON(http.StatusOK, response{Code: 0, Message: "", Data: state}) +} + +func putApplicationPorts(c *gin.Context) { + var request applicationPortsRequest + if err := c.ShouldBindJSON(&request); err != nil { + c.JSON(http.StatusBadRequest, response{Code: requestHeadscaleError, Message: err.Error()}) + return + } + if request.ApplicationPorts == nil { + c.JSON(http.StatusBadRequest, response{Code: requestHeadscaleError, Message: "applicationPorts must be an array"}) + return + } + requested, err := normalizeUserProtocolPorts(*request.ApplicationPorts) + if err != nil { + c.JSON(http.StatusBadRequest, response{Code: requestHeadscaleError, Message: err.Error()}) + return + } + if err := rejectWildcardApplicationPorts(requested); err != nil { + c.JSON(http.StatusBadRequest, response{Code: requestHeadscaleError, Message: err.Error()}) + return + } + requested = subtractUserDefaultPorts(requested, defaultApplicationServicePorts) + + policyUpdateMu.Lock() + defer policyUpdateMu.Unlock() + + current, document, err := loadApplicationPorts() + if err != nil { + c.JSON(http.StatusInternalServerError, response{Code: requestHeadscaleError, Message: err.Error()}) + return + } + + changed := !current.InSync || !userProtocolPortsEqual(current.ApplicationPorts, requested) + if changed { + // Headscale resolves user aliases while validating the policy. Ensure each + // real Olares owner has a corresponding Headscale user before referencing + // owner@. User creation alone grants no network access. + for _, entry := range requested { + if _, err := resolveUserIDString(entry.User); err != nil { + c.JSON(http.StatusBadGateway, response{Code: requestHeadscaleError, Message: fmt.Sprintf("ensure Headscale user %q: %v", entry.User, err)}) + return + } + } + if err := document.setManagedApplicationPorts(defaultApplicationServicePorts, requested); err != nil { + c.JSON(http.StatusConflict, response{Code: requestHeadscaleError, Message: err.Error()}) + return + } + encoded, err := document.marshal() + if err != nil { + c.JSON(http.StatusInternalServerError, response{Code: requestHeadscaleError, Message: err.Error()}) + return + } + updatedAt, err := setHeadscalePolicy(string(encoded)) + if err != nil { + c.JSON(http.StatusBadGateway, response{Code: requestHeadscaleError, Message: err.Error()}) + return + } + current.UpdatedAt = updatedAt + } + + current.DefaultPorts = cloneProtocolPorts(defaultApplicationServicePorts) + current.ApplicationPorts = requested + current.EffectivePorts = effectiveUserPorts(current.DefaultPorts, requested) + current.Revision = userProtocolPortsRevision(requested) + current.InSync = true + current.Changed = changed + c.JSON(http.StatusOK, response{Code: 0, Message: "", Data: current}) +} + +func readApplicationPorts() (applicationPortsState, error) { + state, _, err := loadApplicationPorts() + return state, err +} + +func loadApplicationPorts() (applicationPortsState, *policyDocument, error) { + policy, err := getHeadscalePolicy() + if err != nil { + return applicationPortsState{}, nil, err + } + document, err := parsePolicyDocument([]byte(policy.Policy)) + if err != nil { + return applicationPortsState{}, nil, err + } + policyDefaults, application, err := document.managedApplicationPorts() + if err != nil { + return applicationPortsState{}, nil, err + } + defaults := cloneProtocolPorts(defaultApplicationServicePorts) + return applicationPortsState{ + DefaultPorts: defaults, + ApplicationPorts: application, + EffectivePorts: effectiveUserPorts(defaults, application), + Revision: userProtocolPortsRevision(application), + UpdatedAt: policy.UpdatedAt, + InSync: protocolPortsEqual(policyDefaults, defaults) && + document.applicationRulesCanonical(application), + }, document, nil +} + +func getHeadscalePolicy() (headscalePolicyResponse, error) { + var result headscalePolicyResponse + resp, err := resty.New().SetTimeout(10 * time.Second).R(). + SetHeaders(headers). + SetResult(&result). + Get(url + "/policy") + if err != nil { + return result, fmt.Errorf("get Headscale policy: %w", err) + } + if resp.StatusCode() != http.StatusOK { + return result, fmt.Errorf("get Headscale policy: status %d: %s", resp.StatusCode(), resp.String()) + } + if strings.TrimSpace(result.Policy) == "" { + return result, errors.New("Headscale returned an empty policy") + } + return result, nil +} + +func setHeadscalePolicy(policy string) (string, error) { + var result headscalePolicyResponse + resp, err := resty.New().SetTimeout(10 * time.Second).R(). + SetHeaders(headers). + SetBody(map[string]string{"policy": policy}). + SetResult(&result). + Put(url + "/policy") + if err != nil { + return "", fmt.Errorf("set Headscale policy: %w", err) + } + if resp.StatusCode() != http.StatusOK { + return "", fmt.Errorf("set Headscale policy: status %d: %s", resp.StatusCode(), resp.String()) + } + return result.UpdatedAt, nil +} + +func parsePolicyDocument(policy []byte) (*policyDocument, error) { + standard, err := hujson.Standardize(policy) + if err != nil { + return nil, fmt.Errorf("parse Headscale HuJSON policy: %w", err) + } + fields := make(map[string]json.RawMessage) + if err := json.Unmarshal(standard, &fields); err != nil { + return nil, fmt.Errorf("decode Headscale policy: %w", err) + } + rawACLs, ok := fields["acls"] + if !ok { + return nil, errors.New("Headscale policy has no acls field") + } + var acls []policyACL + if err := json.Unmarshal(rawACLs, &acls); err != nil { + return nil, fmt.Errorf("decode Headscale policy acls: %w", err) + } + return &policyDocument{fields: fields, ACLs: acls}, nil +} + +func (p *policyDocument) marshal() ([]byte, error) { + acls, err := json.Marshal(p.ACLs) + if err != nil { + return nil, fmt.Errorf("encode Headscale policy acls: %w", err) + } + p.fields["acls"] = acls + encoded, err := json.MarshalIndent(p.fields, "", " ") + if err != nil { + return nil, fmt.Errorf("encode Headscale policy: %w", err) + } + return encoded, nil +} + +func (p *policyDocument) managedApplicationPorts() (protocolPorts, []userProtocolPorts, error) { + tcpIndex, udpIndex, err := p.defaultRuleIndexes() + if err != nil { + return protocolPorts{}, nil, err + } + tcp, err := portsFromDestinations(p.ACLs[tcpIndex].Dst, "default TCP") + if err != nil { + return protocolPorts{}, nil, err + } + udp, err := portsFromDestinations(p.ACLs[udpIndex].Dst, "default UDP") + if err != nil { + return protocolPorts{}, nil, err + } + defaults, err := normalizeProtocolPorts(protocolPorts{TCP: tcp, UDP: udp}) + if err != nil { + return protocolPorts{}, nil, err + } + + byUser := make(map[string]protocolPorts) + for _, acl := range p.ACLs { + user, proto, managed := managedUserRule(acl) + if !managed { + continue + } + ports, err := portsFromDestinations(acl.Dst, user+" "+proto) + if err != nil { + return protocolPorts{}, nil, err + } + entry := byUser[user] + if proto == "tcp" { + entry.TCP = append(entry.TCP, ports...) + } else { + entry.UDP = append(entry.UDP, ports...) + } + byUser[user] = entry + } + application := make([]userProtocolPorts, 0, len(byUser)) + for user, ports := range byUser { + normalized, err := normalizeProtocolPorts(ports) + if err != nil { + return protocolPorts{}, nil, fmt.Errorf("invalid managed ports for user %q: %w", user, err) + } + application = append(application, userProtocolPorts{User: user, protocolPorts: normalized}) + } + application, err = normalizeUserProtocolPorts(application) + return defaults, application, err +} + +func (p *policyDocument) setManagedApplicationPorts(defaults protocolPorts, application []userProtocolPorts) error { + tcpIndex, udpIndex, err := p.defaultRuleIndexes() + if err != nil { + return err + } + p.ACLs[tcpIndex].Dst = []string{"tag:olares:" + strings.Join(defaults.TCP, ",")} + p.ACLs[udpIndex].Dst = []string{"tag:olares:" + strings.Join(defaults.UDP, ",")} + + retained := make([]policyACL, 0, len(p.ACLs)+len(application)*2) + for _, acl := range p.ACLs { + if _, _, managed := managedUserRule(acl); managed { + continue + } + retained = append(retained, acl) + } + for _, entry := range application { + if len(entry.TCP) > 0 { + retained = append(retained, newManagedUserACL(entry.User, "tcp", entry.TCP)) + } + if len(entry.UDP) > 0 { + retained = append(retained, newManagedUserACL(entry.User, "udp", entry.UDP)) + } + } + p.ACLs = retained + return nil +} + +func (p *policyDocument) applicationRulesCanonical(application []userProtocolPorts) bool { + expected := make(map[string]string) + for _, entry := range application { + if len(entry.TCP) > 0 { + expected[entry.User+"\x00tcp"] = strings.Join(entry.TCP, ",") + } + if len(entry.UDP) > 0 { + expected[entry.User+"\x00udp"] = strings.Join(entry.UDP, ",") + } + } + actual := make(map[string]string) + for _, acl := range p.ACLs { + user, proto, managed := managedUserRule(acl) + if !managed { + continue + } + key := user + "\x00" + proto + if _, duplicate := actual[key]; duplicate { + return false + } + ports, err := portsFromDestinations(acl.Dst, user+" "+proto) + if err != nil { + return false + } + actual[key] = strings.Join(ports, ",") + } + if len(actual) != len(expected) { + return false + } + for key, ports := range expected { + if actual[key] != ports { + return false + } + } + return true +} + +func (p *policyDocument) defaultRuleIndexes() (int, int, error) { + tcpIndex, udpIndex := -1, -1 + for i, acl := range p.ACLs { + if acl.Action != "accept" || len(acl.Src) != 1 || acl.Src[0] != "autogroup:member" || len(acl.Dst) != 1 || !strings.HasPrefix(acl.Dst[0], "tag:olares:") { + continue + } + switch strings.ToLower(acl.Proto) { + case "tcp": + if tcpIndex != -1 { + return -1, -1, errors.New("Headscale policy has multiple managed default TCP service rules") + } + tcpIndex = i + case "udp": + if udpIndex != -1 { + return -1, -1, errors.New("Headscale policy has multiple managed default UDP service rules") + } + udpIndex = i + } + } + if tcpIndex == -1 || udpIndex == -1 { + return -1, -1, errors.New("Headscale policy is missing the managed default TCP or UDP service rule") + } + return tcpIndex, udpIndex, nil +} + +func managedUserRule(acl policyACL) (string, string, bool) { + if acl.Action != "accept" || len(acl.Src) != 1 || len(acl.Dst) != 1 || !strings.HasPrefix(acl.Dst[0], "tag:olares:") { + return "", "", false + } + source := acl.Src[0] + if !strings.HasSuffix(source, "@") || strings.HasPrefix(source, "autogroup:") { + return "", "", false + } + proto := strings.ToLower(acl.Proto) + if proto != "tcp" && proto != "udp" { + return "", "", false + } + user := strings.TrimSuffix(source, "@") + if validatePolicyUsername(user) != nil { + return "", "", false + } + return user, proto, true +} + +func newManagedUserACL(user, proto string, ports []string) policyACL { + return policyACL{ + Action: "accept", + Src: []string{user + "@"}, + Proto: proto, + Dst: []string{"tag:olares:" + strings.Join(ports, ",")}, + } +} + +func portsFromDestinations(destinations []string, description string) ([]string, error) { + if len(destinations) != 1 { + return nil, fmt.Errorf("managed %s service rule must have exactly one destination", description) + } + const prefix = "tag:olares:" + if !strings.HasPrefix(destinations[0], prefix) { + return nil, fmt.Errorf("managed %s service rule has an unexpected destination", description) + } + value := strings.TrimPrefix(destinations[0], prefix) + if value == "" { + return nil, fmt.Errorf("managed %s service rule has no ports", description) + } + return strings.Split(value, ","), nil +} + +func normalizeUserProtocolPorts(entries []userProtocolPorts) ([]userProtocolPorts, error) { + byUser := make(map[string]protocolPorts) + for _, entry := range entries { + user := strings.TrimSpace(entry.User) + if err := validatePolicyUsername(user); err != nil { + return nil, err + } + ports, err := normalizeProtocolPorts(entry.protocolPorts) + if err != nil { + return nil, fmt.Errorf("invalid ports for user %q: %w", user, err) + } + current := byUser[user] + current.TCP = append(current.TCP, ports.TCP...) + current.UDP = append(current.UDP, ports.UDP...) + byUser[user] = current + } + result := make([]userProtocolPorts, 0, len(byUser)) + for user, ports := range byUser { + normalized, _ := normalizeProtocolPorts(ports) + if len(normalized.TCP) == 0 && len(normalized.UDP) == 0 { + continue + } + result = append(result, userProtocolPorts{User: user, protocolPorts: normalized}) + } + sort.Slice(result, func(i, j int) bool { return result[i].User < result[j].User }) + return result, nil +} + +func rejectWildcardApplicationPorts(entries []userProtocolPorts) error { + for _, entry := range entries { + for _, port := range append(append([]string{}, entry.TCP...), entry.UDP...) { + if port == "*" { + return fmt.Errorf("application ports for user %q must not include wildcard ports", entry.User) + } + } + } + return nil +} + +func validatePolicyUsername(user string) error { + if user == "" { + return errors.New("application port owner is empty") + } + if strings.ContainsAny(user, "@:") || strings.IndexFunc(user, unicode.IsSpace) >= 0 { + return fmt.Errorf("application port owner %q is not a valid Headscale username", user) + } + return nil +} + +func normalizeProtocolPorts(ports protocolPorts) (protocolPorts, error) { + tcp, err := normalizePorts(ports.TCP) + if err != nil { + return protocolPorts{}, fmt.Errorf("invalid TCP ports: %w", err) + } + udp, err := normalizePorts(ports.UDP) + if err != nil { + return protocolPorts{}, fmt.Errorf("invalid UDP ports: %w", err) + } + return protocolPorts{TCP: tcp, UDP: udp}, nil +} + +func normalizePorts(ports []string) ([]string, error) { + set := make(map[string]struct{}) + for _, group := range ports { + for _, raw := range strings.Split(group, ",") { + port, err := normalizePortSpec(raw) + if err != nil { + return nil, err + } + if port == "*" { + return []string{"*"}, nil + } + set[port] = struct{}{} + } + } + result := make([]string, 0, len(set)) + for port := range set { + result = append(result, port) + } + sort.Slice(result, func(i, j int) bool { + leftStart, leftEnd := portSpecBounds(result[i]) + rightStart, rightEnd := portSpecBounds(result[j]) + if leftStart == rightStart { + return leftEnd < rightEnd + } + return leftStart < rightStart + }) + return result, nil +} + +func normalizePortSpec(raw string) (string, error) { + port := strings.TrimSpace(raw) + if port == "" { + return "", errors.New("empty port") + } + if port == "*" { + return port, nil + } + parts := strings.Split(port, "-") + if len(parts) > 2 || len(parts) == 2 && (parts[0] == "" || parts[1] == "") { + return "", fmt.Errorf("%q must be a port or inclusive port range", port) + } + first, err := strconv.ParseUint(parts[0], 10, 16) + if err != nil || first == 0 { + return "", fmt.Errorf("%q must contain ports between 1 and 65535", port) + } + if len(parts) == 1 { + return strconv.FormatUint(first, 10), nil + } + last, err := strconv.ParseUint(parts[1], 10, 16) + if err != nil || last == 0 || first > last { + return "", fmt.Errorf("%q must be an increasing port range between 1 and 65535", port) + } + return strconv.FormatUint(first, 10) + "-" + strconv.FormatUint(last, 10), nil +} + +func portSpecBounds(spec string) (int, int) { + parts := strings.Split(spec, "-") + first, _ := strconv.Atoi(parts[0]) + if len(parts) == 1 { + return first, first + } + last, _ := strconv.Atoi(parts[1]) + return first, last +} + +func subtractUserDefaultPorts(entries []userProtocolPorts, defaults protocolPorts) []userProtocolPorts { + result := make([]userProtocolPorts, 0, len(entries)) + for _, entry := range entries { + ports := subtractProtocolPorts(entry.protocolPorts, defaults) + if len(ports.TCP) == 0 && len(ports.UDP) == 0 { + continue + } + result = append(result, userProtocolPorts{User: entry.User, protocolPorts: ports}) + } + return result +} + +func effectiveUserPorts(defaults protocolPorts, entries []userProtocolPorts) []userProtocolPorts { + result := make([]userProtocolPorts, 0, len(entries)) + for _, entry := range entries { + result = append(result, userProtocolPorts{User: entry.User, protocolPorts: mergeProtocolPorts(defaults, entry.protocolPorts)}) + } + return result +} + +func mergeProtocolPorts(left, right protocolPorts) protocolPorts { + merged, _ := normalizeProtocolPorts(protocolPorts{ + TCP: append(append([]string{}, left.TCP...), right.TCP...), + UDP: append(append([]string{}, left.UDP...), right.UDP...), + }) + return merged +} + +func subtractProtocolPorts(all, remove protocolPorts) protocolPorts { + return protocolPorts{ + TCP: subtractPorts(all.TCP, remove.TCP), + UDP: subtractPorts(all.UDP, remove.UDP), + } +} + +func subtractPorts(all, remove []string) []string { + removed := make(map[string]struct{}, len(remove)) + for _, port := range remove { + removed[port] = struct{}{} + } + result := make([]string, 0, len(all)) + for _, port := range all { + if _, ok := removed[port]; !ok { + result = append(result, port) + } + } + return result +} + +func cloneProtocolPorts(ports protocolPorts) protocolPorts { + return protocolPorts{TCP: append([]string(nil), ports.TCP...), UDP: append([]string(nil), ports.UDP...)} +} + +func protocolPortsEqual(left, right protocolPorts) bool { + return strings.Join(left.TCP, ",") == strings.Join(right.TCP, ",") && strings.Join(left.UDP, ",") == strings.Join(right.UDP, ",") +} + +func userProtocolPortsEqual(left, right []userProtocolPorts) bool { + if len(left) != len(right) { + return false + } + for i := range left { + if left[i].User != right[i].User || !protocolPortsEqual(left[i].protocolPorts, right[i].protocolPorts) { + return false + } + } + return true +} + +func userProtocolPortsRevision(ports []userProtocolPorts) string { + encoded, _ := json.Marshal(ports) + digest := sha256.Sum256(encoded) + return "sha256:" + hex.EncodeToString(digest[:]) +} diff --git a/service_auth.go b/service_auth.go new file mode 100644 index 0000000..8a984d9 --- /dev/null +++ b/service_auth.go @@ -0,0 +1,173 @@ +package main + +import ( + "bytes" + "crypto/tls" + "crypto/x509" + "encoding/json" + "errors" + "fmt" + "io" + "net/http" + "os" + "strings" + "time" + + "github.com/gin-gonic/gin" +) + +const ( + serviceAccountCAPath = "/var/run/secrets/kubernetes.io/serviceaccount/ca.crt" + serviceAccountTokenPath = "/var/run/secrets/kubernetes.io/serviceaccount/token" + defaultServiceAuthNS = "os-framework" + defaultServiceAuthSA = "os-internal" + defaultServiceAudience = "headscale-policy" +) + +type tokenReview struct { + APIVersion string `json:"apiVersion"` + Kind string `json:"kind"` + Spec tokenReviewSpec `json:"spec"` + Status tokenReviewStatus `json:"status,omitempty"` +} + +type tokenReviewSpec struct { + Token string `json:"token"` + Audiences []string `json:"audiences"` +} + +type tokenReviewStatus struct { + Authenticated bool `json:"authenticated"` + User struct { + Username string `json:"username"` + } `json:"user"` + Audiences []string `json:"audiences,omitempty"` + Error string `json:"error,omitempty"` +} + +var reviewServiceAccountToken = kubernetesTokenReview + +func requireServiceAccount() gin.HandlerFunc { + return func(c *gin.Context) { + authorization := strings.Fields(c.GetHeader("Authorization")) + if len(authorization) != 2 || !strings.EqualFold(authorization[0], "Bearer") || authorization[1] == "" { + c.AbortWithStatusJSON(http.StatusUnauthorized, response{Code: requestHeadscaleError, Message: "missing service account bearer token"}) + return + } + token := authorization[1] + + username, err := reviewServiceAccountToken(token) + if err != nil { + c.AbortWithStatusJSON(http.StatusUnauthorized, response{Code: requestHeadscaleError, Message: err.Error()}) + return + } + + expectedNS := strings.TrimSpace(os.Getenv("POLICY_CLIENT_NAMESPACE")) + if expectedNS == "" { + expectedNS = defaultServiceAuthNS + } + expectedSA := strings.TrimSpace(os.Getenv("POLICY_CLIENT_SERVICE_ACCOUNT")) + if expectedSA == "" { + expectedSA = defaultServiceAuthSA + } + expectedUsername := fmt.Sprintf("system:serviceaccount:%s:%s", expectedNS, expectedSA) + if username != expectedUsername { + c.AbortWithStatusJSON(http.StatusForbidden, response{Code: requestHeadscaleError, Message: "service account is not allowed to manage Headscale policy"}) + return + } + + c.Next() + } +} + +func kubernetesTokenReview(token string) (string, error) { + host := strings.TrimSpace(os.Getenv("KUBERNETES_SERVICE_HOST")) + port := strings.TrimSpace(os.Getenv("KUBERNETES_SERVICE_PORT")) + if host == "" || port == "" { + return "", errors.New("Kubernetes API service is unavailable") + } + + caPEM, err := os.ReadFile(serviceAccountCAPath) + if err != nil { + return "", fmt.Errorf("read Kubernetes service account CA: %w", err) + } + roots := x509.NewCertPool() + if !roots.AppendCertsFromPEM(caPEM) { + return "", errors.New("parse Kubernetes service account CA") + } + + audience := strings.TrimSpace(os.Getenv("POLICY_TOKEN_AUDIENCE")) + if audience == "" { + audience = defaultServiceAudience + } + payload, err := json.Marshal(tokenReview{ + APIVersion: "authentication.k8s.io/v1", + Kind: "TokenReview", + Spec: tokenReviewSpec{ + Token: token, + Audiences: []string{audience}, + }, + }) + if err != nil { + return "", fmt.Errorf("encode Kubernetes token review: %w", err) + } + reviewerToken, err := os.ReadFile(serviceAccountTokenPath) + if err != nil { + return "", fmt.Errorf("read wrapper service account token: %w", err) + } + if strings.TrimSpace(string(reviewerToken)) == "" { + return "", errors.New("wrapper service account token is empty") + } + + client := &http.Client{ + Timeout: 10 * time.Second, + Transport: &http.Transport{TLSClientConfig: &tls.Config{ + MinVersion: tls.VersionTLS12, + RootCAs: roots, + }}, + } + req, err := http.NewRequest(http.MethodPost, fmt.Sprintf("https://%s:%s/apis/authentication.k8s.io/v1/tokenreviews", host, port), bytes.NewReader(payload)) + if err != nil { + return "", fmt.Errorf("create Kubernetes token review request: %w", err) + } + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Authorization", "Bearer "+strings.TrimSpace(string(reviewerToken))) + + resp, err := client.Do(req) + if err != nil { + return "", fmt.Errorf("review Kubernetes service account token: %w", err) + } + defer resp.Body.Close() + body, err := io.ReadAll(io.LimitReader(resp.Body, 1<<20)) + if err != nil { + return "", fmt.Errorf("read Kubernetes token review response: %w", err) + } + if resp.StatusCode != http.StatusCreated { + return "", fmt.Errorf("review Kubernetes service account token: status %d", resp.StatusCode) + } + + var reviewed tokenReview + if err := json.Unmarshal(body, &reviewed); err != nil { + return "", fmt.Errorf("decode Kubernetes token review response: %w", err) + } + if !reviewed.Status.Authenticated { + if reviewed.Status.Error != "" { + return "", fmt.Errorf("service account token was not authenticated: %s", reviewed.Status.Error) + } + return "", errors.New("service account token was not authenticated") + } + if reviewed.Status.User.Username == "" { + return "", errors.New("Kubernetes token review returned an empty username") + } + audienceAccepted := false + for _, reviewedAudience := range reviewed.Status.Audiences { + if reviewedAudience == audience { + audienceAccepted = true + break + } + } + if !audienceAccepted { + return "", fmt.Errorf("service account token is not valid for audience %q", audience) + } + return reviewed.Status.User.Username, nil +}