diff --git a/.gitignore b/.gitignore index 97b4ca2..3eabd37 100644 --- a/.gitignore +++ b/.gitignore @@ -17,3 +17,4 @@ worker-data*/ scimesh-worker-data/ coordinator/.demo/ site/ +coordinator/bin/ diff --git a/coordinator/Makefile b/coordinator/Makefile index 371037a..7bcfc2b 100644 --- a/coordinator/Makefile +++ b/coordinator/Makefile @@ -1,6 +1,6 @@ .DEFAULT_GOAL := help -.PHONY: help build run test test-integration vet lint tidy check migrate-up migrate-down up down down-clean logs ps rebuild psql smoke workloads-export demo-ui demo-down demo-reset demo-logs +.PHONY: help build run test test-integration vet lint tidy check migrate-up migrate-down up down down-clean logs ps rebuild psql smoke agent workloads-export demo-ui demo-down demo-reset demo-logs # `check` deliberately uses its own Compose project and host ports. This keeps # it from connecting to or replacing a developer's local PostgreSQL instance. @@ -31,6 +31,13 @@ DEMO_DIR ?= .demo # whenever workloads or their manifests change (requires the Python venv). WORKLOADS_JSON := internal/transport/http/workloads.json +# The Go worker agent: a static coordinator client that executes SDK +# workloads in a Python subprocess per claimed task. The Python worker remains +# the reference implementation. +agent: + CGO_ENABLED=0 go build -trimpath -ldflags="-s -w" -o bin/worker-agent ./cmd/worker-agent + @printf '%s\n' 'Built bin/worker-agent. Configure via environment:' ' COORDINATOR_URL, WORKER_AUTH_TOKEN, WORK_DIR, CPU_COUNT, MEMORY_MB,' ' POLL_INTERVAL, REQUEST_TIMEOUT, HEARTBEAT_INTERVAL, CAPABILITIES,' ' TASK_RUNNER, MAX_TASKS, EXIT_WHEN_IDLE, WORKER_NAME, WORKER_ID' + workloads-export: cd .. && .venv/bin/scimesh workload export -o coordinator/$(WORKLOADS_JSON) diff --git a/coordinator/cmd/worker-agent/main.go b/coordinator/cmd/worker-agent/main.go new file mode 100644 index 0000000..c31ed4f --- /dev/null +++ b/coordinator/cmd/worker-agent/main.go @@ -0,0 +1,26 @@ +// Command worker-agent is the Go worker agent: a coordinator client that +// executes SDK workloads in a Python subprocess per claimed task. +package main + +import ( + "log/slog" + "os" + + "github.com/emil28092005/SciMesh/coordinator/internal/agent" +) + +func main() { + config, err := agent.LoadConfig() + if err != nil { + slog.Error("invalid configuration", "error", err) + os.Exit(2) + } + logger := slog.New(slog.NewTextHandler(os.Stderr, nil)) + client := agent.NewClient(config.CoordinatorURL, config.Token, config.RequestTimeout) + runner := agent.NewTaskRunner(config.TaskRunner) + daemon := agent.NewDaemon(config, client, runner, logger) + if err := daemon.RunForever(); err != nil { + logger.Error("agent stopped", "error", err) + os.Exit(1) + } +} diff --git a/coordinator/internal/agent/client.go b/coordinator/internal/agent/client.go new file mode 100644 index 0000000..c171bbf --- /dev/null +++ b/coordinator/internal/agent/client.go @@ -0,0 +1,339 @@ +package agent + +import ( + "bytes" + "crypto/sha256" + "encoding/hex" + "encoding/json" + "fmt" + "io" + "net/http" + "net/url" + "os" + "path/filepath" + "strconv" + "strings" + "time" +) + +// CoordinatorError is a non-retriable coordinator response. +type CoordinatorError struct{ msg string } + +func (e *CoordinatorError) Error() string { return e.msg } + +// TransientError is a timeout, connection error, or 5xx response. +type TransientError struct{ msg string } + +func (e *TransientError) Error() string { return e.msg } + +// ConflictError means the worker no longer owns the task lease. +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. +// +// 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. +type Client struct { + baseURL string + token string + timeout time.Duration + apiClient *http.Client + dlClient *http.Client +} + +func NewClient(baseURL, token string, timeout time.Duration) *Client { + return &Client{ + baseURL: strings.TrimRight(baseURL, "/"), + token: token, + timeout: timeout, + apiClient: &http.Client{ + Timeout: timeout, + CheckRedirect: func(*http.Request, []*http.Request) error { return http.ErrUseLastResponse }, + }, + dlClient: &http.Client{ + Timeout: timeout, + CheckRedirect: func(req *http.Request, via []*http.Request) error { + if len(via) >= 10 { + return fmt.Errorf("too many redirects") + } + // Go strips Authorization on cross-host redirects by default; + // strip it explicitly on any origin change to be safe. + if len(via) > 0 && origin(req.URL) != origin(via[0].URL) { + req.Header.Del("Authorization") + } + return nil + }, + }, + } +} + +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{} + } + return map[string]string{"Authorization": "Bearer " + c.token} +} + +// Register advertises the worker and returns its identity and heartbeat policy. +func (c *Client) Register(name string, capabilities []string, cpuCount int, memoryMB int) (*RegisteredWorker, error) { + payload := map[string]any{ + "name": name, + "capabilities": capabilities, + "cpu_count": cpuCount, + } + if memoryMB > 0 { + payload["memory_mb"] = memoryMB + } + status, body, err := c.requestJSON("POST", "/workers/register", payload) + if err != nil { + return nil, err + } + if status != http.StatusCreated { + return nil, &CoordinatorError{msg: fmt.Sprintf("worker registration rejected with status %d", status)} + } + return ParseRegistered(body) +} + +// Claim leases one compatible task, or returns nil when the queue is empty. +func (c *Client) Claim(workerID string, capabilities []string) (*Task, error) { + status, body, err := c.requestJSON("POST", "/tasks/claim", map[string]any{ + "worker_id": workerID, + "capabilities": capabilities, + "max_concurrency": 1, + }) + if err != nil { + return nil, err + } + if status == http.StatusNoContent { + return nil, nil + } + if status != http.StatusOK { + return nil, &CoordinatorError{msg: fmt.Sprintf("unexpected claim status %d", status)} + } + return ParseTask(body) +} + +// Heartbeat renews the lease and returns the new deadline. +func (c *Client) Heartbeat(task *Task, workerID string) (time.Time, error) { + status, body, err := c.requestJSON("POST", "/tasks/"+task.TaskID+"/heartbeat", map[string]any{ + "worker_id": workerID, + "attempt": task.Attempt, + }) + if err != nil { + return time.Time{}, err + } + if status != http.StatusOK { + if status == http.StatusConflict { + return time.Time{}, &ConflictError{msg: "heartbeat rejected because the task lease was lost"} + } + return time.Time{}, &CoordinatorError{msg: fmt.Sprintf("heartbeat rejected with status %d", status)} + } + raw, ok := body["lease_expires_at"].(string) + if !ok { + return time.Time{}, &CoordinatorError{msg: "heartbeat response is missing lease_expires_at"} + } + lease, err := time.Parse(time.RFC3339, raw) + if err != nil { + return time.Time{}, &CoordinatorError{msg: "heartbeat returned an invalid lease_expires_at"} + } + task.LeaseExpiresAt = lease + task.leaseExpiresRaw = raw + return lease, nil +} + +// Submit completes a task with the uploaded coordinator-owned artifact. +func (c *Client) Submit(task *Task, workerID string, uploaded *Uploaded, metrics map[string]any) error { + status, _, err := c.requestJSON("POST", "/tasks/"+task.TaskID+"/result", map[string]any{ + "worker_id": workerID, + "attempt": task.Attempt, + "result": map[string]any{"artifact_id": uploaded.ArtifactID}, + "metrics": metrics, + }) + if err != nil { + return err + } + if status != http.StatusOK && status != http.StatusCreated && status != http.StatusAccepted { + if status == http.StatusConflict { + return &ConflictError{msg: "result rejected because the task lease was lost"} + } + return &CoordinatorError{msg: fmt.Sprintf("result rejected with status %d", status)} + } + return nil +} + +// Fail reports a sanitized failure. +func (c *Client) Fail(task *Task, workerID string, code, message string, retryable bool) error { + status, _, err := c.requestJSON("POST", "/tasks/"+task.TaskID+"/failure", map[string]any{ + "worker_id": workerID, + "attempt": task.Attempt, + "error_code": code, + "error_message": message, + "retryable": retryable, + }) + if err != nil { + return err + } + if status != http.StatusOK && status != http.StatusCreated && status != http.StatusAccepted { + if status == http.StatusConflict { + return &ConflictError{msg: "failure rejected because the task lease was lost"} + } + return &CoordinatorError{msg: fmt.Sprintf("failure report rejected with status %d", status)} + } + return nil +} + +// Download streams the task input to destination and returns its SHA-256. +func (c *Client) Download(uri, destination string) (string, error) { + resolved, err := url.Parse(uri) + if err != nil { + return "", fmt.Errorf("invalid input URI: %w", err) + } + if !resolved.IsAbs() { + base, parseErr := url.Parse(c.baseURL) + if parseErr != nil { + return "", fmt.Errorf("invalid coordinator URL") + } + resolved = base.ResolveReference(resolved) + } + if resolved.Scheme != "http" && resolved.Scheme != "https" { + return "", fmt.Errorf("input URI must be an HTTP(S) URL") + } + request, err := http.NewRequest(http.MethodGet, resolved.String(), nil) + if err != nil { + return "", err + } + for name, value := range c.authHeaders() { + request.Header.Set(name, value) + } + response, err := c.dlClient.Do(request) + if err != nil { + return "", &TransientError{msg: "input download failed"} + } + defer response.Body.Close() + if response.StatusCode != http.StatusOK { + return "", &CoordinatorError{msg: fmt.Sprintf("input download rejected with status %d", response.StatusCode)} + } + if err := os.MkdirAll(filepath.Dir(destination), 0o755); err != nil { + return "", err + } + target, err := os.OpenFile(destination, os.O_CREATE|os.O_TRUNC|os.O_WRONLY, 0o644) + if err != nil { + return "", err + } + digest := sha256.New() + _, copyErr := io.Copy(io.MultiWriter(target, digest), response.Body) + closeErr := target.Close() + if copyErr != nil { + os.Remove(destination) + return "", &TransientError{msg: "input download interrupted"} + } + if closeErr != nil { + return "", closeErr + } + return hex.EncodeToString(digest.Sum(nil)), nil +} + +// Upload streams a partial artifact and verifies the returned metadata. +func (c *Client) Upload(task *Task, workerID string, path, contentType string) (*Uploaded, error) { + file, err := os.Open(path) + if err != nil { + return nil, err + } + info, err := file.Stat() + if err != nil { + file.Close() + return nil, err + } + digest := sha256.New() + if _, err := io.Copy(digest, file); err != nil { + file.Close() + return nil, err + } + if _, err := file.Seek(0, io.SeekStart); err != nil { + file.Close() + return nil, err + } + localSHA := hex.EncodeToString(digest.Sum(nil)) + uploadURL := c.baseURL + "/tasks/" + url.PathEscape(task.TaskID) + "/artifacts/" + url.PathEscape(filepath.Base(path)) + request, err := http.NewRequest(http.MethodPut, uploadURL, file) + if err != nil { + file.Close() + return nil, err + } + request.ContentLength = info.Size() + 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() { + request.Header.Set(name, value) + } + response, err := c.apiClient.Do(request) + file.Close() + if err != nil { + return nil, &TransientError{msg: "artifact upload failed"} + } + defer response.Body.Close() + raw, err := io.ReadAll(io.LimitReader(response.Body, 1<<20)) + if err != nil { + return nil, &TransientError{msg: "artifact upload interrupted"} + } + if response.StatusCode == http.StatusConflict { + return nil, &ConflictError{msg: "artifact upload rejected because the task lease was lost"} + } + if response.StatusCode != http.StatusOK { + return nil, &CoordinatorError{msg: fmt.Sprintf("artifact upload rejected with status %d", response.StatusCode)} + } + var payload map[string]any + if err := json.Unmarshal(raw, &payload); err != nil { + return nil, &CoordinatorError{msg: "artifact upload returned invalid metadata"} + } + uploaded, err := ParseUploaded(payload) + if err != nil { + return nil, &CoordinatorError{msg: "artifact upload returned invalid metadata"} + } + if uploaded.SHA256 != localSHA || uploaded.SizeBytes != info.Size() { + return nil, &CoordinatorError{msg: "artifact upload metadata does not match local artifact"} + } + return uploaded, nil +} + +func (c *Client) requestJSON(method, path string, payload any) (int, map[string]any, error) { + body, err := json.Marshal(payload) + if err != nil { + return 0, nil, err + } + request, err := http.NewRequest(method, c.baseURL+path, bytes.NewReader(body)) + if err != nil { + return 0, nil, err + } + request.Header.Set("Content-Type", "application/json") + for name, value := range c.authHeaders() { + request.Header.Set(name, value) + } + response, err := c.apiClient.Do(request) + if err != nil { + return 0, nil, &TransientError{msg: "coordinator request failed"} + } + defer response.Body.Close() + raw, err := io.ReadAll(io.LimitReader(response.Body, 1<<20)) + if err != nil { + return 0, nil, &TransientError{msg: "coordinator request interrupted"} + } + if response.StatusCode >= 500 { + return response.StatusCode, nil, &TransientError{msg: fmt.Sprintf("coordinator returned %d", response.StatusCode)} + } + var decoded map[string]any + if len(raw) > 0 { + if err := json.Unmarshal(raw, &decoded); err != nil { + return response.StatusCode, nil, &CoordinatorError{msg: "coordinator returned invalid JSON"} + } + } + return response.StatusCode, decoded, nil +} diff --git a/coordinator/internal/agent/client_test.go b/coordinator/internal/agent/client_test.go new file mode 100644 index 0000000..61dab58 --- /dev/null +++ b/coordinator/internal/agent/client_test.go @@ -0,0 +1,224 @@ +package agent + +import ( + "crypto/sha256" + "encoding/json" + "fmt" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "strings" + "testing" + "time" +) + +func newTestClient(t *testing.T, server *httptest.Server) *Client { + t.Helper() + return NewClient(server.URL, "test-token", 5*time.Second) +} + +func TestClientRegisterClaimHeartbeat(t *testing.T) { + var registered, claimed, heartbeated bool + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Header.Get("Authorization") != "Bearer test-token" { + http.Error(w, "missing token", http.StatusUnauthorized) + return + } + switch { + case r.Method == http.MethodPost && r.URL.Path == "/workers/register": + registered = true + writeJSON(w, http.StatusCreated, map[string]any{ + "worker_id": "22222222-2222-4222-8222-222222222222", + "heartbeat_interval_seconds": 15, + }) + case r.Method == http.MethodPost && r.URL.Path == "/tasks/claim": + claimed = true + writeJSON(w, http.StatusOK, validTaskPayload()) + case r.Method == http.MethodPost && r.URL.Path == "/tasks/11111111-1111-4111-8111-111111111111/heartbeat": + heartbeated = true + writeJSON(w, http.StatusOK, map[string]any{ + "lease_expires_at": time.Now().Add(2 * time.Minute).UTC().Format(time.RFC3339), + }) + default: + http.NotFound(w, r) + } + })) + defer server.Close() + client := newTestClient(t, server) + + registeredWorker, err := client.Register("test-worker", []string{"similarity-search"}, 2, 1024) + if err != nil { + t.Fatalf("Register: %v", err) + } + if registeredWorker.WorkerID != "22222222-2222-4222-8222-222222222222" { + t.Errorf("worker id = %q", registeredWorker.WorkerID) + } + + task, err := client.Claim("22222222-2222-4222-8222-222222222222", []string{"similarity-search"}) + if err != nil { + t.Fatalf("Claim: %v", err) + } + if task == nil || task.Workload != "similarity-search" { + t.Fatalf("claim = %+v", task) + } + + renewed, err := client.Heartbeat(task, "22222222-2222-4222-8222-222222222222") + if err != nil { + t.Fatalf("Heartbeat: %v", err) + } + if renewed.Before(time.Now()) { + t.Error("renewed lease is in the past") + } + if !registered || !claimed || !heartbeated { + t.Error("some endpoints were not hit") + } +} + +func TestClientClaimEmptyAndConflict(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch r.URL.Path { + case "/tasks/claim": + w.WriteHeader(http.StatusNoContent) + case "/tasks/11111111-1111-4111-8111-111111111111/heartbeat": + w.WriteHeader(http.StatusConflict) + default: + http.NotFound(w, r) + } + })) + defer server.Close() + client := newTestClient(t, server) + + task, err := client.Claim("worker", []string{"similarity-search"}) + if err != nil { + t.Fatalf("Claim: %v", err) + } + if task != nil { + t.Error("expected no task for 204") + } + + claimed, err := ParseTask(validTaskPayload()) + if err != nil { + t.Fatalf("ParseTask: %v", err) + } + if _, err := client.Heartbeat(claimed, "worker"); err == nil { + t.Error("expected conflict error") + } else if _, ok := err.(*ConflictError); !ok { + t.Errorf("error type = %T", err) + } +} + +func TestClientUploadSubmitFail(t *testing.T) { + var uploadedPath string + var server *httptest.Server + server = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch { + case r.Method == http.MethodPut && strings.HasPrefix(r.URL.Path, "/tasks/11111111-1111-4111-8111-111111111111/artifacts/"): + if r.Header.Get("X-Worker-ID") != "worker" || r.Header.Get("X-Task-Attempt") != "1" { + t.Errorf("missing identity headers: %+v", r.Header) + } + uploadedPath = r.URL.Path + writeJSON(w, http.StatusOK, map[string]any{ + "artifact_id": "33333333-3333-4333-8333-333333333333", + "uri": server.URL + "/artifacts/333/download", + "sha256": sha256Of(t, "partial body"), + "size_bytes": int64(len("partial body")), + }) + case r.Method == http.MethodPost && strings.HasSuffix(r.URL.Path, "/result"): + writeJSON(w, http.StatusAccepted, map[string]any{}) + case r.Method == http.MethodPost && strings.HasSuffix(r.URL.Path, "/failure"): + writeJSON(w, http.StatusAccepted, map[string]any{}) + default: + http.NotFound(w, r) + } + })) + defer server.Close() + client := newTestClient(t, server) + task, _ := ParseTask(validTaskPayload()) + + dir := t.TempDir() + partial := filepath.Join(dir, "result.csv") + if err := os.WriteFile(partial, []byte("partial body"), 0o644); err != nil { + t.Fatal(err) + } + uploaded, err := client.Upload(task, "worker", partial, "text/csv") + if err != nil { + t.Fatalf("Upload: %v", err) + } + if uploaded.SizeBytes != int64(len("partial body")) { + t.Errorf("size = %d", uploaded.SizeBytes) + } + if !strings.Contains(uploadedPath, "result.csv") { + t.Errorf("upload path = %q", uploadedPath) + } + + if err := client.Submit(task, "worker", uploaded, map[string]any{"rows": 1}); err != nil { + t.Fatalf("Submit: %v", err) + } + if err := client.Fail(task, "worker", "ValueError", "bad input", false); err != nil { + t.Fatalf("Fail: %v", err) + } +} + +func TestClientDownloadVerifiesChecksumAndStripsAuthOnRedirect(t *testing.T) { + var redirectedAuth string + var bucket *httptest.Server + bucket = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + redirectedAuth = r.Header.Get("Authorization") + _, _ = w.Write([]byte("input bytes")) + })) + defer bucket.Close() + var server *httptest.Server + server = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path == "/tasks/11111111-1111-4111-8111-111111111111/input" { + http.Redirect(w, r, bucket.URL+"/presigned", http.StatusFound) + return + } + http.NotFound(w, r) + })) + defer server.Close() + client := newTestClient(t, server) + + destination := filepath.Join(t.TempDir(), "input") + digest, err := client.Download(server.URL+"/tasks/11111111-1111-4111-8111-111111111111/input", destination) + if err != nil { + t.Fatalf("Download: %v", err) + } + if digest != sha256Of(t, "input bytes") { + t.Errorf("digest = %q", digest) + } + if redirectedAuth != "" { + t.Error("Authorization must be stripped on the redirected download") + } +} + +func TestSanitizeErrorMessageRedactsPaths(t *testing.T) { + message := SanitizeErrorMessage( + "failed at /home/alice/work/attempts/1/input and /private/secret.txt", + "/home/alice/work", + ) + for _, forbidden := range []string{"/home/alice", "/private/secret.txt"} { + if strings.Contains(message, forbidden) { + t.Errorf("message leaks %q: %q", forbidden, message) + } + } + if !strings.Contains(message, "") { + t.Errorf("work dir not redacted: %q", message) + } + long := SanitizeErrorMessage(strings.Repeat("x", 500), "/tmp") + if len(long) != 300 { + t.Errorf("truncated length = %d", len(long)) + } +} + +func writeJSON(w http.ResponseWriter, status int, payload any) { + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(status) + _ = json.NewEncoder(w).Encode(payload) +} + +func sha256Of(t *testing.T, value string) string { + t.Helper() + digest := sha256.Sum256([]byte(value)) + return fmt.Sprintf("%x", digest) +} diff --git a/coordinator/internal/agent/config.go b/coordinator/internal/agent/config.go new file mode 100644 index 0000000..c8782c8 --- /dev/null +++ b/coordinator/internal/agent/config.go @@ -0,0 +1,148 @@ +package agent + +import ( + "encoding/json" + "fmt" + "os" + "path/filepath" + "strconv" + "strings" + "time" +) + +// Config is read only from the environment, mirroring the Python worker. +type Config struct { + CoordinatorURL string + Token string + WorkerName string + WorkerID string // set after registration; overridable for tests + WorkDir string + CPUCount int + MemoryMB int // 0 = not advertised + PollInterval time.Duration + RequestTimeout time.Duration + Heartbeat time.Duration + Capabilities []string + TaskRunner []string // command + args; defaults to python -m scimesh.worker.task + MaxTasks int // 0 = unlimited + ExitWhenIdle bool +} + +func envList(name string) ([]string, error) { + raw := os.Getenv(name) + if raw == "" { + return nil, nil + } + var items []string + if err := json.Unmarshal([]byte(raw), &items); err != nil { + return nil, fmt.Errorf("%s must be a JSON array", name) + } + for _, item := range items { + if strings.TrimSpace(item) == "" { + return nil, fmt.Errorf("%s must not contain empty entries", name) + } + } + return items, nil +} + +// LoadConfig validates the environment and fails fast on invalid values. +func LoadConfig() (*Config, error) { + url := os.Getenv("COORDINATOR_URL") + if url == "" { + return nil, fmt.Errorf("COORDINATOR_URL is required") + } + if !strings.HasPrefix(url, "http://") && !strings.HasPrefix(url, "https://") { + return nil, fmt.Errorf("COORDINATOR_URL must be an absolute HTTP(S) URL") + } + workDir := os.Getenv("WORK_DIR") + if workDir == "" { + workDir = "./scimesh-agent-data" + } + cpu := 1 + if raw := os.Getenv("CPU_COUNT"); raw != "" { + parsed, err := strconv.Atoi(raw) + if err != nil || parsed < 1 { + return nil, fmt.Errorf("CPU_COUNT must be a positive integer") + } + cpu = parsed + } + memoryMB := 0 + if raw := os.Getenv("MEMORY_MB"); raw != "" { + parsed, err := strconv.Atoi(raw) + if err != nil || parsed < 1 { + return nil, fmt.Errorf("MEMORY_MB must be a positive integer") + } + memoryMB = parsed + } + poll, err := durationEnv("POLL_INTERVAL", 2*time.Second) + if err != nil { + return nil, err + } + timeout, err := durationEnv("REQUEST_TIMEOUT", 30*time.Second) + if err != nil { + return nil, err + } + heartbeat, err := durationEnv("HEARTBEAT_INTERVAL", 15*time.Second) + if err != nil { + return nil, err + } + capabilities, err := envList("CAPABILITIES") + if err != nil { + return nil, err + } + if len(capabilities) == 0 { + capabilities = []string{"similarity-search", "similarity_search"} + } + runner, err := envList("TASK_RUNNER") + if err != nil { + return nil, err + } + if len(runner) == 0 { + runner = []string{"python", "-m", "scimesh.worker.task"} + } + maxTasks := 0 + if raw := os.Getenv("MAX_TASKS"); raw != "" { + parsed, err := strconv.Atoi(raw) + if err != nil || parsed < 1 { + return nil, fmt.Errorf("MAX_TASKS must be a positive integer") + } + maxTasks = parsed + } + name := os.Getenv("WORKER_NAME") + if name == "" { + host, _ := os.Hostname() + name = host + } + absWorkDir, err := filepath.Abs(workDir) + if err != nil { + return nil, fmt.Errorf("WORK_DIR must be an absolute path") + } + return &Config{ + CoordinatorURL: strings.TrimRight(url, "/"), + Token: os.Getenv("WORKER_AUTH_TOKEN"), + WorkerName: name, + WorkerID: os.Getenv("WORKER_ID"), + WorkDir: absWorkDir, + CPUCount: cpu, + MemoryMB: memoryMB, + PollInterval: poll, + RequestTimeout: timeout, + Heartbeat: heartbeat, + Capabilities: capabilities, + TaskRunner: runner, + MaxTasks: maxTasks, + ExitWhenIdle: os.Getenv("EXIT_WHEN_IDLE") == "1", + }, nil +} + +func durationEnv(name string, fallback time.Duration) (time.Duration, error) { + raw := os.Getenv(name) + if raw == "" { + return fallback, nil + } + parsed, err := time.ParseDuration(raw) + if err != nil || parsed <= 0 { + return 0, fmt.Errorf("%s must be a positive duration", name) + } + return parsed, nil +} diff --git a/coordinator/internal/agent/daemon.go b/coordinator/internal/agent/daemon.go new file mode 100644 index 0000000..5872a66 --- /dev/null +++ b/coordinator/internal/agent/daemon.go @@ -0,0 +1,313 @@ +package agent + +import ( + "crypto/sha256" + "encoding/hex" + "fmt" + "io" + "log/slog" + "os" + "path/filepath" + "strings" + "sync" + "time" +) + +// Outcome reports whether a claim was made and whether it completed. +type Outcome struct { + Claimed bool + Completed bool +} + +// Daemon is the agent state machine: register, claim, execute via the Python +// task runner, upload, and submit — mirroring the Python worker's lifecycle. +type Daemon struct { + config *Config + client *Client + runner *TaskRunner + log *slog.Logger + workerID string + registered bool + completed int + mu sync.Mutex +} + +func NewDaemon(config *Config, client *Client, runner *TaskRunner, log *slog.Logger) *Daemon { + return &Daemon{config: config, client: client, runner: runner, log: log} +} + +// RunForever loops until interrupted, idle-exit, or max tasks. +func (d *Daemon) RunForever() error { + failures := 0 + for { + if !d.registered { + if err := d.register(); err != nil { + return err + } + } + outcome, err := d.runOnce() + if err != nil { + failures++ + d.log.Warn("agent cycle failed", "error", err) + backoff := d.config.PollInterval + for i := 0; i < failures && i < 6; i++ { + backoff *= 2 + } + if backoff > 60*time.Second { + backoff = 60 * time.Second + } + time.Sleep(backoff) + continue + } + failures = 0 + if outcome.Claimed && outcome.Completed { + d.completed++ + if d.config.MaxTasks > 0 && d.completed >= d.config.MaxTasks { + d.log.Info("max tasks reached") + return nil + } + } + if !outcome.Claimed && d.config.ExitWhenIdle { + d.log.Info("queue empty, exiting") + return nil + } + if outcome.Claimed && d.config.ExitWhenIdle { + d.log.Info("one claim processed, exiting") + return nil + } + if !outcome.Claimed { + time.Sleep(d.config.PollInterval) + } + } +} + +func (d *Daemon) register() error { + registered, err := d.client.Register( + d.config.WorkerName, + d.config.Capabilities, + d.config.CPUCount, + d.config.MemoryMB, + ) + if err != nil { + return err + } + d.mu.Lock() + if d.config.WorkerID != "" { + d.workerID = d.config.WorkerID + } else { + d.workerID = registered.WorkerID + } + d.registered = true + d.mu.Unlock() + d.log.Info("registered", "worker_id", d.workerID) + return nil +} + +func (d *Daemon) workerIDOrEmpty() string { + d.mu.Lock() + defer d.mu.Unlock() + return d.workerID +} + +func (d *Daemon) runOnce() (Outcome, error) { + workerID := d.workerIDOrEmpty() + if workerID == "" { + return Outcome{}, fmt.Errorf("agent is not registered") + } + task, err := d.client.Claim(workerID, d.config.Capabilities) + if err != nil { + return Outcome{}, err + } + if task == nil { + return Outcome{Claimed: false}, nil + } + started := time.Now() + taskDir := filepath.Join(d.config.WorkDir, task.TaskID, fmt.Sprint(task.Attempt)) + if err := os.MkdirAll(taskDir, 0o755); err != nil { + return Outcome{Claimed: true}, err + } + heartbeat := newLeaseHeartbeat(task, workerID, d.client, d.config.Heartbeat) + completed := false + err = heartbeat.Start() + if err != nil { + if _, ok := err.(*ConflictError); ok { + d.log.Warn("lease lost", "task_id", task.TaskID) + return Outcome{Claimed: true}, nil + } + return Outcome{Claimed: true}, err + } + defer heartbeat.Stop() + + // Attempt directory cleanup is deliberately minimal in the prototype: + // attempt directories are retained under the work directory. + failure := d.executeTask(task, workerID, taskDir, started, heartbeat) + if failure != nil { + if _, ok := failure.(*ConflictError); ok { + d.log.Warn("lease lost", "task_id", task.TaskID) + return Outcome{Claimed: true}, nil + } + if err := heartbeat.RaiseIfFailed(); err != nil { + return Outcome{Claimed: true}, nil + } + d.reportFailure(task, workerID, failure) + return Outcome{Claimed: true}, nil + } + if err := heartbeat.RaiseIfFailed(); err != nil { + return Outcome{Claimed: true}, nil + } + completed = true + d.log.Info("task completed", "task_id", task.TaskID, "elapsed_seconds", time.Since(started).Seconds()) + return Outcome{Claimed: true, Completed: completed}, nil +} + +// executeTask returns nil on success or a classified failure. +func (d *Daemon) executeTask(task *Task, workerID, taskDir string, started time.Time, heartbeat *leaseHeartbeat) error { + inputPath := filepath.Join(taskDir, "input") + // Downloads use the coordinator-provided URI verbatim; a relative path is + // resolved against the coordinator by the client. + actualSHA, err := d.client.Download(task.Input.URI, inputPath) + if err != nil { + return err + } + if !strings.EqualFold(actualSHA, task.Input.SHA256) { + return &CoordinatorError{msg: "input checksum mismatch"} + } + if err := heartbeat.RaiseIfFailed(); err != nil { + return err + } + manifestPath := filepath.Join(taskDir, "manifest.json") + manifest, err := d.runner.Run(task, taskDir, manifestPath, nil) + if err != nil { + return err + } + if err := heartbeat.RaiseIfFailed(); err != nil { + return err + } + uploaded, err := d.client.Upload(task, workerID, manifest.ArtifactPath, manifest.ContentType) + if err != nil { + return err + } + if err := heartbeat.RaiseIfFailed(); err != nil { + return err + } + metrics := map[string]any{"elapsed_seconds": roundSeconds(time.Since(started).Seconds())} + for name, value := range manifest.Metrics { + metrics[name] = value + } + return d.client.Submit(task, workerID, uploaded, metrics) +} + +func (d *Daemon) reportFailure(task *Task, workerID string, failure error) { + var code string + switch failure.(type) { + case *CoordinatorError: + code = "ValueError" + default: + code = "TaskRunnerFailed" + } + retryable := IsRetryableError(failure) + message := SanitizeErrorMessage(failure.Error(), d.config.WorkDir) + d.log.Warn("task failed", "task_id", task.TaskID, "error_code", code, "retryable", retryable) + if err := d.client.Fail(task, workerID, code, message, retryable); err != nil { + if _, ok := err.(*ConflictError); ok { + d.log.Warn("lease lost while reporting failure", "task_id", task.TaskID) + return + } + d.log.Warn("failure report rejected", "task_id", task.TaskID, "error", err) + } +} + +// leaseHeartbeat renews the lease from the returned deadline at less than +// half of the remaining TTL, mirroring the Python worker. +type leaseHeartbeat struct { + task *Task + workerID string + client *Client + interval time.Duration + stop chan struct{} + once sync.Once + mu sync.Mutex + lease time.Time + failed error +} + +func newLeaseHeartbeat(task *Task, workerID string, client *Client, interval time.Duration) *leaseHeartbeat { + return &leaseHeartbeat{ + task: task, + workerID: workerID, + client: client, + interval: interval, + stop: make(chan struct{}), + lease: task.LeaseExpiresAt, + } +} + +func (h *leaseHeartbeat) Start() error { + if _, err := h.client.Heartbeat(h.task, h.workerID); err != nil { + return err + } + go h.loop() + return nil +} + +func (h *leaseHeartbeat) Stop() { + h.once.Do(func() { close(h.stop) }) +} + +func (h *leaseHeartbeat) RaiseIfFailed() error { + h.mu.Lock() + defer h.mu.Unlock() + return h.failed +} + +func (h *leaseHeartbeat) loop() { + for { + delay := h.nextDelay() + select { + case <-h.stop: + return + case <-time.After(delay): + } + renewed, err := h.client.Heartbeat(h.task, h.workerID) + h.mu.Lock() + if err != nil { + h.failed = err + h.mu.Unlock() + return + } + h.lease = renewed + h.mu.Unlock() + } +} + +func (h *leaseHeartbeat) nextDelay() time.Duration { + h.mu.Lock() + defer h.mu.Unlock() + remaining := time.Until(h.lease) + if remaining <= 0 { + return 0 + } + half := remaining / 2 + if h.interval < half { + return h.interval + } + return half +} + +func roundSeconds(seconds float64) float64 { + return float64(int64(seconds*1000)) / 1000 +} + +// File-digest helper used by tests. +func sha256File(path string) (string, error) { + file, err := os.Open(path) + if err != nil { + return "", err + } + defer file.Close() + digest := sha256.New() + if _, err := io.Copy(digest, file); err != nil { + return "", err + } + return hex.EncodeToString(digest.Sum(nil)), nil +} diff --git a/coordinator/internal/agent/daemon_test.go b/coordinator/internal/agent/daemon_test.go new file mode 100644 index 0000000..700c88d --- /dev/null +++ b/coordinator/internal/agent/daemon_test.go @@ -0,0 +1,296 @@ +package agent + +import ( + "crypto/sha256" + "encoding/json" + "fmt" + "io" + "log/slog" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "strings" + "testing" + "time" +) + +func sha256HexOf(value string) string { + sum := sha256.Sum256([]byte(value)) + return fmt.Sprintf("%x", sum) +} + +// fakeRunnerScript writes a result manifest for --output and exits with the +// given code. +func fakeRunnerScript(t *testing.T, dir string, exitCode int) string { + t.Helper() + script := filepath.Join(dir, "fake-runner.sh") + content := `#!/bin/sh +out="" +task_dir="" +while [ "$#" -gt 0 ]; do + case "$1" in + --output) out="$2"; shift 2;; + --task-dir) task_dir="$2"; shift 2;; + *) shift;; + esac +done +printf 'id,score\n' > "$task_dir/result.csv" +printf '{"artifact_path":"%s/result.csv","content_type":"text/csv","metrics":{"rows":1}}' "$task_dir" > "$out" +exit ` + fmt.Sprint(exitCode) + "\n" + if err := os.WriteFile(script, []byte(content), 0o755); err != nil { + t.Fatal(err) + } + return script +} + +// fakeCoordinator implements the v1 contract over HTTP and records calls. +type fakeCoordinator struct { + server *httptest.Server + task map[string]any + submits []map[string]any + failures []map[string]any + heartbeats int + uploadSHA string + uploadSize int64 + inputBytes []byte + conflict bool // 409 on heartbeat/upload/result +} + +func newFakeCoordinator(t *testing.T, task map[string]any) *fakeCoordinator { + t.Helper() + fake := &fakeCoordinator{task: task, inputBytes: []byte("input fixture")} + fake.uploadSHA = sha256HexOf(string(fake.inputBytes)) + fake.uploadSize = int64(len(fake.inputBytes)) + var server *httptest.Server + server = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch { + case r.Method == http.MethodPost && r.URL.Path == "/workers/register": + writeJSON(w, http.StatusCreated, map[string]any{ + "worker_id": "22222222-2222-4222-8222-222222222222", + "heartbeat_interval_seconds": 15.0, + }) + case r.Method == http.MethodPost && r.URL.Path == "/tasks/claim": + if fake.task == nil { + w.WriteHeader(http.StatusNoContent) + return + } + writeJSON(w, http.StatusOK, fake.task) + case r.Method == http.MethodGet && strings.HasSuffix(r.URL.Path, "/input"): + _, _ = w.Write(fake.inputBytes) + case r.Method == http.MethodPost && strings.HasSuffix(r.URL.Path, "/heartbeat"): + fake.heartbeats++ + if fake.conflict { + w.WriteHeader(http.StatusConflict) + return + } + writeJSON(w, http.StatusOK, map[string]any{ + "lease_expires_at": time.Now().Add(2 * time.Minute).UTC().Format(time.RFC3339), + }) + case r.Method == http.MethodPut && strings.Contains(r.URL.Path, "/artifacts/"): + raw, _ := io.ReadAll(r.Body) + fake.uploadSize = int64(len(raw)) + fake.uploadSHA = fmt.Sprintf("%x", sha256.Sum256(raw)) + if fake.conflict { + w.WriteHeader(http.StatusConflict) + return + } + writeJSON(w, http.StatusOK, map[string]any{ + "artifact_id": "33333333-3333-4333-8333-333333333333", + "uri": server.URL + "/artifacts/333/download", + "sha256": fake.uploadSHA, + "size_bytes": fake.uploadSize, + }) + case r.Method == http.MethodPost && strings.HasSuffix(r.URL.Path, "/result"): + var payload map[string]any + _ = json.NewDecoder(r.Body).Decode(&payload) + fake.submits = append(fake.submits, payload) + if fake.conflict { + w.WriteHeader(http.StatusConflict) + return + } + w.WriteHeader(http.StatusAccepted) + case r.Method == http.MethodPost && strings.HasSuffix(r.URL.Path, "/failure"): + var payload map[string]any + _ = json.NewDecoder(r.Body).Decode(&payload) + fake.failures = append(fake.failures, payload) + w.WriteHeader(http.StatusAccepted) + default: + http.NotFound(w, r) + } + })) + fake.server = server + return fake +} + +func (f *fakeCoordinator) close() { f.server.Close() } + +func testDaemon(t *testing.T, fake *fakeCoordinator, script string) *Daemon { + t.Helper() + config := &Config{ + CoordinatorURL: fake.server.URL, + WorkerName: "test-worker", + WorkerID: "22222222-2222-4222-8222-222222222222", + WorkDir: t.TempDir(), + CPUCount: 1, + PollInterval: time.Millisecond, + RequestTimeout: 5 * time.Second, + Heartbeat: 15 * time.Second, + Capabilities: []string{"similarity-search"}, + TaskRunner: []string{script}, + } + client := NewClient(fake.server.URL, "test-token", 5*time.Second) + runner := NewTaskRunner(config.TaskRunner) + logger := slog.New(slog.NewTextHandler(io.Discard, nil)) + daemon := NewDaemon(config, client, runner, logger) + if err := daemon.register(); err != nil { + t.Fatalf("register: %v", err) + } + return daemon +} + +func validClaimedTaskPayload() map[string]any { + return map[string]any{ + "task_id": "11111111-1111-4111-8111-111111111111", + "attempt": 1.0, + "lease_expires_at": time.Now().Add(time.Minute).UTC().Format(time.RFC3339), + "workload": "similarity-search", + "input": map[string]any{ + "uri": "/tasks/11111111-1111-4111-8111-111111111111/input", + "sha256": sha256HexOf("input fixture"), + }, + "parameters": map[string]any{"query_smiles": "CCO"}, + } +} + +func TestDaemonCompletesAClaimedTask(t *testing.T) { + fake := newFakeCoordinator(t, validClaimedTaskPayload()) + defer fake.close() + daemon := testDaemon(t, fake, fakeRunnerScript(t, t.TempDir(), 0)) + + outcome, err := daemon.runOnce() + if err != nil { + t.Fatalf("runOnce: %v", err) + } + if !outcome.Claimed || !outcome.Completed { + t.Fatalf("outcome = %+v", outcome) + } + if len(fake.submits) != 1 { + t.Fatalf("submits = %d", len(fake.submits)) + } + result := fake.submits[0]["result"].(map[string]any) + if result["artifact_id"] != "33333333-3333-4333-8333-333333333333" { + t.Errorf("result artifact = %v", result) + } + metrics := fake.submits[0]["metrics"].(map[string]any) + if metrics["rows"] != float64(1) { + t.Errorf("metrics = %v", metrics) + } + if _, ok := metrics["elapsed_seconds"].(float64); !ok { + t.Errorf("missing elapsed_seconds: %v", metrics) + } + if fake.heartbeats < 1 { + t.Error("expected at least one heartbeat") + } + if len(fake.failures) != 0 { + t.Errorf("unexpected failures: %v", fake.failures) + } +} + +func TestDaemonReportsChecksumMismatchAsPermanentFailure(t *testing.T) { + payload := validClaimedTaskPayload() + payload["input"].(map[string]any)["sha256"] = strings.Repeat("b", 64) + fake := newFakeCoordinator(t, payload) + defer fake.close() + daemon := testDaemon(t, fake, fakeRunnerScript(t, t.TempDir(), 0)) + + outcome, err := daemon.runOnce() + if err != nil { + t.Fatalf("runOnce: %v", err) + } + if outcome.Completed { + t.Fatal("task must not complete on checksum mismatch") + } + if len(fake.failures) != 1 { + t.Fatalf("failures = %d", len(fake.failures)) + } + failure := fake.failures[0] + if failure["error_code"] != "ValueError" || failure["retryable"] != false { + t.Errorf("failure = %v", failure) + } + if !strings.Contains(failure["error_message"].(string), "checksum") { + t.Errorf("message = %v", failure["error_message"]) + } + if len(fake.submits) != 0 { + t.Error("no submission expected") + } +} + +func TestDaemonReportsPermanentRunnerFailure(t *testing.T) { + fake := newFakeCoordinator(t, validClaimedTaskPayload()) + defer fake.close() + daemon := testDaemon(t, fake, fakeRunnerScript(t, t.TempDir(), ExitPermanent)) + + outcome, err := daemon.runOnce() + if err != nil { + t.Fatalf("runOnce: %v", err) + } + if outcome.Completed { + t.Fatal("task must not complete") + } + if len(fake.failures) != 1 || fake.failures[0]["retryable"] != false { + t.Fatalf("failures = %v", fake.failures) + } +} + +func TestDaemonReportsRetryableRunnerFailure(t *testing.T) { + fake := newFakeCoordinator(t, validClaimedTaskPayload()) + defer fake.close() + daemon := testDaemon(t, fake, fakeRunnerScript(t, t.TempDir(), 1)) + + outcome, err := daemon.runOnce() + if err != nil { + t.Fatalf("runOnce: %v", err) + } + if outcome.Completed { + t.Fatal("task must not complete") + } + if len(fake.failures) != 1 || fake.failures[0]["retryable"] != true { + t.Fatalf("failures = %v", fake.failures) + } +} + +func TestDaemonLeaseConflictStopsWithoutFailureReport(t *testing.T) { + fake := newFakeCoordinator(t, validClaimedTaskPayload()) + fake.conflict = true + defer fake.close() + daemon := testDaemon(t, fake, fakeRunnerScript(t, t.TempDir(), 0)) + + outcome, err := daemon.runOnce() + if err != nil { + t.Fatalf("runOnce: %v", err) + } + if !outcome.Claimed { + t.Fatal("task was claimed") + } + if len(fake.failures) != 0 { + t.Errorf("no failure report expected after lease loss: %v", fake.failures) + } + if len(fake.submits) != 0 { + t.Errorf("no submission expected after lease loss: %v", fake.submits) + } +} + +func TestDaemonIdleClaimIsNotCompleted(t *testing.T) { + fake := newFakeCoordinator(t, nil) + defer fake.close() + daemon := testDaemon(t, fake, fakeRunnerScript(t, t.TempDir(), 0)) + + outcome, err := daemon.runOnce() + if err != nil { + t.Fatalf("runOnce: %v", err) + } + if outcome.Claimed || outcome.Completed { + t.Fatalf("outcome = %+v", outcome) + } +} diff --git a/coordinator/internal/agent/models.go b/coordinator/internal/agent/models.go new file mode 100644 index 0000000..170346b --- /dev/null +++ b/coordinator/internal/agent/models.go @@ -0,0 +1,205 @@ +// Package agent implements a Go worker agent: a coordinator client and +// task-lifecycle supervisor that executes SDK workloads in a Python +// subprocess. It mirrors the Python worker's v1 wire contract exactly; the +// Python worker remains the reference implementation. +package agent + +import ( + "encoding/json" + "fmt" + "net/url" + "regexp" + "strings" + "time" +) + +var ( + uuidPattern = regexp.MustCompile(`^[0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{12}$`) + sha256Pattern = regexp.MustCompile(`^[0-9a-fA-F]{64}$`) + workloadPattern = regexp.MustCompile(`^[a-z][a-z0-9]*(?:[_-][a-z0-9]+)*$`) +) + +// RegisteredWorker is the coordinator's answer to /workers/register. +type RegisteredWorker struct { + WorkerID string + HeartbeatIntervalSeconds float64 +} + +// Input is the claimed task's input artifact. +type Input struct { + URI string + SHA256 string +} + +// Task is one claimed, leased task. +type Task struct { + TaskID string + Attempt int + LeaseExpiresAt time.Time + Workload string + Input Input + Parameters map[string]any + leaseExpiresRaw string +} + +// Uploaded is the coordinator-owned metadata returned after artifact upload. +type Uploaded struct { + ArtifactID string + URI string + SHA256 string + SizeBytes int64 +} + +func requireString(value any, field string) (string, error) { + text, ok := value.(string) + if !ok || strings.TrimSpace(text) == "" { + return "", fmt.Errorf("%s must be a non-empty string", field) + } + return text, nil +} + +func safeCoordinatorURI(value any, field string) (string, error) { + uri, err := requireString(value, field) + if err != nil { + return "", err + } + if strings.HasPrefix(uri, "/") { + // A network-path reference (//host/path) or dot segments would + // resolve to another origin; reject both. + if strings.HasPrefix(uri, "//") { + return "", fmt.Errorf("%s must be a safe coordinator path", field) + } + for _, segment := range strings.Split(uri, "/") { + if segment == ".." { + return "", fmt.Errorf("%s must be a safe coordinator path", field) + } + } + return uri, nil + } + parsed, err := url.Parse(uri) + if err != nil || (parsed.Scheme != "http" && parsed.Scheme != "https") || parsed.Host == "" { + return "", fmt.Errorf("%s must be an absolute HTTP(S) URL or coordinator path", field) + } + return uri, nil +} + +func sha256Hex(value any, field string) (string, error) { + digest, err := requireString(value, field) + if err != nil { + return "", err + } + digest = strings.ToLower(digest) + if !sha256Pattern.MatchString(digest) { + return "", fmt.Errorf("%s must be a SHA-256 hex digest", field) + } + return digest, nil +} + +func uuid(value any, field string) (string, error) { + text, err := requireString(value, field) + if err != nil { + return "", err + } + if !uuidPattern.MatchString(text) { + return "", fmt.Errorf("%s must be a UUID", field) + } + return strings.ToLower(text), nil +} + +// ParseTask validates a claimed-task response with the same strictness as the +// Python worker's ClaimedTask.from_json. +func ParseTask(payload map[string]any) (*Task, error) { + rawInput, ok := payload["input"].(map[string]any) + if !ok { + return nil, fmt.Errorf("input must be an object") + } + rawAttempt, ok := payload["attempt"].(float64) + if !ok || rawAttempt < 1 || rawAttempt != float64(int(rawAttempt)) { + return nil, fmt.Errorf("attempt must be a positive integer") + } + taskID, err := uuid(payload["task_id"], "task_id") + if err != nil { + return nil, fmt.Errorf("invalid claimed-task response: %w", err) + } + rawLease, err := requireString(payload["lease_expires_at"], "lease_expires_at") + if err != nil { + return nil, fmt.Errorf("invalid claimed-task response: %w", err) + } + lease, err := time.Parse(time.RFC3339, rawLease) + if err != nil || lease.Location() == nil { + return nil, fmt.Errorf("lease_expires_at must include a timezone") + } + workload, err := requireString(payload["workload"], "workload") + if err != nil || !workloadPattern.MatchString(workload) { + return nil, fmt.Errorf("workload must be a canonical name") + } + uri, err := safeCoordinatorURI(rawInput["uri"], "input.uri") + if err != nil { + return nil, fmt.Errorf("invalid claimed-task response: %w", err) + } + digest, err := sha256Hex(rawInput["sha256"], "input.sha256") + if err != nil { + return nil, fmt.Errorf("invalid claimed-task response: %w", err) + } + parameters, ok := payload["parameters"].(map[string]any) + if !ok { + parameters = map[string]any{} + } + return &Task{ + TaskID: taskID, + Attempt: int(rawAttempt), + LeaseExpiresAt: lease, + leaseExpiresRaw: rawLease, + Workload: workload, + Input: Input{URI: uri, SHA256: digest}, + Parameters: parameters, + }, nil +} + +// LeaseExpiresRaw returns the original lease timestamp string for +// round-tripping in heartbeat deadlines. +func (t *Task) LeaseExpiresRaw() string { return t.leaseExpiresRaw } + +// ParseRegistered validates a registration response. +func ParseRegistered(payload map[string]any) (*RegisteredWorker, error) { + workerID, err := uuid(payload["worker_id"], "worker_id") + if err != nil { + return nil, fmt.Errorf("invalid worker registration response: %w", err) + } + interval, ok := payload["heartbeat_interval_seconds"].(float64) + if !ok || interval <= 0 { + return nil, fmt.Errorf("heartbeat_interval_seconds must be positive") + } + return &RegisteredWorker{WorkerID: workerID, HeartbeatIntervalSeconds: interval}, nil +} + +// ParseUploaded validates an artifact upload response. +func ParseUploaded(payload map[string]any) (*Uploaded, error) { + artifactID, err := uuid(payload["artifact_id"], "artifact_id") + if err != nil { + return nil, fmt.Errorf("invalid artifact upload response: %w", err) + } + uri, err := safeCoordinatorURI(payload["uri"], "uri") + if err != nil { + return nil, fmt.Errorf("invalid artifact upload response: %w", err) + } + digest, err := sha256Hex(payload["sha256"], "sha256") + if err != nil { + return nil, fmt.Errorf("invalid artifact upload response: %w", err) + } + rawSize, ok := payload["size_bytes"].(float64) + if !ok || rawSize < 0 || rawSize != float64(int64(rawSize)) { + return nil, fmt.Errorf("artifact size_bytes must be a non-negative integer") + } + return &Uploaded{ArtifactID: artifactID, URI: uri, SHA256: digest, SizeBytes: int64(rawSize)}, nil +} + +// TaskRunnerManifest is what the Python task entry writes on success. +type TaskRunnerManifest struct { + ArtifactPath string `json:"artifact_path"` + ContentType string `json:"content_type"` + Metrics map[string]any `json:"metrics"` +} + +// Encode serializes a claim payload for /tasks/claim. +func Encode(v any) ([]byte, error) { return json.Marshal(v) } diff --git a/coordinator/internal/agent/models_test.go b/coordinator/internal/agent/models_test.go new file mode 100644 index 0000000..3fa10e1 --- /dev/null +++ b/coordinator/internal/agent/models_test.go @@ -0,0 +1,111 @@ +package agent + +import ( + "testing" + "time" +) + +func validTaskPayload() map[string]any { + return map[string]any{ + "task_id": "11111111-1111-4111-8111-111111111111", + "attempt": 1.0, + "lease_expires_at": "2026-08-02T00:00:00Z", + "workload": "similarity-search", + "input": map[string]any{ + "uri": "/tasks/11111111-1111-4111-8111-111111111111/input", + "sha256": "aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa", + }, + "parameters": map[string]any{"query_smiles": "CCO"}, + } +} + +func TestParseTaskAcceptsValidPayload(t *testing.T) { + task, err := ParseTask(validTaskPayload()) + if err != nil { + t.Fatalf("ParseTask: %v", err) + } + if task.TaskID != "11111111-1111-4111-8111-111111111111" { + t.Errorf("task id = %q", task.TaskID) + } + if task.Attempt != 1 || task.Workload != "similarity-search" { + t.Errorf("attempt/workload = %d/%q", task.Attempt, task.Workload) + } + if task.Input.SHA256 != "aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa" { + t.Errorf("sha256 = %q", task.Input.SHA256) + } + if task.LeaseExpiresAt.IsZero() { + t.Error("lease must parse") + } +} + +func TestParseTaskRejectsInvalidPayloads(t *testing.T) { + tests := []struct { + name string + mutate func(map[string]any) + }{ + {"non-uuid task id", func(p map[string]any) { p["task_id"] = "../outside" }}, + {"zero attempt", func(p map[string]any) { p["attempt"] = 0 }}, + {"naive lease", func(p map[string]any) { p["lease_expires_at"] = "2026-08-02T00:00:00" }}, + {"network-path uri", func(p map[string]any) { + p["input"].(map[string]any)["uri"] = "//outside.example/input" + }}, + {"dot-segment uri", func(p map[string]any) { + p["input"].(map[string]any)["uri"] = "/tasks/../outside/input" + }}, + {"short sha256", func(p map[string]any) { + p["input"].(map[string]any)["sha256"] = "abc" + }}, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + payload := validTaskPayload() + test.mutate(payload) + if _, err := ParseTask(payload); err == nil { + t.Error("expected ParseTask to reject the payload") + } + }) + } +} + +func TestParseRegisteredAndUploaded(t *testing.T) { + registered, err := ParseRegistered(map[string]any{ + "worker_id": "22222222-2222-4222-8222-222222222222", + "heartbeat_interval_seconds": 15.0, + }) + if err != nil { + t.Fatalf("ParseRegistered: %v", err) + } + if registered.HeartbeatIntervalSeconds != 15 { + t.Errorf("interval = %v", registered.HeartbeatIntervalSeconds) + } + + uploaded, err := ParseUploaded(map[string]any{ + "artifact_id": "33333333-3333-4333-8333-333333333333", + "uri": "https://coordinator.example/artifacts/333/download", + "sha256": "bbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbb", + "size_bytes": 12.0, + }) + if err != nil { + t.Fatalf("ParseUploaded: %v", err) + } + if uploaded.SizeBytes != 12 { + t.Errorf("size = %d", uploaded.SizeBytes) + } + + if _, err := ParseUploaded(map[string]any{"artifact_id": "missing"}); err == nil { + t.Error("expected invalid upload metadata to fail") + } +} + +func TestLeaseHeartbeatDelayIsBelowHalfTTL(t *testing.T) { + task, err := ParseTask(validTaskPayload()) + if err != nil { + t.Fatal(err) + } + task.LeaseExpiresAt = time.Now().Add(60 * time.Second) + heartbeat := newLeaseHeartbeat(task, "worker", nil, 15*time.Second) + delay := heartbeat.nextDelay() + if delay > 30*time.Second || delay <= 0 { + t.Errorf("delay = %v, want < 30s", delay) + } +} diff --git a/coordinator/internal/agent/sanitize.go b/coordinator/internal/agent/sanitize.go new file mode 100644 index 0000000..7972b4a --- /dev/null +++ b/coordinator/internal/agent/sanitize.go @@ -0,0 +1,58 @@ +package agent + +import ( + "fmt" + "os" + "regexp" + "strings" +) + +var ( + // Go's regexp (RE2) has no lookbehind, so these patterns conservatively + // anchor on the characters that typically precede a local path: + // whitespace, quotes, parens, brackets, or the start of the message. + windowsPathPattern = regexp.MustCompile(`[A-Za-z]:\\[^\s'"\],)]+`) + posixPathPattern = regexp.MustCompile(`(^|[\s'"(\[=])/(?:[^\s'"\],)]+)`) +) + +// SanitizeErrorMessage keeps coordinator-visible failures useful without +// exposing local paths. It mirrors the Python worker's sanitizer: local work +// directories and absolute paths are redacted, and the message is truncated +// to 300 characters. +func SanitizeErrorMessage(message string, workDir string) string { + message = strings.ReplaceAll(message, workDir, "") + message = windowsPathPattern.ReplaceAllString(message, "") + message = posixPathPattern.ReplaceAllString(message, "${1}") + if len(message) > 300 { + message = message[:300] + } + return message +} + +// IsRetryableError classifies failures for the coordinator. Invalid scientific +// input and missing local tools are permanent; everything else (transient +// transport errors, unexpected runner failures) may be retried. +func IsRetryableError(err error) bool { + if err == nil { + return false + } + switch err.(type) { + case *CoordinatorError: + return false + case *os.PathError: + return false + } + return true +} + +// TaskRunnerExit classifies subprocess exits. +const ( + ExitPermanent = 3 // runner classified the failure as invalid input +) + +func runnerExitError(exit int, stderr string) error { + if exit == ExitPermanent { + return &CoordinatorError{msg: stderr} + } + return fmt.Errorf("task runner failed with exit code %d: %s", exit, stderr) +} diff --git a/coordinator/internal/agent/taskrunner.go b/coordinator/internal/agent/taskrunner.go new file mode 100644 index 0000000..e362be6 --- /dev/null +++ b/coordinator/internal/agent/taskrunner.go @@ -0,0 +1,81 @@ +package agent + +import ( + "bytes" + "encoding/json" + "fmt" + "os" + "os/exec" + "path/filepath" +) + +// TaskRunner spawns the Python task entry and returns the sealed partial +// artifact manifest it produced. +type TaskRunner struct { + command []string +} + +func NewTaskRunner(command []string) *TaskRunner { + return &TaskRunner{command: command} +} + +// Run executes one task: the task payload is written as JSON into the attempt +// directory, the Python entry computes and seals the partial, and the written +// manifest is parsed back. stderr is captured for failure reporting. +func (r *TaskRunner) Run(task *Task, taskDir string, manifestPath string, extraEnv []string) (*TaskRunnerManifest, error) { + if err := os.MkdirAll(taskDir, 0o755); err != nil { + return nil, err + } + payload := map[string]any{ + "task_id": task.TaskID, + "attempt": task.Attempt, + "lease_expires_at": task.leaseExpiresRaw, + "workload": task.Workload, + "input": map[string]any{ + "uri": task.Input.URI, + "sha256": task.Input.SHA256, + }, + "parameters": task.Parameters, + } + payloadBytes, err := json.Marshal(payload) + if err != nil { + return nil, err + } + taskJSONPath := filepath.Join(taskDir, "task.json") + if err := os.WriteFile(taskJSONPath, payloadBytes, 0o600); err != nil { + return nil, err + } + args := append([]string{}, r.command[1:]...) + args = append(args, + "--task-json", taskJSONPath, + "--task-dir", taskDir, + "--output", manifestPath, + ) + command := exec.Command(r.command[0], args...) + command.Dir = taskDir + command.Env = append(os.Environ(), extraEnv...) + var stderr bytes.Buffer + command.Stderr = &stderr + if err := command.Run(); err != nil { + if exitErr, ok := err.(*exec.ExitError); ok { + return nil, runnerExitError(exitErr.ExitCode(), stderr.String()) + } + return nil, fmt.Errorf("task runner could not be started: %w", err) + } + raw, err := os.ReadFile(manifestPath) + if err != nil { + return nil, fmt.Errorf("task runner produced no result manifest") + } + var manifest TaskRunnerManifest + if err := json.Unmarshal(raw, &manifest); err != nil { + return nil, fmt.Errorf("task runner produced an invalid result manifest") + } + if manifest.ArtifactPath == "" || manifest.ContentType == "" { + return nil, fmt.Errorf("task runner produced an incomplete result manifest") + } + info, err := os.Stat(manifest.ArtifactPath) + if err != nil || !info.Mode().IsRegular() { + return nil, fmt.Errorf("task runner produced no artifact file") + } + return &manifest, nil +} diff --git a/docs/sdk-handoff.md b/docs/sdk-handoff.md index eb778f0..b639f20 100644 --- a/docs/sdk-handoff.md +++ b/docs/sdk-handoff.md @@ -28,6 +28,8 @@ resolves `query_id` per task and rejects plan-time `max_rows`. **MkDocs site (2026-08-02):** the standalone documentation site lives in `mkdocs/` (`docs_dir: mkdocs`) and does not use the project's `docs/` directory. It contains guides (`mkdocs/sdk/`: overview, authoring-workloads, cli, worker-integration), the full auto-generated API reference for all `scimesh.sdk` modules (`mkdocs/api/`, mkdocstrings `::: scimesh.sdk.` — set `show_if_no_docstring: true`), and the writing rules (`mkdocs/approach.md`). `make docs` builds it; the UI serves it at `/ui/docs/`. All public SDK members now carry Google-style docstrings. +**Go worker agent prototype (2026-08-02):** `coordinator/internal/agent/` (config, models, client, sanitize, taskrunner, daemon) + `coordinator/cmd/worker-agent`, built with `make agent`. It mirrors the Python worker's v1 lifecycle; per-task SDK execution happens in a Python subprocess (`scimesh/worker/task.py`: exits 0 on success, 3 permanent, 1 retryable). Default `TASK_RUNNER` is `python -m scimesh.worker.task` — set it to the venv python in source checkouts. Verified E2E against the demo coordinator. Open items: JWT refresh, resource slots/limits, attempt-dir cleanup, protocol-v2 features. + **Authoring scaffold (2026-08-01):** `scimesh/sdk/batch.py` adds `MapReduceWorkload` — the primary authoring surface for `core-batch-v1`. A subclass declares identity/parameters/ports and three scientific hooks diff --git a/scimesh/worker/task.py b/scimesh/worker/task.py new file mode 100644 index 0000000..4b27b6b --- /dev/null +++ b/scimesh/worker/task.py @@ -0,0 +1,71 @@ +"""Per-task entry point executed by the Go worker agent. + +Reads a claimed task from JSON, runs the SDK-built workload with the same +bridge the Python daemon uses (``SciMeshRunner``), and writes a result +manifest consumed by the agent: + +- on success: ``{artifact_path, content_type, metrics}`` and exit 0; +- on invalid input/validation failures: a sanitized message on stderr and + exit 3 (permanent, non-retryable); +- on any other failure: exit 1 (retryable by the agent). + +The agent passes task parameters through unchanged; all scientific policy +lives in the workload. +""" + +from __future__ import annotations + +import argparse +import json +import sys +from pathlib import Path + +from .models import ClaimedTask +from .runners import SciMeshRunner + +_EXIT_RETRYABLE = 1 +_EXIT_PERMANENT = 3 + + +def _load_task(path: Path) -> ClaimedTask: + try: + payload = json.loads(path.read_text(encoding="utf-8")) + except (OSError, ValueError) as error: + raise ValueError("task payload is unreadable") from error + return ClaimedTask.from_json(payload) + + +def main(argv: list[str] | None = None) -> int: + parser = argparse.ArgumentParser(prog="scimesh.worker.task") + parser.add_argument("--task-json", required=True, type=Path) + parser.add_argument("--task-dir", required=True, type=Path) + parser.add_argument("--output", required=True, type=Path) + args = parser.parse_args(argv) + try: + task = _load_task(args.task_json) + task_dir = args.task_dir.resolve() + task_dir.mkdir(parents=True, exist_ok=True) + result = SciMeshRunner().run(task, task_dir) + if len(result.artifacts) != 1: + raise ValueError("runner must produce exactly one result artifact") + artifact = result.artifacts[0] + if not artifact.path.is_file(): + raise ValueError("runner artifact is missing") + manifest = { + "artifact_path": str(artifact.path), + "content_type": artifact.content_type, + "metrics": dict(result.metrics), + } + args.output.parent.mkdir(parents=True, exist_ok=True) + args.output.write_text(json.dumps(manifest), encoding="utf-8") + return 0 + except ValueError as error: + print(f"permanent: {error}", file=sys.stderr) + return _EXIT_PERMANENT + except Exception as error: # noqa: BLE001 - exit code is the contract + print(f"retryable: {error}", file=sys.stderr) + return _EXIT_RETRYABLE + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/tests/test_worker_task.py b/tests/test_worker_task.py new file mode 100644 index 0000000..e8d8253 --- /dev/null +++ b/tests/test_worker_task.py @@ -0,0 +1,129 @@ +"""Tests for the per-task entry point used by the Go worker agent.""" + +from __future__ import annotations + +import hashlib +import json +import subprocess +import sys +from pathlib import Path + +import pytest + +from scimesh.worker.task import _EXIT_PERMANENT, _EXIT_RETRYABLE + + +def _task_json(path: Path, *, query_smiles: str = "CCO") -> dict[str, object]: + return { + "task_id": "11111111-1111-4111-8111-111111111111", + "attempt": 1, + "lease_expires_at": "2026-08-02T00:00:00Z", + "workload": "similarity-search", + "input": { + "uri": "/tasks/11111111-1111-4111-8111-111111111111/input", + "sha256": "a" * 64, + }, + "parameters": {"query_smiles": query_smiles, "top_k": 5}, + } + + +def _run_task(root: Path, task: dict[str, object]) -> subprocess.CompletedProcess[str]: + task_path = root / "task.json" + task_path.write_text(json.dumps(task), encoding="utf-8") + output = root / "manifest.json" + return subprocess.run( + [ + sys.executable, + "-m", + "scimesh.worker.task", + "--task-json", + str(task_path), + "--task-dir", + str(root), + "--output", + str(output), + ], + capture_output=True, + text=True, + ) + + +def test_task_runner_writes_a_result_manifest(tmp_path: Path) -> None: + content = b"chembl_id\tcanonical_smiles\nQUERY\tCCO\nMATCH\tCCCO\n" + (tmp_path / "input").write_bytes(content) + result = _run_task(tmp_path, _task_json(tmp_path)) + + assert result.returncode == 0, result.stderr + manifest = json.loads((tmp_path / "manifest.json").read_text(encoding="utf-8")) + assert manifest["content_type"] == "text/csv" + assert manifest["metrics"]["matches_emitted"] == 1 + artifact = Path(manifest["artifact_path"]) + assert artifact.is_file() + header = artifact.read_text(encoding="utf-8").splitlines()[0] + assert header == "rank,chembl_id,canonical_smiles,similarity" + + +def test_task_runner_exits_permanent_on_invalid_input(tmp_path: Path) -> None: + (tmp_path / "input").write_text( + "chembl_id\tcanonical_smiles\nQUERY\tCCO\n", encoding="utf-8" + ) + result = _run_task(tmp_path, _task_json(tmp_path, query_smiles="not-a-smiles")) + assert result.returncode == _EXIT_PERMANENT + assert "permanent:" in result.stderr + assert not (tmp_path / "manifest.json").exists() + + +def test_task_runner_exits_retryable_on_unexpected_failure( + tmp_path: Path, +) -> None: + content = b"chembl_id\tcanonical_smiles\nQUERY\tCCO\n" + (tmp_path / "input").write_bytes(content) + task_path = tmp_path / "task.json" + task_path.write_text(json.dumps(_task_json(tmp_path)), encoding="utf-8") + output = tmp_path / "manifest.json" + args = [ + "--task-json", + str(task_path), + "--task-dir", + str(tmp_path), + "--output", + str(output), + ] + code = ( + "import sys\n" + "import scimesh.worker.task as t\n" + "from scimesh.worker.runners import SciMeshRunner\n" + "def boom(self, task, task_dir):\n" + " raise RuntimeError('simulated')\n" + "SciMeshRunner.run = boom\n" + f"sys.exit(t.main({args!r}))\n" + ) + result = subprocess.run( + [sys.executable, "-c", code], + capture_output=True, + text=True, + ) + assert result.returncode == _EXIT_RETRYABLE + assert "retryable:" in result.stderr + assert not output.exists() + + +def test_task_runner_rejects_malformed_task_payload(tmp_path: Path) -> None: + task_path = tmp_path / "task.json" + task_path.write_text('{"broken":', encoding="utf-8") + result = subprocess.run( + [ + sys.executable, + "-m", + "scimesh.worker.task", + "--task-json", + str(task_path), + "--task-dir", + str(tmp_path), + "--output", + str(tmp_path / "manifest.json"), + ], + capture_output=True, + text=True, + ) + assert result.returncode == _EXIT_PERMANENT