Replace Python worker daemon with the Go worker agent

This commit is contained in:
Emil
2026-08-02 16:53:04 +03:00
parent 9a8221163a
commit 706bc85e17
31 changed files with 516 additions and 2399 deletions
+7 -1
View File
@@ -16,7 +16,13 @@ func main() {
os.Exit(2)
}
logger := slog.New(slog.NewTextHandler(os.Stderr, nil))
client := agent.NewClient(config.CoordinatorURL, config.Token, config.RequestTimeout)
tokens := agent.NewTokenProvider(
config.WorkerKey,
config.UserserviceURL,
config.Token,
config.RequestTimeout,
)
client := agent.NewClient(config.CoordinatorURL, tokens, config.RequestTimeout)
runner := agent.NewTaskRunner(config.TaskRunner)
daemon := agent.NewDaemon(config, client, runner, logger)
if err := daemon.RunForever(); err != nil {
+120
View File
@@ -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}
}
+102
View File
@@ -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())
}
}
+43 -11
View File
@@ -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)}
}
+1 -1
View File
@@ -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) {
+14 -3
View File
@@ -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
}
+36
View File
@@ -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 == "" {
+1 -1
View File
@@ -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)
+26 -13
View File
@@ -28,7 +28,7 @@ case "$demo_dir" in
/*) ;;
*) demo_dir="$coordinator_dir/$demo_dir" ;;
esac
worker_bin=${SCIMESH_WORKER_BIN:-"$repo_dir/.venv/bin/scimesh-worker"}
agent_bin=${SCIMESH_AGENT_BIN:-"$coordinator_dir/bin/worker-agent"}
pid_file="$demo_dir/workers.pids"
logs_dir="$demo_dir/logs"
@@ -79,13 +79,23 @@ stop_workers() {
[[ "$pid" =~ ^[0-9]+$ ]] || continue
command_line=$(ps -p "$pid" -o args= 2>/dev/null || true)
# Never kill a recycled PID or a worker launched outside this demo.
if [[ "$command_line" == *"$demo_dir/worker-"* ]]; then
if [[ "$command_line" == *"worker-agent"* ]]; then
kill "$pid" 2>/dev/null || true
fi
done < "$pid_file"
rm -f "$pid_file"
}
build_agent() {
if [[ ! -x "$agent_bin" ]]; then
echo "Building the Go worker agent..." >&2
make -C "$coordinator_dir" agent >&2 || {
echo "Failed to build the Go worker agent." >&2
exit 2
}
fi
}
wait_for_coordinator() {
local attempt=0
until curl --fail --silent --show-error "http://localhost:$coordinator_port/health" >/dev/null; do
@@ -145,11 +155,7 @@ start() {
echo "DEMO_WORKERS must be a positive integer (got $workers)." >&2
exit 2
fi
if [[ ! -x "$worker_bin" ]]; then
echo "Reference worker not found: $worker_bin" >&2
echo "Create it first from the repository root: python3 -m venv .venv && .venv/bin/pip install -e '.[dev]'" >&2
exit 2
fi
build_agent
command -v docker >/dev/null || { echo "Docker is required." >&2; exit 2; }
command -v curl >/dev/null || { echo "curl is required." >&2; exit 2; }
@@ -162,14 +168,21 @@ start() {
wait_for_userservice
: > "$pid_file"
task_runner="[\"$repo_dir/.venv/bin/python\",\"-m\",\"scimesh.worker.task\"]"
for index in $(seq 1 "$workers"); do
work_dir="$demo_dir/worker-$index"
mkdir -p "$work_dir"
SCIMESH_COORDINATOR_URL="http://localhost:$coordinator_port" \
SCIMESH_BEARER_TOKEN="$worker_token" \
"$worker_bin" \
--worker-name "demo-worker-$index" \
--work-dir "$work_dir" \
COORDINATOR_URL="http://localhost:$coordinator_port" \
WORKER_AUTH_TOKEN="$worker_token" \
WORKER_NAME="demo-worker-$index" \
WORK_DIR="$work_dir" \
CPU_COUNT=1 \
MEMORY_MB=1024 \
POLL_INTERVAL=0.5s \
REQUEST_TIMEOUT=15s \
HEARTBEAT_INTERVAL=15s \
TASK_RUNNER="$task_runner" \
"$agent_bin" \
>"$logs_dir/worker-$index.log" 2>&1 &
echo "$!" >> "$pid_file"
done
@@ -184,7 +197,7 @@ SciMesh manual demo is ready.
Userservice: http://localhost:$userservice_port
Grafana: http://localhost:$grafana_port (anonymous view; admin/${GRAFANA_PASSWORD:-admin} to edit)
Prometheus: http://localhost:$prometheus_port
Workers: $workers local reference workers
Workers: $workers Go worker agents (Python task execution)
Sign in with the admin above, or register a new account from the login page.
The admin sees every job; a plain user sees only their own. Upload a small