Implement authenticated federation worker transport

This commit is contained in:
kami
2026-07-26 20:38:59 +04:00
parent 4364dff9c2
commit 8c9b6499e2
3 changed files with 277 additions and 15 deletions
+168 -11
View File
@@ -92,6 +92,17 @@ func main() {
}
mux := http.NewServeMux()
workers := &federation.Registry{}
workers.OnOffline = func(w federation.Worker) {
for _, t := range s.Tasks() {
if t.State == domain.StateLeased && t.Lease != nil && t.Lease.HarnessID == w.ID {
p, _ := json.Marshal(map[string]any{"reason": "worker_offline", "harness_id": w.ID})
e := domain.Event{ID: id(), Type: "TaskReleased", TaskID: t.ID, Version: t.Version + 1, Payload: p}
if err := s.Append(e); err == nil && rt != nil {
_, _ = rt.HandleEvent(e)
}
}
}
}
providerHealth := map[string]*provider.Supervisor{}
surface := func(r *http.Request) authz.Surface {
v := authz.ParseSurface(r.Header.Get("X-Orchestra-Surface"))
@@ -173,9 +184,7 @@ func main() {
json.NewEncoder(w).Encode(operations.BuildBrief(s.Events(0), from, to, operations.GitState(dir)))
})
standup := func() (domain.Event, error) {
items, _ := json.Marshal(map[string]any{"items": operations.StandupItems(s.Tasks()), "generated_at": time.Now().UTC()})
e := domain.Event{ID: id(), Type: "StandupAdvisory", TaskID: "system", Version: 0, Payload: items}
return e, s.Append(e)
return operations.GenerateStandupAdvisory(s, time.Now().UTC())
}
mux.HandleFunc("/v1/standup", func(w http.ResponseWriter, r *http.Request) {
if r.Method == http.MethodGet {
@@ -197,6 +206,53 @@ func main() {
}
json.NewEncoder(w).Encode(e)
})
mux.HandleFunc("/v1/standup/approve", func(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
http.Error(w, "method not allowed", 405)
return
}
if err := authz.AuthorizeEvent(surface(r), "ApprovalGranted"); err != nil {
http.Error(w, err.Error(), 403)
return
}
var p struct {
AdvisoryID string `json:"advisory_id"`
}
if json.NewDecoder(r.Body).Decode(&p) != nil || p.AdvisoryID == "" {
http.Error(w, "advisory_id required", 400)
return
}
b, _ := json.Marshal(map[string]any{"subject_ref": p.AdvisoryID})
e := domain.Event{ID: id(), Type: "ApprovalGranted", TaskID: "system", Version: 0, Payload: b}
if err := s.Append(e); err != nil {
http.Error(w, err.Error(), 400)
return
}
json.NewEncoder(w).Encode(e)
})
mux.HandleFunc("/v1/standup/apply", func(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
http.Error(w, "method not allowed", 405)
return
}
if err := authz.AuthorizeEvent(surface(r), "TaskAmended"); err != nil {
http.Error(w, err.Error(), 403)
return
}
var p struct {
AdvisoryID string `json:"advisory_id"`
}
if json.NewDecoder(r.Body).Decode(&p) != nil || p.AdvisoryID == "" {
http.Error(w, "advisory_id required", 400)
return
}
out, err := operations.ApplyAdvisory(s, p.AdvisoryID)
if err != nil {
http.Error(w, err.Error(), 400)
return
}
json.NewEncoder(w).Encode(out)
})
mux.HandleFunc("/metrics", func(w http.ResponseWriter, r *http.Request) {
tasks := s.Tasks()
counts := map[domain.TaskState]int{}
@@ -375,11 +431,20 @@ func main() {
http.Error(w, "method not allowed", 405)
return
}
var worker federation.Worker
if json.NewDecoder(http.MaxBytesReader(w, r.Body, 64<<10)).Decode(&worker) != nil {
var registration struct {
federation.Worker
Token string `json:"token"`
}
if json.NewDecoder(http.MaxBytesReader(w, r.Body, 64<<10)).Decode(&registration) != nil {
http.Error(w, "invalid worker", 400)
return
}
worker := registration.Worker
worker.Token = registration.Token
if worker.Token == "" {
http.Error(w, "token required", 400)
return
}
if err := workers.Register(worker); err != nil {
http.Error(w, err.Error(), 400)
return
@@ -387,16 +452,57 @@ func main() {
w.WriteHeader(http.StatusCreated)
json.NewEncoder(w).Encode(worker)
})
workerAuth := func(r *http.Request) (string, error) {
wid := r.Header.Get("X-Orchestra-Worker")
tok := strings.TrimPrefix(r.Header.Get("Authorization"), "Bearer ")
if err := workers.Authenticate(wid, tok); err != nil {
return "", err
}
return wid, nil
}
mux.HandleFunc("/v1/federation/events", func(w http.ResponseWriter, r *http.Request) {
wid, err := workerAuth(r)
if err != nil {
http.Error(w, err.Error(), http.StatusUnauthorized)
return
}
if r.Method != http.MethodGet {
http.Error(w, "method not allowed", 405)
return
}
cursor, _ := strconv.ParseUint(r.URL.Query().Get("since"), 10, 64)
json.NewEncoder(w).Encode(s.Events(cursor))
cursor, _ := workers.Cursor(wid)
if q := r.URL.Query().Get("since"); q != "" {
cursor, _ = strconv.ParseUint(q, 10, 64)
}
events := s.Events(cursor)
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(map[string]any{"cursor": cursor, "events": events})
})
mux.HandleFunc("/v1/federation/events/ack", func(w http.ResponseWriter, r *http.Request) {
wid, err := workerAuth(r)
if err != nil {
http.Error(w, err.Error(), 401)
return
}
if r.Method != http.MethodPost {
http.Error(w, "method not allowed", 405)
return
}
var body struct {
Cursor uint64 `json:"cursor"`
}
if json.NewDecoder(r.Body).Decode(&body) != nil {
http.Error(w, "invalid cursor", 400)
return
}
if err := workers.Ack(wid, body.Cursor); err != nil {
http.Error(w, err.Error(), 409)
return
}
w.WriteHeader(http.StatusNoContent)
})
mux.HandleFunc("/v1/federation/workers/", func(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost || !strings.HasSuffix(r.URL.Path, "/heartbeat") {
if r.Method != http.MethodPost || (!strings.HasSuffix(r.URL.Path, "/heartbeat") && !strings.HasSuffix(r.URL.Path, "/claim") && !strings.HasSuffix(r.URL.Path, "/handoff")) {
http.Error(w, "not found", 404)
return
}
@@ -405,11 +511,62 @@ func main() {
http.Error(w, "not found", 404)
return
}
if err := workers.Heartbeat(parts[3]); err != nil {
http.Error(w, err.Error(), 404)
if _, err := workerAuth(r); err != nil {
http.Error(w, err.Error(), 401)
return
}
w.WriteHeader(http.StatusNoContent)
if strings.HasSuffix(r.URL.Path, "/heartbeat") {
if err := workers.Heartbeat(parts[3]); err != nil {
http.Error(w, err.Error(), 404)
return
}
w.WriteHeader(http.StatusNoContent)
return
}
var b struct {
TaskID string `json:"task_id"`
TTLSeconds int `json:"ttl_seconds"`
HandoffRef string `json:"handoff_ref"`
}
if json.NewDecoder(r.Body).Decode(&b) != nil || b.TaskID == "" {
http.Error(w, "invalid lease body", 400)
return
}
t, ok := s.Task(b.TaskID)
if !ok {
http.Error(w, "task not found", 404)
return
}
if strings.HasSuffix(r.URL.Path, "/claim") {
if b.TTLSeconds <= 0 {
b.TTLSeconds = 1800
}
e, err := s.Lease(b.TaskID, parts[3], time.Duration(b.TTLSeconds)*time.Second)
if err != nil {
http.Error(w, err.Error(), 409)
return
}
json.NewEncoder(w).Encode(e)
return
}
if t.State != domain.StateLeased || t.Lease == nil || t.Lease.HarnessID != parts[3] {
http.Error(w, "lease not owned", 409)
return
}
if b.HandoffRef == "" {
http.Error(w, "handoff_ref required", 400)
return
}
p, _ := json.Marshal(map[string]any{"handoff_ref": b.HandoffRef, "harness_id": parts[3]})
e := domain.Event{ID: id(), Type: "TaskReleased", TaskID: b.TaskID, Version: t.Version + 1, Payload: p}
if err := s.Append(e); err != nil {
http.Error(w, err.Error(), 409)
return
}
if rt != nil {
_, _ = rt.HandleEvent(e)
}
json.NewEncoder(w).Encode(e)
})
if base := os.Getenv("ORCHESTRA_GITEA_URL"); base != "" {
g := provider.Gitea{BaseURL: base, Token: os.Getenv("ORCHESTRA_GITEA_TOKEN"), WebhookSecret: os.Getenv("ORCHESTRA_GITEA_WEBHOOK_SECRET"), Owner: os.Getenv("ORCHESTRA_GITEA_OWNER"), Repo: os.Getenv("ORCHESTRA_GITEA_REPO")}
+53 -4
View File
@@ -7,6 +7,7 @@ import (
)
var ErrUnknownWorker = errors.New("unknown worker")
var ErrUnauthorized = errors.New("worker authentication failed")
type Worker struct {
ID string `json:"id"`
@@ -14,12 +15,15 @@ type Worker struct {
Capacity int `json:"capacity"`
LastSeen time.Time `json:"last_seen"`
Online bool `json:"online"`
Token string `json:"-"`
}
type Registry struct {
mu sync.Mutex
workers map[string]Worker
TTL time.Duration
mu sync.Mutex
workers map[string]Worker
TTL time.Duration
OnOffline func(Worker)
cursors map[string]uint64
}
func (r *Registry) init() {
@@ -29,6 +33,9 @@ func (r *Registry) init() {
if r.workers == nil {
r.workers = map[string]Worker{}
}
if r.cursors == nil {
r.cursors = map[string]uint64{}
}
}
func (r *Registry) Register(w Worker) error {
if w.ID == "" {
@@ -40,6 +47,44 @@ func (r *Registry) Register(w Worker) error {
w.LastSeen = time.Now().UTC()
w.Online = true
r.workers[w.ID] = w
if _, ok := r.cursors[w.ID]; !ok {
r.cursors[w.ID] = 0
}
return nil
}
func (r *Registry) Authenticate(id, token string) error {
r.mu.Lock()
defer r.mu.Unlock()
r.init()
w, ok := r.workers[id]
if !ok {
return ErrUnknownWorker
}
if w.Token == "" || token == "" || w.Token != token {
return ErrUnauthorized
}
return nil
}
func (r *Registry) Cursor(id string) (uint64, error) {
r.mu.Lock()
defer r.mu.Unlock()
r.init()
if _, ok := r.workers[id]; !ok {
return 0, ErrUnknownWorker
}
return r.cursors[id], nil
}
func (r *Registry) Ack(id string, cursor uint64) error {
r.mu.Lock()
defer r.mu.Unlock()
r.init()
if _, ok := r.workers[id]; !ok {
return ErrUnknownWorker
}
if cursor < r.cursors[id] {
return errors.New("cursor moved backwards")
}
r.cursors[id] = cursor
return nil
}
func (r *Registry) Heartbeat(id string) error {
@@ -59,11 +104,15 @@ func (r *Registry) Snapshot() []Worker {
r.mu.Lock()
defer r.mu.Unlock()
r.init()
now := time.Now()
now := time.Now().UTC()
out := make([]Worker, 0, len(r.workers))
for id, w := range r.workers {
wasOnline := w.Online
w.Online = now.Sub(w.LastSeen) <= r.TTL
r.workers[id] = w
if wasOnline && !w.Online && r.OnOffline != nil {
go r.OnOffline(w)
}
out = append(out, w)
}
return out
+56
View File
@@ -0,0 +1,56 @@
package federation
import (
"testing"
"time"
)
func TestCursorIsMonotonicAndAuthenticationIsRequired(t *testing.T) {
r := &Registry{}
if err := r.Register(Worker{ID: "workpc", Token: "secret"}); err != nil {
t.Fatal(err)
}
if err := r.Authenticate("workpc", "wrong"); err != ErrUnauthorized {
t.Fatalf("got %v", err)
}
if err := r.Authenticate("workpc", "secret"); err != nil {
t.Fatal(err)
}
if err := r.Ack("workpc", 7); err != nil {
t.Fatal(err)
}
if err := r.Ack("workpc", 6); err == nil {
t.Fatal("backwards cursor accepted")
}
if got, _ := r.Cursor("workpc"); got != 7 {
t.Fatalf("cursor = %d", got)
}
}
func TestOfflineHookRunsOnceOnTransition(t *testing.T) {
called := make(chan Worker, 1)
r := &Registry{TTL: time.Millisecond, OnOffline: func(w Worker) { called <- w }}
if err := r.Register(Worker{ID: "workpc"}); err != nil {
t.Fatal(err)
}
r.mu.Lock()
w := r.workers["workpc"]
w.LastSeen = time.Now().Add(-time.Second)
r.workers["workpc"] = w
r.mu.Unlock()
r.Snapshot()
select {
case got := <-called:
if got.ID != "workpc" {
t.Fatal(got.ID)
}
case <-time.After(time.Second):
t.Fatal("offline hook not called")
}
r.Snapshot()
select {
case <-called:
t.Fatal("offline hook called twice")
case <-time.After(10 * time.Millisecond):
}
}