Implement authenticated federation worker transport
This commit is contained in:
+168
-11
@@ -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(®istration) != 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")}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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):
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user