Replace Python worker daemon with the Go worker agent
This commit is contained in:
@@ -0,0 +1,120 @@
|
||||
package agent
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
// TokenProvider supplies the current bearer token. A static token is served
|
||||
// forever; a worker key is exchanged at the userservice for short-lived JWTs
|
||||
// and refreshed before they expire (mirroring the former Python worker).
|
||||
type TokenProvider interface {
|
||||
Token() (string, error)
|
||||
Refresh() error
|
||||
}
|
||||
|
||||
// StaticToken serves a fixed token forever; empty means no Authorization.
|
||||
type StaticToken struct{ token string }
|
||||
|
||||
func (s *StaticToken) Token() (string, error) { return s.token, nil }
|
||||
func (s *StaticToken) Refresh() error { return nil }
|
||||
|
||||
// WorkerKeyToken exchanges a long-lived worker key for short-lived JWTs.
|
||||
type WorkerKeyToken struct {
|
||||
userserviceURL string
|
||||
workerKey string
|
||||
timeout time.Duration
|
||||
leeway float64
|
||||
mu sync.Mutex
|
||||
token string
|
||||
refreshAt time.Time
|
||||
}
|
||||
|
||||
func NewWorkerKeyToken(userserviceURL, workerKey string, timeout time.Duration) *WorkerKeyToken {
|
||||
return &WorkerKeyToken{
|
||||
userserviceURL: strings.TrimRight(userserviceURL, "/"),
|
||||
workerKey: workerKey,
|
||||
timeout: timeout,
|
||||
leeway: 0.2,
|
||||
}
|
||||
}
|
||||
|
||||
// Token returns the current token, exchanging first when missing or stale.
|
||||
func (p *WorkerKeyToken) Token() (string, error) {
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
if p.token == "" || time.Now().After(p.refreshAt) {
|
||||
if err := p.exchangeLocked(); err != nil {
|
||||
return "", err
|
||||
}
|
||||
}
|
||||
return p.token, nil
|
||||
}
|
||||
|
||||
// Refresh forces an immediate exchange.
|
||||
func (p *WorkerKeyToken) Refresh() error {
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
return p.exchangeLocked()
|
||||
}
|
||||
|
||||
func (p *WorkerKeyToken) exchangeLocked() error {
|
||||
payload, err := json.Marshal(map[string]string{"key": p.workerKey})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
request, err := http.NewRequest(http.MethodPost, p.userserviceURL+"/worker-tokens/exchange", bytes.NewReader(payload))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
request.Header.Set("Content-Type", "application/json")
|
||||
client := &http.Client{Timeout: p.timeout}
|
||||
response, err := client.Do(request)
|
||||
if err != nil {
|
||||
return fmt.Errorf("worker key exchange request failed")
|
||||
}
|
||||
defer response.Body.Close()
|
||||
if response.StatusCode != http.StatusOK {
|
||||
return fmt.Errorf("worker key exchange rejected with status %d", response.StatusCode)
|
||||
}
|
||||
raw, err := io.ReadAll(io.LimitReader(response.Body, 1<<20))
|
||||
if err != nil {
|
||||
return fmt.Errorf("worker key exchange request failed")
|
||||
}
|
||||
var data map[string]any
|
||||
if err := json.Unmarshal(raw, &data); err != nil {
|
||||
return fmt.Errorf("worker key exchange response is invalid")
|
||||
}
|
||||
token, _ := data["token"].(string)
|
||||
if token == "" {
|
||||
return fmt.Errorf("worker key exchange response is missing a token")
|
||||
}
|
||||
var ttl time.Duration
|
||||
switch value := data["expires_in"].(type) {
|
||||
case float64:
|
||||
ttl = time.Duration(value * float64(time.Second))
|
||||
case int:
|
||||
ttl = time.Duration(value) * time.Second
|
||||
}
|
||||
p.token = token
|
||||
p.refreshAt = time.Time{}
|
||||
if ttl > 0 {
|
||||
p.refreshAt = time.Now().Add(time.Duration(float64(ttl) * (1.0 - p.leeway)))
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// NewTokenProvider picks the strategy: a worker key (with userservice) wins
|
||||
// over a static bearer token.
|
||||
func NewTokenProvider(workerKey, userserviceURL, bearerToken string, timeout time.Duration) TokenProvider {
|
||||
if workerKey != "" && userserviceURL != "" {
|
||||
return NewWorkerKeyToken(userserviceURL, workerKey, timeout)
|
||||
}
|
||||
return &StaticToken{token: bearerToken}
|
||||
}
|
||||
@@ -0,0 +1,102 @@
|
||||
package agent
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestWorkerKeyTokenExchangesAndCaches(t *testing.T) {
|
||||
var exchanges atomic.Int64
|
||||
var server *httptest.Server
|
||||
server = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.URL.Path != "/worker-tokens/exchange" {
|
||||
http.NotFound(w, r)
|
||||
return
|
||||
}
|
||||
var payload map[string]string
|
||||
if err := json.NewDecoder(r.Body).Decode(&payload); err != nil || payload["key"] != "scimesh_wk_live_x" {
|
||||
http.Error(w, "bad key", http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
exchanges.Add(1)
|
||||
writeJSON(w, http.StatusOK, map[string]any{
|
||||
"token": "jwt-1",
|
||||
"expires_in": 100,
|
||||
})
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
provider := NewWorkerKeyToken(server.URL, "scimesh_wk_live_x", 5*time.Second)
|
||||
token, err := provider.Token()
|
||||
if err != nil || token != "jwt-1" {
|
||||
t.Fatalf("token = %q, err = %v", token, err)
|
||||
}
|
||||
// The second call within the TTL reuses the cache.
|
||||
again, err := provider.Token()
|
||||
if err != nil || again != "jwt-1" {
|
||||
t.Fatalf("cached token = %q, err = %v", again, err)
|
||||
}
|
||||
if exchanges.Load() != 1 {
|
||||
t.Errorf("exchanges = %d, want 1", exchanges.Load())
|
||||
}
|
||||
// An explicit refresh re-exchanges.
|
||||
if err := provider.Refresh(); err != nil {
|
||||
t.Fatalf("Refresh: %v", err)
|
||||
}
|
||||
if exchanges.Load() != 2 {
|
||||
t.Errorf("exchanges after refresh = %d, want 2", exchanges.Load())
|
||||
}
|
||||
}
|
||||
|
||||
func TestWorkerKeyTokenRejectsBadKey(t *testing.T) {
|
||||
var server *httptest.Server
|
||||
server = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusUnauthorized)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
provider := NewWorkerKeyToken(server.URL, "bad", 5*time.Second)
|
||||
if _, err := provider.Token(); err == nil {
|
||||
t.Error("expected exchange failure for a rejected key")
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewTokenProviderSelectsStrategy(t *testing.T) {
|
||||
if _, ok := NewTokenProvider("", "", "static", time.Second).(*StaticToken); !ok {
|
||||
t.Error("expected a static token provider")
|
||||
}
|
||||
if _, ok := NewTokenProvider("key", "http://users", "", time.Second).(*WorkerKeyToken); !ok {
|
||||
t.Error("expected a worker-key provider")
|
||||
}
|
||||
}
|
||||
|
||||
func TestClientRefreshesTokenOnceOn401(t *testing.T) {
|
||||
var attempts atomic.Int64
|
||||
var server *httptest.Server
|
||||
server = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
attempts.Add(1)
|
||||
if attempts.Load() == 1 {
|
||||
w.WriteHeader(http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
writeJSON(w, http.StatusOK, map[string]any{})
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
provider := &StaticToken{token: "t"}
|
||||
client := NewClient(server.URL, provider, 5*time.Second)
|
||||
status, _, err := client.requestJSON(http.MethodGet, "/ok", map[string]any{})
|
||||
if err != nil {
|
||||
t.Fatalf("request: %v", err)
|
||||
}
|
||||
if status != http.StatusOK {
|
||||
t.Errorf("status = %d", status)
|
||||
}
|
||||
if attempts.Load() != 2 {
|
||||
t.Errorf("attempts = %d, want 2 (401 then retry)", attempts.Load())
|
||||
}
|
||||
}
|
||||
@@ -31,23 +31,25 @@ type ConflictError struct{ msg string }
|
||||
|
||||
func (e *ConflictError) Error() string { return e.msg }
|
||||
|
||||
// Client speaks the v1 worker contract over HTTP with a static bearer token.
|
||||
// Client speaks the v1 worker contract over HTTP with a token provider.
|
||||
//
|
||||
// API calls never follow redirects (a redirect is a contract violation); the
|
||||
// artifact download follows redirects but strips the Authorization header on
|
||||
// cross-origin hops, matching the Python worker's SameOriginAuthRedirectHandler.
|
||||
// A 401 response refreshes the token exactly once and retries, so a lapsed JWT
|
||||
// does not fail an in-flight task.
|
||||
type Client struct {
|
||||
baseURL string
|
||||
token string
|
||||
tokens TokenProvider
|
||||
timeout time.Duration
|
||||
apiClient *http.Client
|
||||
dlClient *http.Client
|
||||
}
|
||||
|
||||
func NewClient(baseURL, token string, timeout time.Duration) *Client {
|
||||
func NewClient(baseURL string, tokens TokenProvider, timeout time.Duration) *Client {
|
||||
return &Client{
|
||||
baseURL: strings.TrimRight(baseURL, "/"),
|
||||
token: token,
|
||||
tokens: tokens,
|
||||
timeout: timeout,
|
||||
apiClient: &http.Client{
|
||||
Timeout: timeout,
|
||||
@@ -74,11 +76,19 @@ func origin(u *url.URL) string {
|
||||
return u.Scheme + "://" + u.Host
|
||||
}
|
||||
|
||||
func (c *Client) authHeaders() map[string]string {
|
||||
if c.token == "" {
|
||||
return map[string]string{}
|
||||
func (c *Client) authHeaders() (map[string]string, error) {
|
||||
token, err := c.tokens.Token()
|
||||
if err != nil {
|
||||
return nil, &CoordinatorError{msg: "token refresh failed: " + err.Error()}
|
||||
}
|
||||
return map[string]string{"Authorization": "Bearer " + c.token}
|
||||
if token == "" {
|
||||
return map[string]string{}, nil
|
||||
}
|
||||
return map[string]string{"Authorization": "Bearer " + token}, nil
|
||||
}
|
||||
|
||||
func (c *Client) refreshAndRetry() bool {
|
||||
return c.tokens.Refresh() == nil
|
||||
}
|
||||
|
||||
// Register advertises the worker and returns its identity and heartbeat policy.
|
||||
@@ -209,7 +219,11 @@ func (c *Client) Download(uri, destination string) (string, error) {
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
for name, value := range c.authHeaders() {
|
||||
headers, err := c.authHeaders()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
for name, value := range headers {
|
||||
request.Header.Set(name, value)
|
||||
}
|
||||
response, err := c.dlClient.Do(request)
|
||||
@@ -217,6 +231,9 @@ func (c *Client) Download(uri, destination string) (string, error) {
|
||||
return "", &TransientError{msg: "input download failed"}
|
||||
}
|
||||
defer response.Body.Close()
|
||||
if response.StatusCode == http.StatusUnauthorized && c.refreshAndRetry() {
|
||||
return c.Download(uri, destination)
|
||||
}
|
||||
if response.StatusCode != http.StatusOK {
|
||||
return "", &CoordinatorError{msg: fmt.Sprintf("input download rejected with status %d", response.StatusCode)}
|
||||
}
|
||||
@@ -271,7 +288,12 @@ func (c *Client) Upload(task *Task, workerID string, path, contentType string) (
|
||||
request.Header.Set("Content-Type", contentType)
|
||||
request.Header.Set("X-Worker-ID", workerID)
|
||||
request.Header.Set("X-Task-Attempt", strconv.Itoa(task.Attempt))
|
||||
for name, value := range c.authHeaders() {
|
||||
headers, err := c.authHeaders()
|
||||
if err != nil {
|
||||
file.Close()
|
||||
return nil, err
|
||||
}
|
||||
for name, value := range headers {
|
||||
request.Header.Set(name, value)
|
||||
}
|
||||
response, err := c.apiClient.Do(request)
|
||||
@@ -287,6 +309,9 @@ func (c *Client) Upload(task *Task, workerID string, path, contentType string) (
|
||||
if response.StatusCode == http.StatusConflict {
|
||||
return nil, &ConflictError{msg: "artifact upload rejected because the task lease was lost"}
|
||||
}
|
||||
if response.StatusCode == http.StatusUnauthorized && c.refreshAndRetry() {
|
||||
return c.Upload(task, workerID, path, contentType)
|
||||
}
|
||||
if response.StatusCode != http.StatusOK {
|
||||
return nil, &CoordinatorError{msg: fmt.Sprintf("artifact upload rejected with status %d", response.StatusCode)}
|
||||
}
|
||||
@@ -314,7 +339,11 @@ func (c *Client) requestJSON(method, path string, payload any) (int, map[string]
|
||||
return 0, nil, err
|
||||
}
|
||||
request.Header.Set("Content-Type", "application/json")
|
||||
for name, value := range c.authHeaders() {
|
||||
headers, err := c.authHeaders()
|
||||
if err != nil {
|
||||
return 0, nil, err
|
||||
}
|
||||
for name, value := range headers {
|
||||
request.Header.Set(name, value)
|
||||
}
|
||||
response, err := c.apiClient.Do(request)
|
||||
@@ -326,6 +355,9 @@ func (c *Client) requestJSON(method, path string, payload any) (int, map[string]
|
||||
if err != nil {
|
||||
return 0, nil, &TransientError{msg: "coordinator request interrupted"}
|
||||
}
|
||||
if response.StatusCode == http.StatusUnauthorized && c.refreshAndRetry() {
|
||||
return c.requestJSON(method, path, payload)
|
||||
}
|
||||
if response.StatusCode >= 500 {
|
||||
return response.StatusCode, nil, &TransientError{msg: fmt.Sprintf("coordinator returned %d", response.StatusCode)}
|
||||
}
|
||||
|
||||
@@ -15,7 +15,7 @@ import (
|
||||
|
||||
func newTestClient(t *testing.T, server *httptest.Server) *Client {
|
||||
t.Helper()
|
||||
return NewClient(server.URL, "test-token", 5*time.Second)
|
||||
return NewClient(server.URL, &StaticToken{token: "test-token"}, 5*time.Second)
|
||||
}
|
||||
|
||||
func TestClientRegisterClaimHeartbeat(t *testing.T) {
|
||||
|
||||
@@ -10,10 +10,13 @@ import (
|
||||
"time"
|
||||
)
|
||||
|
||||
// Config is read only from the environment, mirroring the Python worker.
|
||||
// Config is read only from the environment, mirroring the former Python
|
||||
// worker's configuration surface.
|
||||
type Config struct {
|
||||
CoordinatorURL string
|
||||
Token string
|
||||
WorkerKey string
|
||||
UserserviceURL string
|
||||
WorkerName string
|
||||
WorkerID string // set after registration; overridable for tests
|
||||
WorkDir string
|
||||
@@ -22,6 +25,7 @@ type Config struct {
|
||||
PollInterval time.Duration
|
||||
RequestTimeout time.Duration
|
||||
Heartbeat time.Duration
|
||||
CleanupAfter time.Duration // 0 = keep attempt directories
|
||||
Capabilities []string
|
||||
TaskRunner []string // command + args; defaults to python -m scimesh.worker.task
|
||||
MaxTasks int // 0 = unlimited
|
||||
@@ -86,6 +90,10 @@ func LoadConfig() (*Config, error) {
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
cleanup, err := durationEnv("CLEANUP_AFTER_SECONDS", 0)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
capabilities, err := envList("CAPABILITIES")
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -120,6 +128,8 @@ func LoadConfig() (*Config, error) {
|
||||
return &Config{
|
||||
CoordinatorURL: strings.TrimRight(url, "/"),
|
||||
Token: os.Getenv("WORKER_AUTH_TOKEN"),
|
||||
WorkerKey: os.Getenv("WORKER_KEY"),
|
||||
UserserviceURL: strings.TrimRight(os.Getenv("USERSERVICE_URL"), "/"),
|
||||
WorkerName: name,
|
||||
WorkerID: os.Getenv("WORKER_ID"),
|
||||
WorkDir: absWorkDir,
|
||||
@@ -128,6 +138,7 @@ func LoadConfig() (*Config, error) {
|
||||
PollInterval: poll,
|
||||
RequestTimeout: timeout,
|
||||
Heartbeat: heartbeat,
|
||||
CleanupAfter: cleanup,
|
||||
Capabilities: capabilities,
|
||||
TaskRunner: runner,
|
||||
MaxTasks: maxTasks,
|
||||
@@ -141,8 +152,8 @@ func durationEnv(name string, fallback time.Duration) (time.Duration, error) {
|
||||
return fallback, nil
|
||||
}
|
||||
parsed, err := time.ParseDuration(raw)
|
||||
if err != nil || parsed <= 0 {
|
||||
return 0, fmt.Errorf("%s must be a positive duration", name)
|
||||
if err != nil || parsed < 0 {
|
||||
return 0, fmt.Errorf("%s must be a non-negative duration", name)
|
||||
}
|
||||
return parsed, nil
|
||||
}
|
||||
|
||||
@@ -45,6 +45,7 @@ func (d *Daemon) RunForever() error {
|
||||
return err
|
||||
}
|
||||
}
|
||||
d.cleanupExpiredDirectories()
|
||||
outcome, err := d.runOnce()
|
||||
if err != nil {
|
||||
failures++
|
||||
@@ -109,6 +110,41 @@ func (d *Daemon) workerIDOrEmpty() string {
|
||||
return d.workerID
|
||||
}
|
||||
|
||||
// cleanupExpiredDirectories removes task attempt directories older than the
|
||||
// configured retention, mirroring the former Python worker's cleanup.
|
||||
func (d *Daemon) cleanupExpiredDirectories() {
|
||||
if d.config.CleanupAfter <= 0 {
|
||||
return
|
||||
}
|
||||
cutoff := time.Now().Add(-d.config.CleanupAfter)
|
||||
tasks, err := os.ReadDir(d.config.WorkDir)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
for _, taskEntry := range tasks {
|
||||
if !taskEntry.IsDir() {
|
||||
continue
|
||||
}
|
||||
taskDir := filepath.Join(d.config.WorkDir, taskEntry.Name())
|
||||
attempts, err := os.ReadDir(taskDir)
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
for _, attemptEntry := range attempts {
|
||||
if !attemptEntry.IsDir() {
|
||||
continue
|
||||
}
|
||||
info, err := attemptEntry.Info()
|
||||
if err == nil && info.ModTime().Before(cutoff) {
|
||||
_ = os.RemoveAll(filepath.Join(taskDir, attemptEntry.Name()))
|
||||
}
|
||||
}
|
||||
if entries, err := os.ReadDir(taskDir); err == nil && len(entries) == 0 {
|
||||
_ = os.Remove(taskDir)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (d *Daemon) runOnce() (Outcome, error) {
|
||||
workerID := d.workerIDOrEmpty()
|
||||
if workerID == "" {
|
||||
|
||||
@@ -139,7 +139,7 @@ func testDaemon(t *testing.T, fake *fakeCoordinator, script string) *Daemon {
|
||||
Capabilities: []string{"similarity-search"},
|
||||
TaskRunner: []string{script},
|
||||
}
|
||||
client := NewClient(fake.server.URL, "test-token", 5*time.Second)
|
||||
client := NewClient(fake.server.URL, &StaticToken{token: "test-token"}, 5*time.Second)
|
||||
runner := NewTaskRunner(config.TaskRunner)
|
||||
logger := slog.New(slog.NewTextHandler(io.Discard, nil))
|
||||
daemon := NewDaemon(config, client, runner, logger)
|
||||
|
||||
Reference in New Issue
Block a user