Merge branch 'fix/coordinator-worker-integration'
This commit is contained in:
@@ -1,7 +1,7 @@
|
|||||||
# SciMesh Status
|
# SciMesh Status
|
||||||
|
|
||||||
**Updated:** 2026-07-23
|
**Updated:** 2026-07-23
|
||||||
**Branch baseline:** `planning` at `13f9a0b`
|
**Branch baseline:** `main` at `b4a89dd` (coordinator merge)
|
||||||
|
|
||||||
## Current state
|
## Current state
|
||||||
|
|
||||||
@@ -15,39 +15,43 @@ the reference behaviour for future distributed execution:
|
|||||||
- Python Worker skeleton: claim, heartbeat, input checksum validation,
|
- Python Worker skeleton: claim, heartbeat, input checksum validation,
|
||||||
artifact upload, completion and failure reporting.
|
artifact upload, completion and failure reporting.
|
||||||
|
|
||||||
The Go coordinator, PostgreSQL schema, coordinator artifact storage, planner,
|
The Go coordinator and its PostgreSQL-backed task lifecycle are implemented:
|
||||||
reducer, and end-to-end distributed execution are **not implemented yet**.
|
registration, atomic claiming, lease renewal, artifact storage, dataset
|
||||||
|
chunking, result/failure reporting, and job progress. The Python worker now
|
||||||
|
uses the live coordinator contract; its HTTP path was exercised against a real
|
||||||
|
Docker PostgreSQL stack on 2026-07-23.
|
||||||
|
|
||||||
## Milestone tracker
|
## Milestone tracker
|
||||||
|
|
||||||
| CTX | Status | Notes |
|
| CTX | Status | Notes |
|
||||||
| --- | --- | --- |
|
| --- | --- | --- |
|
||||||
| CTX-00 API and error contract | Ready to implement | `docs/api-contract.md` created; needs owner review/freeze. |
|
| CTX-00 API and error contract | Implemented | Contract, OpenAPI, and request examples are in `docs/`. |
|
||||||
| CTX-01 Go coordinator bootstrap | Not started | Depends on CTX-00. |
|
| CTX-01 Go coordinator bootstrap | Implemented | Go service and Docker runtime in `coordinator/`. |
|
||||||
| CTX-02 PostgreSQL migrations | Not started | Depends on CTX-00 and CTX-01. |
|
| CTX-02 PostgreSQL migrations | Implemented | Applied by the Compose migration service. |
|
||||||
| CTX-03 Transactional queue | Not started | Depends on CTX-02. |
|
| CTX-03 Transactional queue | Implemented | Real-PostgreSQL integration tests cover atomic claims and concurrency. |
|
||||||
| CTX-04 Worker registry and HTTP API | Not started | Depends on CTX-03. |
|
| CTX-04 Worker registry and HTTP API | Implemented | Registration, claim, heartbeat, result, failure, and status endpoints. |
|
||||||
| CTX-05 Artifact storage | Not started | Depends on CTX-02 and CTX-04. |
|
| CTX-05 Artifact storage | Implemented | Coordinator-owned inputs/results, checksum verification, and upload flow. |
|
||||||
| CTX-06 Python Worker live-contract alignment | Partially prepared | Worker skeleton exists; needs real Go contract tests. |
|
| CTX-06 Python Worker live-contract alignment | Implemented | Worker completed a real uploaded shard via HTTP on 2026-07-23. |
|
||||||
| CTX-07 Distributed workload protocol | Not started | Depends on artifact and Worker contracts. |
|
| CTX-07 Distributed workload protocol | Not started | Depends on artifact and Worker contracts. |
|
||||||
| CTX-08 Distributed similarity-search | Not started | Local reference exists. |
|
| CTX-08 Distributed similarity-search | Not started | Local reference exists. |
|
||||||
| CTX-09 Reducer and final-result API | Not started | Depends on CTX-07 and CTX-08. |
|
| CTX-09 Reducer and final-result API | Not started | Depends on CTX-07 and CTX-08. |
|
||||||
| CTX-10 Distributed similarity-graph | Not started | Local reference exists. |
|
| CTX-10 Distributed similarity-graph | Not started | Local reference exists. |
|
||||||
| CTX-11 Dashboard/operator view | Not started | Deferred until API and reducer work. |
|
| CTX-11 Dashboard/operator view | Not started | Deferred until API and reducer work. |
|
||||||
| CTX-12 Reliability, security, CI | Not started | Final milestone. |
|
| CTX-12 Reliability, security, CI | In progress | Unit, race, PostgreSQL integration, and smoke checks exist; CI hardening remains. |
|
||||||
|
|
||||||
## Next recommended assignment
|
## Next recommended assignment
|
||||||
|
|
||||||
Assign **CTX-00** to the coordinator role in `.agents/coordinator.md`: review
|
Assign **CTX-07** to the workload role: define distributed job planning and
|
||||||
and freeze `docs/api-contract.md` against `PLAN.md`. Do not begin coordinator
|
reduction boundaries before implementing distributed search or graph execution.
|
||||||
or Worker API implementation until the contract owner accepts it.
|
|
||||||
|
|
||||||
## Known constraints
|
## Known constraints
|
||||||
|
|
||||||
- Distributed execution is not available; use the local `scimesh` CLI.
|
- Planner/reducer semantics are not implemented; use the local `scimesh` CLI
|
||||||
- No Go module, PostgreSQL migrations, runtime configuration, or integration
|
for complete workload results.
|
||||||
environment exists yet.
|
- The worker/coordinator flow currently accepts both underscore API workload
|
||||||
- Local worker unit tests do not prove interoperability with a live coordinator.
|
names and hyphenated CLI names while the contract is consolidated.
|
||||||
|
- A real-stack worker test uses a small `query_smiles` shard. Resolving a
|
||||||
|
`query_id` once and sharing it across shards belongs to CTX-07.
|
||||||
|
|
||||||
## Update rule
|
## Update rule
|
||||||
|
|
||||||
|
|||||||
+18
-4
@@ -1,5 +1,15 @@
|
|||||||
.PHONY: build run test test-integration vet lint tidy check migrate-up migrate-down up down down-clean logs ps rebuild psql smoke
|
.PHONY: build run test test-integration vet lint tidy check migrate-up migrate-down up down down-clean logs ps rebuild psql smoke
|
||||||
|
|
||||||
|
# `check` deliberately uses its own Compose project and host ports. This keeps
|
||||||
|
# it from connecting to or replacing a developer's local PostgreSQL instance.
|
||||||
|
CHECK_PROJECT ?= scimesh-check
|
||||||
|
CHECK_POSTGRES_PORT ?= 55432
|
||||||
|
CHECK_COORDINATOR_PORT ?= 18080
|
||||||
|
CHECK_HOST ?= http://localhost:$(CHECK_COORDINATOR_PORT)
|
||||||
|
CHECK_TOKEN ?= dev-token
|
||||||
|
CHECK_DATABASE_URL ?= postgres://scimesh:scimesh@localhost:$(CHECK_POSTGRES_PORT)/scimesh?sslmode=disable
|
||||||
|
CHECK_COMPOSE = POSTGRES_PORT=$(CHECK_POSTGRES_PORT) COORDINATOR_PORT=$(CHECK_COORDINATOR_PORT) docker compose -p $(CHECK_PROJECT)
|
||||||
|
|
||||||
# --- build / run ---------------------------------------------------------
|
# --- build / run ---------------------------------------------------------
|
||||||
build:
|
build:
|
||||||
go build ./...
|
go build ./...
|
||||||
@@ -23,12 +33,16 @@ vet:
|
|||||||
# Needs Docker. Hand this to a reviewer.
|
# Needs Docker. Hand this to a reviewer.
|
||||||
check: vet lint
|
check: vet lint
|
||||||
go test -race ./...
|
go test -race ./...
|
||||||
docker compose up -d --build
|
$(CHECK_COMPOSE) up -d --build
|
||||||
@echo "waiting for the coordinator to be ready..."
|
@echo "waiting for the coordinator to be ready..."
|
||||||
@sleep 6
|
@attempt=0; until curl -fsS "$(CHECK_HOST)/health" >/dev/null; do \
|
||||||
TEST_DATABASE_URL='postgres://scimesh:scimesh@localhost:5432/scimesh?sslmode=disable' \
|
attempt=$$((attempt + 1)); \
|
||||||
|
if [ $$attempt -ge 30 ]; then $(CHECK_COMPOSE) logs coordinator; exit 1; fi; \
|
||||||
|
sleep 1; \
|
||||||
|
done
|
||||||
|
TEST_DATABASE_URL="$(CHECK_DATABASE_URL)" \
|
||||||
go test -tags=integration ./internal/storage/postgres/ -v
|
go test -tags=integration ./internal/storage/postgres/ -v
|
||||||
./scripts/smoke.sh
|
HOST="$(CHECK_HOST)" TOKEN="$(CHECK_TOKEN)" ./scripts/smoke.sh
|
||||||
@echo "\nall checks passed ✓"
|
@echo "\nall checks passed ✓"
|
||||||
|
|
||||||
# Runs golangci-lint without installing it system-wide. Install it for speed:
|
# Runs golangci-lint without installing it system-wide. Install it for speed:
|
||||||
|
|||||||
@@ -169,3 +169,14 @@ make test-integration TEST_DATABASE_URL='postgres://scimesh:scimesh@localhost:54
|
|||||||
|
|
||||||
CI (`.github/workflows/coordinator.yml`) runs vet, gofmt, race tests, lint, and
|
CI (`.github/workflows/coordinator.yml`) runs vet, gofmt, race tests, lint, and
|
||||||
the integration suite against a Postgres service on every push and PR.
|
the integration suite against a Postgres service on every push and PR.
|
||||||
|
|
||||||
|
For the complete local verification, including an isolated Docker PostgreSQL
|
||||||
|
and the HTTP smoke flow, run:
|
||||||
|
|
||||||
|
```sh
|
||||||
|
make check
|
||||||
|
```
|
||||||
|
|
||||||
|
It uses Compose project `scimesh-check` and ports `55432`/`18080` by default,
|
||||||
|
so it does not connect to a PostgreSQL already running on `5432`. Override
|
||||||
|
`CHECK_POSTGRES_PORT`, `CHECK_COORDINATOR_PORT`, or `CHECK_PROJECT` if needed.
|
||||||
|
|||||||
@@ -23,6 +23,7 @@ type Artifact struct {
|
|||||||
ID uuid.UUID
|
ID uuid.UUID
|
||||||
JobID uuid.UUID
|
JobID uuid.UUID
|
||||||
TaskID *uuid.UUID // nil for a job-level input
|
TaskID *uuid.UUID // nil for a job-level input
|
||||||
|
Attempt *int // required for a partial result; nil for non-worker artifacts
|
||||||
Kind ArtifactKind
|
Kind ArtifactKind
|
||||||
Filename string
|
Filename string
|
||||||
StorageKey string
|
StorageKey string
|
||||||
|
|||||||
@@ -117,8 +117,10 @@ const DefaultMaxAttempts = 3
|
|||||||
func (t *Task) CanRetry() bool { return t.Attempt < t.MaxAttempts }
|
func (t *Task) CanRetry() bool { return t.Attempt < t.MaxAttempts }
|
||||||
|
|
||||||
// IsLeaseHeldBy reports whether worker currently holds this task at attempt.
|
// IsLeaseHeldBy reports whether worker currently holds this task at attempt.
|
||||||
func (t *Task) IsLeaseHeldBy(worker string, attempt int) bool {
|
func (t *Task) IsLeaseHeldBy(worker string, attempt int, now time.Time) bool {
|
||||||
return t.LeaseOwner != nil && *t.LeaseOwner == worker && t.Attempt == attempt
|
return t.LeaseOwner != nil && t.LeaseExpiresAt != nil && now.Before(*t.LeaseExpiresAt) &&
|
||||||
|
*t.LeaseOwner == worker && t.Attempt == attempt &&
|
||||||
|
(t.Status == TaskLeased || t.Status == TaskRunning)
|
||||||
}
|
}
|
||||||
|
|
||||||
// AsClaimed projects the task into the trimmed view handed to a worker:
|
// AsClaimed projects the task into the trimmed view handed to a worker:
|
||||||
@@ -146,7 +148,7 @@ func (t *Task) AsClaimed() ClaimedTask {
|
|||||||
|
|
||||||
// verifyLease is the guard every worker-driven transition shares: the caller
|
// verifyLease is the guard every worker-driven transition shares: the caller
|
||||||
// must own the lease and reference the attempt it was granted.
|
// must own the lease and reference the attempt it was granted.
|
||||||
func (t *Task) verifyLease(worker string, attempt int) error {
|
func (t *Task) verifyLease(worker string, attempt int, now time.Time) error {
|
||||||
// A task is worker-owned while leased or running: the first heartbeat moves
|
// A task is worker-owned while leased or running: the first heartbeat moves
|
||||||
// it from leased to running, but ownership rules are identical for both.
|
// it from leased to running, but ownership rules are identical for both.
|
||||||
if t.Status != TaskLeased && t.Status != TaskRunning {
|
if t.Status != TaskLeased && t.Status != TaskRunning {
|
||||||
@@ -158,13 +160,16 @@ func (t *Task) verifyLease(worker string, attempt int) error {
|
|||||||
if t.Attempt != attempt {
|
if t.Attempt != attempt {
|
||||||
return ErrStaleAttempt
|
return ErrStaleAttempt
|
||||||
}
|
}
|
||||||
|
if t.LeaseExpiresAt == nil || !now.Before(*t.LeaseExpiresAt) {
|
||||||
|
return ErrLeaseConflict
|
||||||
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// RenewLease extends the lease of the worker that holds it. The first heartbeat
|
// RenewLease extends the lease of the worker that holds it. The first heartbeat
|
||||||
// also acknowledges start, moving the task from leased to running.
|
// also acknowledges start, moving the task from leased to running.
|
||||||
func (t *Task) RenewLease(worker string, attempt int, until time.Time) error {
|
func (t *Task) RenewLease(worker string, attempt int, now, until time.Time) error {
|
||||||
if err := t.verifyLease(worker, attempt); err != nil {
|
if err := t.verifyLease(worker, attempt, now); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
t.LeaseExpiresAt = &until
|
t.LeaseExpiresAt = &until
|
||||||
@@ -195,7 +200,7 @@ func (t *Task) CompleteWith(resultArtifactID uuid.UUID, metrics map[string]any,
|
|||||||
return ErrResultConflict
|
return ErrResultConflict
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := t.verifyLease(worker, attempt); err != nil {
|
if err := t.verifyLease(worker, attempt, now); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -214,7 +219,7 @@ func (t *Task) CompleteWith(resultArtifactID uuid.UUID, metrics map[string]any,
|
|||||||
// Fail records a worker-reported failure. A retryable failure with attempts
|
// Fail records a worker-reported failure. A retryable failure with attempts
|
||||||
// left returns the task to the queue; otherwise it terminates as failed.
|
// left returns the task to the queue; otherwise it terminates as failed.
|
||||||
func (t *Task) Fail(worker string, attempt int, code, message string, retryable bool, now time.Time) error {
|
func (t *Task) Fail(worker string, attempt int, code, message string, retryable bool, now time.Time) error {
|
||||||
if err := t.verifyLease(worker, attempt); err != nil {
|
if err := t.verifyLease(worker, attempt, now); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
t.ErrorCode = &code
|
t.ErrorCode = &code
|
||||||
|
|||||||
@@ -172,14 +172,14 @@ func TestFirstHeartbeatMovesLeasedToRunning(t *testing.T) {
|
|||||||
task := leasedTask(1, 3)
|
task := leasedTask(1, 3)
|
||||||
until := testLater.Add(time.Hour)
|
until := testLater.Add(time.Hour)
|
||||||
|
|
||||||
if err := task.RenewLease(testWorker, 1, until); err != nil {
|
if err := task.RenewLease(testWorker, 1, testNow, until); err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
if task.Status != TaskRunning {
|
if task.Status != TaskRunning {
|
||||||
t.Errorf("status = %q, want running after first heartbeat", task.Status)
|
t.Errorf("status = %q, want running after first heartbeat", task.Status)
|
||||||
}
|
}
|
||||||
// A second heartbeat keeps it running.
|
// A second heartbeat keeps it running.
|
||||||
if err := task.RenewLease(testWorker, 1, until); err != nil {
|
if err := task.RenewLease(testWorker, 1, testNow, until); err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
if task.Status != TaskRunning {
|
if task.Status != TaskRunning {
|
||||||
@@ -190,15 +190,15 @@ func TestFirstHeartbeatMovesLeasedToRunning(t *testing.T) {
|
|||||||
func TestRunningTaskCanBeCompletedAndExpired(t *testing.T) {
|
func TestRunningTaskCanBeCompletedAndExpired(t *testing.T) {
|
||||||
// Complete works from running.
|
// Complete works from running.
|
||||||
task := leasedTask(1, 3)
|
task := leasedTask(1, 3)
|
||||||
_ = task.RenewLease(testWorker, 1, testLater) // -> running
|
_ = task.RenewLease(testWorker, 1, testNow, testLater) // -> running
|
||||||
if err := task.CompleteWith(testResult, nil, testWorker, 1, testNow); err != nil {
|
if err := task.CompleteWith(testResult, nil, testWorker, 1, testNow); err != nil {
|
||||||
t.Errorf("complete from running: %v", err)
|
t.Errorf("complete from running: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Expire reclaims a running task too.
|
// Expire reclaims a running task too.
|
||||||
task2 := leasedTask(1, 3)
|
task2 := leasedTask(1, 3)
|
||||||
_ = task2.RenewLease(testWorker, 1, testLater) // -> running
|
_ = task2.RenewLease(testWorker, 1, testNow, testLater) // -> running
|
||||||
task2.ExpireLease(testNow)
|
task2.ExpireLease(testLater)
|
||||||
if task2.Status != TaskPending {
|
if task2.Status != TaskPending {
|
||||||
t.Errorf("status = %q, want pending after a running lease expires", task2.Status)
|
t.Errorf("status = %q, want pending after a running lease expires", task2.Status)
|
||||||
}
|
}
|
||||||
@@ -208,14 +208,29 @@ func TestRenewLeaseExtendsOnlyForHolder(t *testing.T) {
|
|||||||
task := leasedTask(1, 3)
|
task := leasedTask(1, 3)
|
||||||
until := testLater.Add(time.Hour)
|
until := testLater.Add(time.Hour)
|
||||||
|
|
||||||
if err := task.RenewLease(testWorker, 1, until); err != nil {
|
if err := task.RenewLease(testWorker, 1, testNow, until); err != nil {
|
||||||
t.Fatalf("unexpected error: %v", err)
|
t.Fatalf("unexpected error: %v", err)
|
||||||
}
|
}
|
||||||
if !task.LeaseExpiresAt.Equal(until) {
|
if !task.LeaseExpiresAt.Equal(until) {
|
||||||
t.Error("lease must be extended")
|
t.Error("lease must be extended")
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := task.RenewLease("worker-2", 1, until); !errors.Is(err, ErrLeaseConflict) {
|
if err := task.RenewLease("worker-2", 1, testNow, until); !errors.Is(err, ErrLeaseConflict) {
|
||||||
t.Errorf("err = %v, want ErrLeaseConflict", err)
|
t.Errorf("err = %v, want ErrLeaseConflict", err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestExpiredLeaseRejectsRenewalCompletionAndFailure(t *testing.T) {
|
||||||
|
task := leasedTask(1, 3)
|
||||||
|
expired := testLater.Add(time.Nanosecond)
|
||||||
|
|
||||||
|
if err := task.RenewLease(testWorker, 1, expired, expired.Add(time.Minute)); !errors.Is(err, ErrLeaseConflict) {
|
||||||
|
t.Errorf("renew expired lease: err = %v, want ErrLeaseConflict", err)
|
||||||
|
}
|
||||||
|
if err := task.CompleteWith(testResult, nil, testWorker, 1, expired); !errors.Is(err, ErrLeaseConflict) {
|
||||||
|
t.Errorf("complete expired lease: err = %v, want ErrLeaseConflict", err)
|
||||||
|
}
|
||||||
|
if err := task.Fail(testWorker, 1, "timeout", "expired", true, expired); !errors.Is(err, ErrLeaseConflict) {
|
||||||
|
t.Errorf("fail expired lease: err = %v, want ErrLeaseConflict", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -26,14 +26,14 @@ func NewArtifactRepo(pool *pgxpool.Pool) *ArtifactRepo {
|
|||||||
var _ usecase.ArtifactRepository = (*ArtifactRepo)(nil)
|
var _ usecase.ArtifactRepository = (*ArtifactRepo)(nil)
|
||||||
|
|
||||||
var artifactColumns = []string{
|
var artifactColumns = []string{
|
||||||
"id", "job_id", "task_id", "kind", "filename", "storage_key",
|
"id", "job_id", "task_id", "attempt", "kind", "filename", "storage_key",
|
||||||
"content_type", "size_bytes", "sha256", "created_at",
|
"content_type", "size_bytes", "sha256", "created_at",
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *ArtifactRepo) Insert(ctx context.Context, a *domain.Artifact) error {
|
func (r *ArtifactRepo) Insert(ctx context.Context, a *domain.Artifact) error {
|
||||||
sql, args, err := psql.Insert("artifacts").
|
sql, args, err := psql.Insert("artifacts").
|
||||||
Columns(artifactColumns...).
|
Columns(artifactColumns...).
|
||||||
Values(a.ID, a.JobID, a.TaskID, string(a.Kind), a.Filename, a.StorageKey,
|
Values(a.ID, a.JobID, a.TaskID, a.Attempt, string(a.Kind), a.Filename, a.StorageKey,
|
||||||
a.ContentType, a.SizeBytes, a.SHA256, a.CreatedAt).
|
a.ContentType, a.SizeBytes, a.SHA256, a.CreatedAt).
|
||||||
ToSql()
|
ToSql()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -59,7 +59,7 @@ func (r *ArtifactRepo) Get(ctx context.Context, id uuid.UUID) (*domain.Artifact,
|
|||||||
kind string
|
kind string
|
||||||
)
|
)
|
||||||
err = conn(ctx, r.pool).QueryRow(ctx, sql, args...).Scan(
|
err = conn(ctx, r.pool).QueryRow(ctx, sql, args...).Scan(
|
||||||
&a.ID, &a.JobID, &a.TaskID, &kind, &a.Filename, &a.StorageKey,
|
&a.ID, &a.JobID, &a.TaskID, &a.Attempt, &kind, &a.Filename, &a.StorageKey,
|
||||||
&a.ContentType, &a.SizeBytes, &a.SHA256, &a.CreatedAt)
|
&a.ContentType, &a.SizeBytes, &a.SHA256, &a.CreatedAt)
|
||||||
if errors.Is(err, pgx.ErrNoRows) {
|
if errors.Is(err, pgx.ErrNoRows) {
|
||||||
return nil, domain.ErrArtifactNotFound
|
return nil, domain.ErrArtifactNotFound
|
||||||
|
|||||||
@@ -223,7 +223,7 @@ func TestUpdateRejectsStaleVersion(t *testing.T) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
if err := fresh.RenewLease("worker-1", fresh.Attempt, now.Add(2*time.Minute)); err != nil {
|
if err := fresh.RenewLease("worker-1", fresh.Attempt, now, now.Add(2*time.Minute)); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
return repo.Update(ctx, fresh)
|
return repo.Update(ctx, fresh)
|
||||||
@@ -259,6 +259,8 @@ func TestListCompletedIsOrderedByChunkIndex(t *testing.T) {
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
art.SetContent(fmt.Sprintf("rsha-%d", task.ChunkIndex), 1)
|
art.SetContent(fmt.Sprintf("rsha-%d", task.ChunkIndex), 1)
|
||||||
|
attempt := 1
|
||||||
|
art.Attempt = &attempt
|
||||||
if err := artifacts.Insert(ctx, art); err != nil {
|
if err := artifacts.Insert(ctx, art); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
@@ -269,6 +271,7 @@ func TestListCompletedIsOrderedByChunkIndex(t *testing.T) {
|
|||||||
}
|
}
|
||||||
owner := "worker-1"
|
owner := "worker-1"
|
||||||
fresh.Status = domain.TaskLeased
|
fresh.Status = domain.TaskLeased
|
||||||
|
fresh.Attempt = attempt
|
||||||
fresh.LeaseOwner = &owner
|
fresh.LeaseOwner = &owner
|
||||||
expires := now.Add(time.Minute)
|
expires := now.Add(time.Minute)
|
||||||
fresh.LeaseExpiresAt = &expires
|
fresh.LeaseExpiresAt = &expires
|
||||||
@@ -350,6 +353,10 @@ func seedArtifact(t *testing.T, pool *pgxpool.Pool, jobID uuid.UUID, taskID *uui
|
|||||||
t.Fatalf("build artifact: %v", err)
|
t.Fatalf("build artifact: %v", err)
|
||||||
}
|
}
|
||||||
art.SetContent(fmt.Sprintf("sha-%s", art.ID), 3)
|
art.SetContent(fmt.Sprintf("sha-%s", art.ID), 3)
|
||||||
|
if kind == domain.ArtifactPartialResult {
|
||||||
|
attempt := 1
|
||||||
|
art.Attempt = &attempt
|
||||||
|
}
|
||||||
if err := NewArtifactRepo(pool).Insert(context.Background(), art); err != nil {
|
if err := NewArtifactRepo(pool).Insert(context.Background(), art); err != nil {
|
||||||
t.Fatalf("insert artifact: %v", err)
|
t.Fatalf("insert artifact: %v", err)
|
||||||
}
|
}
|
||||||
@@ -446,6 +453,32 @@ func TestArtifactRepoRoundTrip(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestPartialResultArtifactRoundTripsAttempt(t *testing.T) {
|
||||||
|
pool := testPool(t)
|
||||||
|
ctx := context.Background()
|
||||||
|
job, tasks := seedJob(t, pool, 1)
|
||||||
|
taskID := tasks[0].ID
|
||||||
|
art, err := domain.NewArtifact(job.ID, &taskID, domain.ArtifactPartialResult,
|
||||||
|
"result.csv", "text/csv", time.Now().UTC())
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
attempt := 2
|
||||||
|
art.Attempt = &attempt
|
||||||
|
art.SetContent("sha", 3)
|
||||||
|
repo := NewArtifactRepo(pool)
|
||||||
|
if err := repo.Insert(ctx, art); err != nil {
|
||||||
|
t.Fatalf("insert: %v", err)
|
||||||
|
}
|
||||||
|
got, err := repo.Get(ctx, art.ID)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("get: %v", err)
|
||||||
|
}
|
||||||
|
if got.Attempt == nil || *got.Attempt != attempt {
|
||||||
|
t.Fatalf("attempt = %v, want %d", got.Attempt, attempt)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// A shard task stores its input as an artifact and no URI: this exercises the
|
// A shard task stores its input as an artifact and no URI: this exercises the
|
||||||
// nullable input_uri column, the input_artifact_id round-trip, and the
|
// nullable input_uri column, the input_artifact_id round-trip, and the
|
||||||
// ck_tasks_has_input check that requires one or the other.
|
// ck_tasks_has_input check that requires one or the other.
|
||||||
|
|||||||
@@ -3,6 +3,7 @@ package http
|
|||||||
import (
|
import (
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"errors"
|
"errors"
|
||||||
|
"io"
|
||||||
"net/http"
|
"net/http"
|
||||||
|
|
||||||
"github.com/emil28092005/SciMesh/coordinator/internal/domain"
|
"github.com/emil28092005/SciMesh/coordinator/internal/domain"
|
||||||
@@ -24,7 +25,13 @@ func decodeJSON(r *http.Request, dst any) error {
|
|||||||
// Reject unknown fields: silently ignoring a misspelled "worker_ID" would
|
// Reject unknown fields: silently ignoring a misspelled "worker_ID" would
|
||||||
// surface later as a baffling validation failure.
|
// surface later as a baffling validation failure.
|
||||||
dec.DisallowUnknownFields()
|
dec.DisallowUnknownFields()
|
||||||
return dec.Decode(dst)
|
if err := dec.Decode(dst); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := dec.Decode(&struct{}{}); !errors.Is(err, io.EOF) {
|
||||||
|
return errors.New("request body must contain exactly one JSON value")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// writeError translates domain errors into status codes. This mapping is the
|
// writeError translates domain errors into status codes. This mapping is the
|
||||||
|
|||||||
@@ -194,11 +194,14 @@ func (s *Server) handleUploadDataset(w http.ResponseWriter, r *http.Request) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
var (
|
var (
|
||||||
workload string
|
workload string
|
||||||
params map[string]any
|
params map[string]any
|
||||||
rows = defaultChunkRows
|
rows = defaultChunkRows
|
||||||
result usecase.SubmitDatasetResult
|
result usecase.SubmitDatasetResult
|
||||||
gotDataset bool
|
gotDataset bool
|
||||||
|
gotWorkload bool
|
||||||
|
gotParams bool
|
||||||
|
gotRows bool
|
||||||
)
|
)
|
||||||
|
|
||||||
for {
|
for {
|
||||||
@@ -213,9 +216,18 @@ func (s *Server) handleUploadDataset(w http.ResponseWriter, r *http.Request) {
|
|||||||
|
|
||||||
switch part.FormName() {
|
switch part.FormName() {
|
||||||
case "workload":
|
case "workload":
|
||||||
|
if gotDataset || gotWorkload {
|
||||||
|
s.writeError(w, r, domain.ErrInvalidInput)
|
||||||
|
return
|
||||||
|
}
|
||||||
b, _ := io.ReadAll(io.LimitReader(part, 1<<10))
|
b, _ := io.ReadAll(io.LimitReader(part, 1<<10))
|
||||||
workload = strings.TrimSpace(string(b))
|
workload = strings.TrimSpace(string(b))
|
||||||
|
gotWorkload = true
|
||||||
case "parameters":
|
case "parameters":
|
||||||
|
if gotDataset || gotParams {
|
||||||
|
s.writeError(w, r, domain.ErrInvalidInput)
|
||||||
|
return
|
||||||
|
}
|
||||||
b, _ := io.ReadAll(io.LimitReader(part, 1<<16))
|
b, _ := io.ReadAll(io.LimitReader(part, 1<<16))
|
||||||
if len(b) > 0 {
|
if len(b) > 0 {
|
||||||
if err := json.Unmarshal(b, ¶ms); err != nil {
|
if err := json.Unmarshal(b, ¶ms); err != nil {
|
||||||
@@ -223,12 +235,25 @@ func (s *Server) handleUploadDataset(w http.ResponseWriter, r *http.Request) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
gotParams = true
|
||||||
case "chunk_rows":
|
case "chunk_rows":
|
||||||
b, _ := io.ReadAll(io.LimitReader(part, 32))
|
if gotDataset || gotRows {
|
||||||
if n, err := strconv.Atoi(strings.TrimSpace(string(b))); err == nil {
|
s.writeError(w, r, domain.ErrInvalidInput)
|
||||||
rows = n
|
return
|
||||||
}
|
}
|
||||||
|
b, _ := io.ReadAll(io.LimitReader(part, 32))
|
||||||
|
n, err := strconv.Atoi(strings.TrimSpace(string(b)))
|
||||||
|
if err != nil || n < 1 {
|
||||||
|
s.writeError(w, r, domain.ErrInvalidInput)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
rows = n
|
||||||
|
gotRows = true
|
||||||
case "file", "dataset":
|
case "file", "dataset":
|
||||||
|
if gotDataset || workload == "" {
|
||||||
|
s.writeError(w, r, domain.ErrInvalidInput)
|
||||||
|
return
|
||||||
|
}
|
||||||
filename := part.FileName()
|
filename := part.FileName()
|
||||||
if filename == "" {
|
if filename == "" {
|
||||||
filename = "dataset"
|
filename = "dataset"
|
||||||
@@ -246,6 +271,9 @@ func (s *Server) handleUploadDataset(w http.ResponseWriter, r *http.Request) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
gotDataset = true
|
gotDataset = true
|
||||||
|
default:
|
||||||
|
s.writeError(w, r, domain.ErrInvalidInput)
|
||||||
|
return
|
||||||
}
|
}
|
||||||
_ = part.Close()
|
_ = part.Close()
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -259,6 +259,37 @@ func TestErrorMappings(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestJSONRejectsTrailingValue(t *testing.T) {
|
||||||
|
e := newEnv(t, healthy)
|
||||||
|
if code, _ := e.do(t, "POST", "/workers/register",
|
||||||
|
`{"name":"lab","capabilities":["w"]} {}`); code != http.StatusBadRequest {
|
||||||
|
t.Errorf("status = %d, want 400", code)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestUploadDatasetRejectsAmbiguousMultipartInput(t *testing.T) {
|
||||||
|
e := newEnv(t, healthy)
|
||||||
|
var buf bytes.Buffer
|
||||||
|
mw := multipart.NewWriter(&buf)
|
||||||
|
_ = mw.WriteField("workload", "w")
|
||||||
|
_ = mw.WriteField("chunk_rows", "not-a-number")
|
||||||
|
fw, _ := mw.CreateFormFile("file", "chembl.tsv")
|
||||||
|
_, _ = io.Copy(fw, strings.NewReader("id\tsmiles\nA\tCC\n"))
|
||||||
|
_ = mw.Close()
|
||||||
|
|
||||||
|
req, _ := http.NewRequestWithContext(context.Background(), "POST", e.ts.URL+"/jobs/upload", &buf)
|
||||||
|
req.Header.Set("Authorization", "Bearer "+token)
|
||||||
|
req.Header.Set("Content-Type", mw.FormDataContentType())
|
||||||
|
resp, err := http.DefaultClient.Do(req)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
defer resp.Body.Close()
|
||||||
|
if resp.StatusCode != http.StatusBadRequest {
|
||||||
|
t.Errorf("status = %d, want 400", resp.StatusCode)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// --- helpers -------------------------------------------------------------
|
// --- helpers -------------------------------------------------------------
|
||||||
|
|
||||||
func (e *env) putArtifact(t *testing.T, taskID, worker string, attempt int, data string) string {
|
func (e *env) putArtifact(t *testing.T, taskID, worker string, attempt int, data string) string {
|
||||||
|
|||||||
@@ -29,7 +29,7 @@ func (uc *UploadArtifact) Execute(ctx context.Context, in UploadArtifactInput) (
|
|||||||
}
|
}
|
||||||
// Only the worker holding the current lease at this attempt may upload the
|
// Only the worker holding the current lease at this attempt may upload the
|
||||||
// task's output — the coordinator never trusts an ownership claim on faith.
|
// task's output — the coordinator never trusts an ownership claim on faith.
|
||||||
if !task.IsLeaseHeldBy(in.WorkerID, in.Attempt) {
|
if !task.IsLeaseHeldBy(in.WorkerID, in.Attempt, uc.clk.Now()) {
|
||||||
return nil, domain.ErrLeaseConflict
|
return nil, domain.ErrLeaseConflict
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -39,6 +39,8 @@ func (uc *UploadArtifact) Execute(ctx context.Context, in UploadArtifactInput) (
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
attempt := in.Attempt
|
||||||
|
art.Attempt = &attempt
|
||||||
|
|
||||||
// Stream to storage first: size and checksum are measured here, by us, not
|
// Stream to storage first: size and checksum are measured here, by us, not
|
||||||
// taken from the worker. A large shard never sits in memory.
|
// taken from the worker. A large shard never sits in memory.
|
||||||
@@ -48,6 +50,19 @@ func (uc *UploadArtifact) Execute(ctx context.Context, in UploadArtifactInput) (
|
|||||||
}
|
}
|
||||||
art.SetContent(sum, size)
|
art.SetContent(sum, size)
|
||||||
|
|
||||||
|
// The stream may take longer than the lease. Re-check after it finishes so
|
||||||
|
// an expired worker cannot leave a durable result record behind. Completion
|
||||||
|
// performs the same ownership check under its transaction.
|
||||||
|
current, err := uc.tasks.Get(ctx, in.TaskID)
|
||||||
|
if err != nil {
|
||||||
|
_ = uc.blobs.Delete(ctx, art.StorageKey)
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if !current.IsLeaseHeldBy(in.WorkerID, in.Attempt, uc.clk.Now()) {
|
||||||
|
_ = uc.blobs.Delete(ctx, art.StorageKey)
|
||||||
|
return nil, domain.ErrLeaseConflict
|
||||||
|
}
|
||||||
|
|
||||||
// Persist the record. If that fails the blob would be an orphan, so remove it.
|
// Persist the record. If that fails the blob would be an orphan, so remove it.
|
||||||
if err := uc.artifacts.Insert(ctx, art); err != nil {
|
if err := uc.artifacts.Insert(ctx, art); err != nil {
|
||||||
_ = uc.blobs.Delete(ctx, art.StorageKey)
|
_ = uc.blobs.Delete(ctx, art.StorageKey)
|
||||||
|
|||||||
@@ -91,7 +91,8 @@ func (uc *RenewLease) Execute(ctx context.Context, in RenewLeaseInput) (*domain.
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
if err := task.RenewLease(in.WorkerID, in.Attempt, uc.clock.Now().Add(uc.leaseDuration)); err != nil {
|
now := uc.clock.Now()
|
||||||
|
if err := task.RenewLease(in.WorkerID, in.Attempt, now, now.Add(uc.leaseDuration)); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
if err := uc.tasks.Update(ctx, task); err != nil {
|
if err := uc.tasks.Update(ctx, task); err != nil {
|
||||||
@@ -143,7 +144,7 @@ func (uc *CompleteTask) Execute(ctx context.Context, in CompleteTaskInput) (*dom
|
|||||||
}
|
}
|
||||||
// Rule 10: never trust a worker-supplied artifact reference. The result
|
// Rule 10: never trust a worker-supplied artifact reference. The result
|
||||||
// must be an artifact the coordinator itself stored for *this* task.
|
// must be an artifact the coordinator itself stored for *this* task.
|
||||||
if err := uc.verifyResultArtifact(ctx, in.TaskID, in.ResultArtifactID); err != nil {
|
if err := uc.verifyResultArtifact(ctx, in.TaskID, in.Attempt, in.ResultArtifactID); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
now := uc.clock.Now()
|
now := uc.clock.Now()
|
||||||
@@ -175,12 +176,12 @@ func (uc *CompleteTask) Execute(ctx context.Context, in CompleteTaskInput) (*dom
|
|||||||
// verifyResultArtifact enforces that the referenced artifact was stored by the
|
// verifyResultArtifact enforces that the referenced artifact was stored by the
|
||||||
// coordinator for this exact task. It stops a worker from completing task B with
|
// coordinator for this exact task. It stops a worker from completing task B with
|
||||||
// an artifact it uploaded for task A, and from naming an id that isn't a result.
|
// an artifact it uploaded for task A, and from naming an id that isn't a result.
|
||||||
func (uc *CompleteTask) verifyResultArtifact(ctx context.Context, taskID, artifactID uuid.UUID) error {
|
func (uc *CompleteTask) verifyResultArtifact(ctx context.Context, taskID uuid.UUID, attempt int, artifactID uuid.UUID) error {
|
||||||
art, err := uc.artifacts.Get(ctx, artifactID)
|
art, err := uc.artifacts.Get(ctx, artifactID)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
if art.TaskID == nil || *art.TaskID != taskID || art.Kind != domain.ArtifactPartialResult {
|
if art.TaskID == nil || *art.TaskID != taskID || art.Attempt == nil || *art.Attempt != attempt || art.Kind != domain.ArtifactPartialResult {
|
||||||
return domain.ErrResultConflict
|
return domain.ErrResultConflict
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
|
|||||||
@@ -4,6 +4,7 @@ import (
|
|||||||
"context"
|
"context"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"io"
|
||||||
"strings"
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
@@ -19,6 +20,17 @@ var ctx = context.Background()
|
|||||||
|
|
||||||
const lease = 2 * time.Minute
|
const lease = 2 * time.Minute
|
||||||
|
|
||||||
|
type expiringBlobStore struct {
|
||||||
|
*memstore.BlobStore
|
||||||
|
clock *memstore.Clock
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s expiringBlobStore) Put(ctx context.Context, key string, body io.Reader) (string, int64, error) {
|
||||||
|
sum, size, err := s.BlobStore.Put(ctx, key, body)
|
||||||
|
s.clock.Advance(lease + time.Second)
|
||||||
|
return sum, size, err
|
||||||
|
}
|
||||||
|
|
||||||
// harness wires every use case to in-memory stores so orchestration can be
|
// harness wires every use case to in-memory stores so orchestration can be
|
||||||
// tested without a database.
|
// tested without a database.
|
||||||
type harness struct {
|
type harness struct {
|
||||||
@@ -239,6 +251,46 @@ func TestCompleteRejectsForeignArtifact(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestCompleteRejectsArtifactFromExpiredAttempt(t *testing.T) {
|
||||||
|
h := newHarness()
|
||||||
|
h.seedJob(t, "w", 1)
|
||||||
|
taskID, attemptOne := h.leaseOne(t, "w1", "w")
|
||||||
|
staleArtifact := h.uploadResult(t, taskID, "w1", attemptOne)
|
||||||
|
|
||||||
|
h.clk.Advance(lease + time.Second)
|
||||||
|
if _, err := h.expire.Execute(ctx); err != nil {
|
||||||
|
t.Fatalf("expire lease: %v", err)
|
||||||
|
}
|
||||||
|
_, attemptTwo := h.leaseOne(t, "w2", "w")
|
||||||
|
if attemptTwo != attemptOne+1 {
|
||||||
|
t.Fatalf("attempt = %d, want %d", attemptTwo, attemptOne+1)
|
||||||
|
}
|
||||||
|
|
||||||
|
_, err := h.complete.Execute(ctx, usecase.CompleteTaskInput{
|
||||||
|
TaskID: taskID, WorkerID: "w2", Attempt: attemptTwo, ResultArtifactID: staleArtifact,
|
||||||
|
})
|
||||||
|
if !errors.Is(err, domain.ErrResultConflict) {
|
||||||
|
t.Errorf("stale-attempt artifact: err = %v, want ErrResultConflict", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestUploadRejectsLeaseThatExpiresDuringStreaming(t *testing.T) {
|
||||||
|
h := newHarness()
|
||||||
|
h.seedJob(t, "w", 1)
|
||||||
|
taskID, attempt := h.leaseOne(t, "w1", "w")
|
||||||
|
h.uploadArt = usecase.NewUploadArtifact(
|
||||||
|
h.tasks, h.arts, expiringBlobStore{BlobStore: h.blobs, clock: h.clk}, h.clk,
|
||||||
|
)
|
||||||
|
|
||||||
|
_, err := h.uploadArt.Execute(ctx, usecase.UploadArtifactInput{
|
||||||
|
TaskID: taskID, WorkerID: "w1", Attempt: attempt,
|
||||||
|
Filename: "result.csv", ContentType: "text/csv", Body: strings.NewReader("result"),
|
||||||
|
})
|
||||||
|
if !errors.Is(err, domain.ErrLeaseConflict) {
|
||||||
|
t.Errorf("upload after lease expiry: err = %v, want ErrLeaseConflict", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestCompleteIsIdempotentOnReplay(t *testing.T) {
|
func TestCompleteIsIdempotentOnReplay(t *testing.T) {
|
||||||
h := newHarness()
|
h := newHarness()
|
||||||
h.seedJob(t, "w", 1)
|
h.seedJob(t, "w", 1)
|
||||||
|
|||||||
@@ -0,0 +1,7 @@
|
|||||||
|
BEGIN;
|
||||||
|
|
||||||
|
ALTER TABLE artifacts DROP CONSTRAINT IF EXISTS ck_artifact_attempt_positive;
|
||||||
|
ALTER TABLE artifacts DROP CONSTRAINT IF EXISTS ck_partial_result_attempt;
|
||||||
|
ALTER TABLE artifacts DROP COLUMN IF EXISTS attempt;
|
||||||
|
|
||||||
|
COMMIT;
|
||||||
@@ -0,0 +1,34 @@
|
|||||||
|
BEGIN;
|
||||||
|
|
||||||
|
-- A partial result belongs to the lease attempt that uploaded it. Without this
|
||||||
|
-- binding a worker holding a later retry could complete a task with stale bytes
|
||||||
|
-- uploaded by an expired attempt of that same task.
|
||||||
|
ALTER TABLE artifacts ADD COLUMN attempt integer;
|
||||||
|
|
||||||
|
-- A completed task never gets a later lease, so its current attempt is also
|
||||||
|
-- the attempt that produced the stored result.
|
||||||
|
UPDATE artifacts AS a
|
||||||
|
SET attempt = t.attempt
|
||||||
|
FROM tasks AS t
|
||||||
|
WHERE a.task_id = t.id
|
||||||
|
AND a.kind = 'partial_result'::artifact_kind
|
||||||
|
AND t.status = 'completed'::task_status
|
||||||
|
AND a.attempt IS NULL;
|
||||||
|
|
||||||
|
-- For unfinished tasks the old schema cannot tell which attempt uploaded a
|
||||||
|
-- partial result. Keeping it would let a later retry claim stale bytes, so the
|
||||||
|
-- worker must upload again. Blob garbage is harmless and follows the existing
|
||||||
|
-- coordinator-owned storage cleanup policy.
|
||||||
|
DELETE FROM artifacts AS a
|
||||||
|
USING tasks AS t
|
||||||
|
WHERE a.task_id = t.id
|
||||||
|
AND a.kind = 'partial_result'::artifact_kind
|
||||||
|
AND t.status <> 'completed'::task_status
|
||||||
|
AND a.attempt IS NULL;
|
||||||
|
|
||||||
|
ALTER TABLE artifacts ADD CONSTRAINT ck_partial_result_attempt
|
||||||
|
CHECK (kind <> 'partial_result'::artifact_kind OR attempt IS NOT NULL);
|
||||||
|
ALTER TABLE artifacts ADD CONSTRAINT ck_artifact_attempt_positive
|
||||||
|
CHECK (attempt IS NULL OR attempt > 0);
|
||||||
|
|
||||||
|
COMMIT;
|
||||||
@@ -56,6 +56,8 @@ Response: `{ "worker_id": "<uuid>", "heartbeat_interval_seconds": 15 }`.
|
|||||||
- **Keep `worker_id`**. Use it as your identity in every later call. Using the
|
- **Keep `worker_id`**. Use it as your identity in every later call. Using the
|
||||||
registered UUID is what lets the coordinator track your liveness (it marks
|
registered UUID is what lets the coordinator track your liveness (it marks
|
||||||
workers offline after they go silent).
|
workers offline after they go silent).
|
||||||
|
- Current coordinator jobs use `similarity_search` / `similarity_graph`; the
|
||||||
|
reference Python worker also accepts the public CLI spellings with hyphens.
|
||||||
|
|
||||||
## 2. Claim a task
|
## 2. Claim a task
|
||||||
|
|
||||||
@@ -185,7 +187,8 @@ next claim.
|
|||||||
Per the worker contract, at minimum:
|
Per the worker contract, at minimum:
|
||||||
|
|
||||||
- `SCIMESH_COORDINATOR_URL` (e.g. `http://coordinator:8080`)
|
- `SCIMESH_COORDINATOR_URL` (e.g. `http://coordinator:8080`)
|
||||||
- `SCIMESH_WORKER_ID` (or derive from hostname)
|
- worker name (the coordinator returns its `worker_id` at registration;
|
||||||
|
`SCIMESH_WORKER_ID` is only a legacy/test override)
|
||||||
- the bearer token
|
- the bearer token
|
||||||
- poll interval and request timeout
|
- poll interval and request timeout
|
||||||
- a working directory for downloaded inputs and generated outputs
|
- a working directory for downloaded inputs and generated outputs
|
||||||
|
|||||||
+32
-41
@@ -7,38 +7,23 @@ import http.client
|
|||||||
import json
|
import json
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Protocol
|
from typing import Protocol
|
||||||
from urllib.parse import quote, urlsplit
|
from urllib.parse import quote, urljoin, urlsplit
|
||||||
from urllib.request import HTTPRedirectHandler, Request, build_opener
|
from urllib.request import Request, build_opener
|
||||||
|
|
||||||
from .models import ClaimedTask, ProducedArtifact
|
from .coordinator import CoordinatorConflictError
|
||||||
|
from .models import ClaimedTask, ProducedArtifact, UploadedArtifact
|
||||||
|
from .transport import SameOriginAuthRedirectHandler, origin
|
||||||
|
|
||||||
|
# Compatibility aliases for focused transport tests.
|
||||||
def _origin(uri: str) -> tuple[str, str, int | None]:
|
_SameOriginAuthRedirectHandler = SameOriginAuthRedirectHandler
|
||||||
parsed = urlsplit(uri)
|
_origin = origin
|
||||||
scheme = parsed.scheme.lower()
|
|
||||||
default_port = {"http": 80, "https": 443}.get(scheme)
|
|
||||||
return scheme, (parsed.hostname or "").lower(), parsed.port or default_port
|
|
||||||
|
|
||||||
|
|
||||||
class _SameOriginAuthRedirectHandler(HTTPRedirectHandler):
|
|
||||||
"""Do not forward the coordinator token when a download changes origin."""
|
|
||||||
|
|
||||||
def __init__(self, coordinator_origin: tuple[str, str, int | None]) -> None:
|
|
||||||
super().__init__()
|
|
||||||
self.coordinator_origin = coordinator_origin
|
|
||||||
|
|
||||||
def redirect_request(self, req: Request, fp: object, code: int, msg: str, headers: object, newurl: str) -> Request | None:
|
|
||||||
redirected = super().redirect_request(req, fp, code, msg, headers, newurl)
|
|
||||||
if redirected and _origin(newurl) != self.coordinator_origin:
|
|
||||||
redirected.remove_header("Authorization")
|
|
||||||
return redirected
|
|
||||||
|
|
||||||
class ArtifactClient(Protocol):
|
class ArtifactClient(Protocol):
|
||||||
def download(self, uri: str, destination: Path) -> None: ...
|
def download(self, uri: str, destination: Path) -> None: ...
|
||||||
|
|
||||||
def upload(
|
def upload(
|
||||||
self, task: ClaimedTask, worker_id: str, artifact: ProducedArtifact
|
self, task: ClaimedTask, worker_id: str, artifact: ProducedArtifact
|
||||||
) -> str: ...
|
) -> UploadedArtifact: ...
|
||||||
|
|
||||||
|
|
||||||
class HttpArtifactClient:
|
class HttpArtifactClient:
|
||||||
@@ -48,18 +33,21 @@ class HttpArtifactClient:
|
|||||||
self.coordinator_url = coordinator_url.rstrip("/")
|
self.coordinator_url = coordinator_url.rstrip("/")
|
||||||
self.timeout = timeout
|
self.timeout = timeout
|
||||||
self.bearer_token = bearer_token
|
self.bearer_token = bearer_token
|
||||||
self.coordinator_origin = _origin(coordinator_url)
|
self.coordinator_origin = origin(coordinator_url)
|
||||||
self._opener = build_opener(_SameOriginAuthRedirectHandler(self.coordinator_origin))
|
self._opener = build_opener(SameOriginAuthRedirectHandler(self.coordinator_origin))
|
||||||
|
|
||||||
def download(self, uri: str, destination: Path) -> None:
|
def download(self, uri: str, destination: Path) -> None:
|
||||||
destination.parent.mkdir(parents=True, exist_ok=True)
|
destination.parent.mkdir(parents=True, exist_ok=True)
|
||||||
request = Request(uri, headers=self._auth_headers_for(uri))
|
resolved_uri = urljoin(f"{self.coordinator_url}/", uri)
|
||||||
|
request = Request(resolved_uri, headers=self._auth_headers_for(resolved_uri))
|
||||||
with self._opener.open(request, timeout=self.timeout) as response, destination.open("wb") as target:
|
with self._opener.open(request, timeout=self.timeout) as response, destination.open("wb") as target:
|
||||||
while chunk := response.read(1024 * 1024):
|
while chunk := response.read(1024 * 1024):
|
||||||
target.write(chunk)
|
target.write(chunk)
|
||||||
|
|
||||||
def upload(self, task: ClaimedTask, worker_id: str, artifact: ProducedArtifact) -> str:
|
def upload(
|
||||||
"""Stream one result artifact to the coordinator and return its stable URI."""
|
self, task: ClaimedTask, worker_id: str, artifact: ProducedArtifact
|
||||||
|
) -> UploadedArtifact:
|
||||||
|
"""Stream an artifact and require durable coordinator-owned metadata."""
|
||||||
url = (
|
url = (
|
||||||
f"{self.coordinator_url}/tasks/{quote(task.task_id, safe='')}/artifacts/"
|
f"{self.coordinator_url}/tasks/{quote(task.task_id, safe='')}/artifacts/"
|
||||||
f"{quote(artifact.path.name, safe='')}"
|
f"{quote(artifact.path.name, safe='')}"
|
||||||
@@ -71,11 +59,13 @@ class HttpArtifactClient:
|
|||||||
http.client.HTTPSConnection if parsed.scheme == "https" else http.client.HTTPConnection
|
http.client.HTTPSConnection if parsed.scheme == "https" else http.client.HTTPConnection
|
||||||
)
|
)
|
||||||
connection = connection_class(parsed.hostname, parsed.port, timeout=self.timeout)
|
connection = connection_class(parsed.hostname, parsed.port, timeout=self.timeout)
|
||||||
|
local_size = artifact.path.stat().st_size
|
||||||
|
local_sha256 = sha256_file(artifact.path)
|
||||||
try:
|
try:
|
||||||
path = parsed.path + (f"?{parsed.query}" if parsed.query else "")
|
path = parsed.path + (f"?{parsed.query}" if parsed.query else "")
|
||||||
connection.putrequest("PUT", path)
|
connection.putrequest("PUT", path)
|
||||||
connection.putheader("Content-Type", artifact.content_type)
|
connection.putheader("Content-Type", artifact.content_type)
|
||||||
connection.putheader("Content-Length", str(artifact.path.stat().st_size))
|
connection.putheader("Content-Length", str(local_size))
|
||||||
connection.putheader("X-Worker-ID", worker_id)
|
connection.putheader("X-Worker-ID", worker_id)
|
||||||
connection.putheader("X-Task-Attempt", str(task.attempt))
|
connection.putheader("X-Task-Attempt", str(task.attempt))
|
||||||
for name, value in self._auth_headers_for(url).items():
|
for name, value in self._auth_headers_for(url).items():
|
||||||
@@ -86,23 +76,24 @@ class HttpArtifactClient:
|
|||||||
connection.send(chunk)
|
connection.send(chunk)
|
||||||
response = connection.getresponse()
|
response = connection.getresponse()
|
||||||
body = response.read()
|
body = response.read()
|
||||||
if not 200 <= response.status < 300:
|
if response.status == 409:
|
||||||
|
raise CoordinatorConflictError("artifact upload rejected because the task lease was lost")
|
||||||
|
if response.status != 200:
|
||||||
raise RuntimeError(f"artifact upload rejected with status {response.status}")
|
raise RuntimeError(f"artifact upload rejected with status {response.status}")
|
||||||
if body:
|
try:
|
||||||
try:
|
response_data = json.loads(body)
|
||||||
response_data = json.loads(body)
|
uploaded = UploadedArtifact.from_json(response_data)
|
||||||
except json.JSONDecodeError as error:
|
except (ValueError, json.JSONDecodeError) as error:
|
||||||
raise RuntimeError("artifact upload returned invalid JSON") from error
|
raise RuntimeError("artifact upload returned invalid metadata") from error
|
||||||
response_uri = response_data.get("uri") if isinstance(response_data, dict) else None
|
if uploaded.sha256 != local_sha256 or uploaded.size_bytes != local_size:
|
||||||
if isinstance(response_uri, str) and response_uri:
|
raise RuntimeError("artifact upload metadata does not match local artifact")
|
||||||
return response_uri
|
return uploaded
|
||||||
return url
|
|
||||||
finally:
|
finally:
|
||||||
connection.close()
|
connection.close()
|
||||||
|
|
||||||
def _auth_headers_for(self, uri: str) -> dict[str, str]:
|
def _auth_headers_for(self, uri: str) -> dict[str, str]:
|
||||||
"""Only coordinator-owned URLs receive the coordinator bearer token."""
|
"""Only coordinator-owned URLs receive the coordinator bearer token."""
|
||||||
if self.bearer_token and _origin(uri) == self.coordinator_origin:
|
if self.bearer_token and origin(uri) == self.coordinator_origin:
|
||||||
return {"Authorization": f"Bearer {self.bearer_token}"}
|
return {"Authorization": f"Bearer {self.bearer_token}"}
|
||||||
return {}
|
return {}
|
||||||
|
|
||||||
|
|||||||
+24
-4
@@ -14,22 +14,42 @@ from .runners import SciMeshRunner
|
|||||||
|
|
||||||
|
|
||||||
def main(argv: list[str] | None = None) -> int:
|
def main(argv: list[str] | None = None) -> int:
|
||||||
parser = argparse.ArgumentParser(prog="scimesh-worker")
|
parser = argparse.ArgumentParser(
|
||||||
|
prog="scimesh-worker",
|
||||||
|
epilog=(
|
||||||
|
"Environment: SCIMESH_COORDINATOR_URL, SCIMESH_WORK_DIR, "
|
||||||
|
"SCIMESH_WORKER_NAME, SCIMESH_CPU_COUNT, SCIMESH_MEMORY_MB, "
|
||||||
|
"SCIMESH_POLL_INTERVAL, SCIMESH_REQUEST_TIMEOUT, "
|
||||||
|
"SCIMESH_HEARTBEAT_INTERVAL, SCIMESH_CLEANUP_AFTER_SECONDS, and "
|
||||||
|
"SCIMESH_BEARER_TOKEN. SCIMESH_WORKER_ID is a legacy/test override."
|
||||||
|
),
|
||||||
|
)
|
||||||
parser.add_argument("--coordinator-url")
|
parser.add_argument("--coordinator-url")
|
||||||
parser.add_argument("--worker-id")
|
parser.add_argument("--worker-id")
|
||||||
parser.add_argument("--work-dir")
|
parser.add_argument("--work-dir")
|
||||||
|
parser.add_argument("--worker-name")
|
||||||
|
parser.add_argument("--cpu-count", type=int)
|
||||||
|
parser.add_argument("--memory-mb", type=int)
|
||||||
parser.add_argument("--poll-interval", type=float)
|
parser.add_argument("--poll-interval", type=float)
|
||||||
parser.add_argument("--request-timeout", type=float)
|
parser.add_argument("--request-timeout", type=float)
|
||||||
parser.add_argument("--heartbeat-interval", type=float)
|
parser.add_argument("--heartbeat-interval", type=float)
|
||||||
|
parser.add_argument("--cleanup-after-seconds", type=float)
|
||||||
args = parser.parse_args(argv)
|
args = parser.parse_args(argv)
|
||||||
config = WorkerConfig.from_environment()
|
|
||||||
overrides = {key: value for key, value in vars(args).items() if value is not None}
|
overrides = {key: value for key, value in vars(args).items() if value is not None}
|
||||||
if "work_dir" in overrides:
|
if "work_dir" in overrides:
|
||||||
overrides["work_dir"] = Path(overrides["work_dir"])
|
overrides["work_dir"] = Path(overrides["work_dir"])
|
||||||
config = WorkerConfig(**{**config.__dict__, **overrides})
|
try:
|
||||||
|
config = WorkerConfig.from_environment(overrides)
|
||||||
|
except (TypeError, ValueError) as error:
|
||||||
|
parser.error(str(error))
|
||||||
logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s")
|
logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s")
|
||||||
client = HttpCoordinatorClient(config.coordinator_url, config.request_timeout, config.bearer_token)
|
client = HttpCoordinatorClient(config.coordinator_url, config.request_timeout, config.bearer_token)
|
||||||
WorkerDaemon(config, client, HttpArtifactClient(config.coordinator_url, config.request_timeout, config.bearer_token), SciMeshRunner()).run_forever()
|
WorkerDaemon(
|
||||||
|
config,
|
||||||
|
client,
|
||||||
|
HttpArtifactClient(config.coordinator_url, config.request_timeout, config.bearer_token),
|
||||||
|
SciMeshRunner(),
|
||||||
|
).run_forever()
|
||||||
return 0
|
return 0
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
+75
-20
@@ -3,44 +3,99 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
|
from math import isfinite
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
import os
|
import os
|
||||||
|
import socket
|
||||||
|
from typing import Mapping
|
||||||
|
from urllib.parse import urlsplit
|
||||||
|
|
||||||
|
|
||||||
|
def _positive_number(value: object, name: str, *, allow_zero: bool = False) -> None:
|
||||||
|
if (
|
||||||
|
isinstance(value, bool)
|
||||||
|
or not isinstance(value, (int, float))
|
||||||
|
or not isfinite(value)
|
||||||
|
or value < 0
|
||||||
|
or (not allow_zero and value == 0)
|
||||||
|
):
|
||||||
|
qualifier = "non-negative" if allow_zero else "positive"
|
||||||
|
raise ValueError(f"{name} must be {qualifier}")
|
||||||
|
|
||||||
|
|
||||||
@dataclass(frozen=True)
|
@dataclass(frozen=True)
|
||||||
class WorkerConfig:
|
class WorkerConfig:
|
||||||
coordinator_url: str
|
coordinator_url: str
|
||||||
worker_id: str
|
worker_id: str | None
|
||||||
work_dir: Path
|
work_dir: Path
|
||||||
|
worker_name: str = "scimesh-worker"
|
||||||
|
cpu_count: int = 1
|
||||||
|
memory_mb: int | None = None
|
||||||
poll_interval: float = 2.0
|
poll_interval: float = 2.0
|
||||||
request_timeout: float = 30.0
|
request_timeout: float = 30.0
|
||||||
heartbeat_interval: float = 15.0
|
heartbeat_interval: float = 15.0
|
||||||
bearer_token: str | None = None
|
bearer_token: str | None = None
|
||||||
cleanup_after_seconds: float | None = None
|
cleanup_after_seconds: float | None = None
|
||||||
capabilities: tuple[str, ...] = ("similarity-search", "similarity-graph")
|
# The local CLI uses hyphens; the first coordinator contract used
|
||||||
|
# underscores. Advertise both stable spellings while jobs are migrated.
|
||||||
|
capabilities: tuple[str, ...] = (
|
||||||
|
"similarity-search",
|
||||||
|
"similarity-graph",
|
||||||
|
"similarity_search",
|
||||||
|
"similarity_graph",
|
||||||
|
)
|
||||||
|
|
||||||
def __post_init__(self) -> None:
|
def __post_init__(self) -> None:
|
||||||
if self.poll_interval <= 0:
|
parsed = urlsplit(self.coordinator_url)
|
||||||
raise ValueError("poll_interval must be positive")
|
if parsed.scheme not in {"http", "https"} or not parsed.hostname:
|
||||||
if self.request_timeout <= 0:
|
raise ValueError("coordinator_url must be an absolute HTTP(S) URL")
|
||||||
raise ValueError("request_timeout must be positive")
|
if not isinstance(self.worker_name, str) or not self.worker_name.strip():
|
||||||
if self.heartbeat_interval <= 0:
|
raise ValueError("worker_name must be non-empty")
|
||||||
raise ValueError("heartbeat_interval must be positive")
|
if isinstance(self.cpu_count, bool) or not isinstance(self.cpu_count, int) or self.cpu_count < 1:
|
||||||
|
raise ValueError("cpu_count must be positive")
|
||||||
|
if self.worker_id is not None and not isinstance(self.worker_id, str):
|
||||||
|
raise ValueError("worker_id must be a string when set")
|
||||||
|
if self.memory_mb is not None and (
|
||||||
|
isinstance(self.memory_mb, bool)
|
||||||
|
or not isinstance(self.memory_mb, int)
|
||||||
|
or self.memory_mb < 1
|
||||||
|
):
|
||||||
|
raise ValueError("memory_mb must be positive when set")
|
||||||
|
_positive_number(self.poll_interval, "poll_interval")
|
||||||
|
_positive_number(self.request_timeout, "request_timeout")
|
||||||
|
_positive_number(self.heartbeat_interval, "heartbeat_interval")
|
||||||
|
if self.cleanup_after_seconds is not None:
|
||||||
|
_positive_number(self.cleanup_after_seconds, "cleanup_after_seconds", allow_zero=True)
|
||||||
|
if not self.capabilities:
|
||||||
|
raise ValueError("capabilities cannot be empty")
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def from_environment(cls) -> "WorkerConfig":
|
def from_environment(
|
||||||
url = os.getenv("SCIMESH_COORDINATOR_URL")
|
cls, overrides: Mapping[str, object] | None = None
|
||||||
worker_id = os.getenv("SCIMESH_WORKER_ID")
|
) -> "WorkerConfig":
|
||||||
if not url or not worker_id:
|
"""Build config from environment, allowing typed CLI values to override it."""
|
||||||
raise ValueError("SCIMESH_COORDINATOR_URL and SCIMESH_WORKER_ID are required")
|
values = overrides or {}
|
||||||
cleanup = os.getenv("SCIMESH_CLEANUP_AFTER_SECONDS")
|
|
||||||
|
def value(name: str, environment: str, default: object | None = None) -> object | None:
|
||||||
|
override = values.get(name)
|
||||||
|
return override if override is not None else os.getenv(environment, default)
|
||||||
|
|
||||||
|
url = value("coordinator_url", "SCIMESH_COORDINATOR_URL")
|
||||||
|
if not isinstance(url, str) or not url:
|
||||||
|
raise ValueError("SCIMESH_COORDINATOR_URL or --coordinator-url is required")
|
||||||
|
cleanup = value("cleanup_after_seconds", "SCIMESH_CLEANUP_AFTER_SECONDS")
|
||||||
|
cpu_count = value("cpu_count", "SCIMESH_CPU_COUNT", os.cpu_count() or 1)
|
||||||
|
memory_mb = value("memory_mb", "SCIMESH_MEMORY_MB")
|
||||||
return cls(
|
return cls(
|
||||||
coordinator_url=url.rstrip("/"),
|
coordinator_url=url.rstrip("/"),
|
||||||
worker_id=worker_id,
|
worker_id=value("worker_id", "SCIMESH_WORKER_ID"),
|
||||||
work_dir=Path(os.getenv("SCIMESH_WORK_DIR", "./scimesh-worker-data")),
|
work_dir=Path(value("work_dir", "SCIMESH_WORK_DIR", "./scimesh-worker-data")),
|
||||||
poll_interval=float(os.getenv("SCIMESH_POLL_INTERVAL", "2")),
|
worker_name=str(value("worker_name", "SCIMESH_WORKER_NAME", socket.gethostname())),
|
||||||
request_timeout=float(os.getenv("SCIMESH_REQUEST_TIMEOUT", "30")),
|
cpu_count=int(cpu_count),
|
||||||
heartbeat_interval=float(os.getenv("SCIMESH_HEARTBEAT_INTERVAL", "15")),
|
memory_mb=int(memory_mb) if memory_mb is not None else None,
|
||||||
bearer_token=os.getenv("SCIMESH_BEARER_TOKEN"),
|
poll_interval=float(value("poll_interval", "SCIMESH_POLL_INTERVAL", "2")),
|
||||||
|
request_timeout=float(value("request_timeout", "SCIMESH_REQUEST_TIMEOUT", "30")),
|
||||||
|
heartbeat_interval=float(value("heartbeat_interval", "SCIMESH_HEARTBEAT_INTERVAL", "15")),
|
||||||
|
bearer_token=value("bearer_token", "SCIMESH_BEARER_TOKEN"),
|
||||||
cleanup_after_seconds=float(cleanup) if cleanup else None,
|
cleanup_after_seconds=float(cleanup) if cleanup else None,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -5,9 +5,10 @@ from __future__ import annotations
|
|||||||
import json
|
import json
|
||||||
from typing import Any, Protocol
|
from typing import Any, Protocol
|
||||||
from urllib.error import HTTPError, URLError
|
from urllib.error import HTTPError, URLError
|
||||||
from urllib.request import Request, urlopen
|
from urllib.request import Request, build_opener
|
||||||
|
|
||||||
from .models import ClaimedTask
|
from .models import ClaimedTask, RegisteredWorker
|
||||||
|
from .transport import NoRedirectHandler
|
||||||
|
|
||||||
|
|
||||||
class CoordinatorError(RuntimeError):
|
class CoordinatorError(RuntimeError):
|
||||||
@@ -18,7 +19,15 @@ class CoordinatorTransientError(CoordinatorError):
|
|||||||
"""A timeout, connection error, or 5xx coordinator response."""
|
"""A timeout, connection error, or 5xx coordinator response."""
|
||||||
|
|
||||||
|
|
||||||
|
class CoordinatorConflictError(CoordinatorError):
|
||||||
|
"""The worker no longer owns the task lease or attempted a conflicting mutation."""
|
||||||
|
|
||||||
|
|
||||||
class CoordinatorClient(Protocol):
|
class CoordinatorClient(Protocol):
|
||||||
|
def register(
|
||||||
|
self, name: str, capabilities: tuple[str, ...], cpu_count: int, memory_mb: int | None
|
||||||
|
) -> RegisteredWorker: ...
|
||||||
|
|
||||||
def claim(self, worker_id: str, capabilities: tuple[str, ...]) -> ClaimedTask | None: ...
|
def claim(self, worker_id: str, capabilities: tuple[str, ...]) -> ClaimedTask | None: ...
|
||||||
|
|
||||||
def submit(self, task: ClaimedTask, payload: dict[str, Any]) -> None: ...
|
def submit(self, task: ClaimedTask, payload: dict[str, Any]) -> None: ...
|
||||||
@@ -33,6 +42,25 @@ class HttpCoordinatorClient:
|
|||||||
self.base_url = base_url.rstrip("/")
|
self.base_url = base_url.rstrip("/")
|
||||||
self.timeout = timeout
|
self.timeout = timeout
|
||||||
self.bearer_token = bearer_token
|
self.bearer_token = bearer_token
|
||||||
|
self._opener = build_opener(NoRedirectHandler())
|
||||||
|
|
||||||
|
def register(
|
||||||
|
self, name: str, capabilities: tuple[str, ...], cpu_count: int, memory_mb: int | None
|
||||||
|
) -> RegisteredWorker:
|
||||||
|
payload: dict[str, Any] = {
|
||||||
|
"name": name,
|
||||||
|
"capabilities": list(capabilities),
|
||||||
|
"cpu_count": cpu_count,
|
||||||
|
}
|
||||||
|
if memory_mb is not None:
|
||||||
|
payload["memory_mb"] = memory_mb
|
||||||
|
status, body = self._request("POST", "/workers/register", payload)
|
||||||
|
if status != 201:
|
||||||
|
raise CoordinatorError(f"worker registration rejected with status {status}")
|
||||||
|
try:
|
||||||
|
return RegisteredWorker.from_json(body)
|
||||||
|
except ValueError as error:
|
||||||
|
raise CoordinatorError("invalid worker registration response") from error
|
||||||
|
|
||||||
def claim(self, worker_id: str, capabilities: tuple[str, ...]) -> ClaimedTask | None:
|
def claim(self, worker_id: str, capabilities: tuple[str, ...]) -> ClaimedTask | None:
|
||||||
status, body = self._request("POST", "/tasks/claim", {
|
status, body = self._request("POST", "/tasks/claim", {
|
||||||
@@ -48,11 +76,15 @@ class HttpCoordinatorClient:
|
|||||||
status, _ = self._request("POST", f"/tasks/{task.task_id}/result", payload)
|
status, _ = self._request("POST", f"/tasks/{task.task_id}/result", payload)
|
||||||
# 200/201/202 include a successful or idempotent duplicate result response.
|
# 200/201/202 include a successful or idempotent duplicate result response.
|
||||||
if status not in (200, 201, 202):
|
if status not in (200, 201, 202):
|
||||||
|
if status == 409:
|
||||||
|
raise CoordinatorConflictError("result rejected because the task lease was lost")
|
||||||
raise CoordinatorError(f"result rejected with status {status}")
|
raise CoordinatorError(f"result rejected with status {status}")
|
||||||
|
|
||||||
def fail(self, task: ClaimedTask, payload: dict[str, Any]) -> None:
|
def fail(self, task: ClaimedTask, payload: dict[str, Any]) -> None:
|
||||||
status, _ = self._request("POST", f"/tasks/{task.task_id}/failure", payload)
|
status, _ = self._request("POST", f"/tasks/{task.task_id}/failure", payload)
|
||||||
if status not in (200, 201, 202):
|
if status not in (200, 201, 202):
|
||||||
|
if status == 409:
|
||||||
|
raise CoordinatorConflictError("failure rejected because the task lease was lost")
|
||||||
raise CoordinatorError(f"failure report rejected with status {status}")
|
raise CoordinatorError(f"failure report rejected with status {status}")
|
||||||
|
|
||||||
def heartbeat(self, task: ClaimedTask, worker_id: str) -> str:
|
def heartbeat(self, task: ClaimedTask, worker_id: str) -> str:
|
||||||
@@ -61,6 +93,8 @@ class HttpCoordinatorClient:
|
|||||||
{"worker_id": worker_id, "attempt": task.attempt},
|
{"worker_id": worker_id, "attempt": task.attempt},
|
||||||
)
|
)
|
||||||
if status != 200:
|
if status != 200:
|
||||||
|
if status == 409:
|
||||||
|
raise CoordinatorConflictError("heartbeat rejected because the task lease was lost")
|
||||||
raise CoordinatorError(f"heartbeat rejected with status {status}")
|
raise CoordinatorError(f"heartbeat rejected with status {status}")
|
||||||
lease_expires_at = body.get("lease_expires_at")
|
lease_expires_at = body.get("lease_expires_at")
|
||||||
if not isinstance(lease_expires_at, str):
|
if not isinstance(lease_expires_at, str):
|
||||||
@@ -73,9 +107,12 @@ class HttpCoordinatorClient:
|
|||||||
headers={"Content-Type": "application/json", **self._auth_header()},
|
headers={"Content-Type": "application/json", **self._auth_header()},
|
||||||
)
|
)
|
||||||
try:
|
try:
|
||||||
with urlopen(request, timeout=self.timeout) as response:
|
with self._opener.open(request, timeout=self.timeout) as response:
|
||||||
raw = response.read()
|
raw = response.read()
|
||||||
return response.status, json.loads(raw) if raw else {}
|
try:
|
||||||
|
return response.status, json.loads(raw) if raw else {}
|
||||||
|
except json.JSONDecodeError as error:
|
||||||
|
raise CoordinatorError("coordinator returned invalid JSON") from error
|
||||||
except HTTPError as error:
|
except HTTPError as error:
|
||||||
if error.code >= 500:
|
if error.code >= 500:
|
||||||
raise CoordinatorTransientError(f"coordinator returned {error.code}") from error
|
raise CoordinatorTransientError(f"coordinator returned {error.code}") from error
|
||||||
|
|||||||
+62
-20
@@ -3,6 +3,7 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
|
from dataclasses import replace
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
import random
|
import random
|
||||||
import shutil
|
import shutil
|
||||||
@@ -12,8 +13,8 @@ from datetime import datetime, timezone
|
|||||||
|
|
||||||
from .artifacts import ArtifactClient, sha256_file
|
from .artifacts import ArtifactClient, sha256_file
|
||||||
from .config import WorkerConfig
|
from .config import WorkerConfig
|
||||||
from .coordinator import CoordinatorClient, CoordinatorTransientError
|
from .coordinator import CoordinatorClient, CoordinatorConflictError, CoordinatorTransientError
|
||||||
from .models import ClaimedTask
|
from .models import ClaimedTask, UploadedArtifact
|
||||||
from .runners import Runner
|
from .runners import Runner
|
||||||
|
|
||||||
|
|
||||||
@@ -32,6 +33,7 @@ class LeaseHeartbeat:
|
|||||||
self._lease_expires_at = self.coordinator.heartbeat(
|
self._lease_expires_at = self.coordinator.heartbeat(
|
||||||
self.task, self.config.worker_id
|
self.task, self.config.worker_id
|
||||||
)
|
)
|
||||||
|
self._next_delay()
|
||||||
self._thread = threading.Thread(target=self._run, name=f"lease-{self.task.task_id}", daemon=True)
|
self._thread = threading.Thread(target=self._run, name=f"lease-{self.task.task_id}", daemon=True)
|
||||||
self._thread.start()
|
self._thread.start()
|
||||||
|
|
||||||
@@ -45,18 +47,19 @@ class LeaseHeartbeat:
|
|||||||
raise self._error
|
raise self._error
|
||||||
|
|
||||||
def _run(self) -> None:
|
def _run(self) -> None:
|
||||||
delay = min(self.config.heartbeat_interval, self._seconds_until_expiry() / 2)
|
delay = self._next_delay()
|
||||||
while not self._stop.wait(max(delay, 0.01)):
|
while not self._stop.wait(max(delay, 0.01)):
|
||||||
try:
|
try:
|
||||||
self._lease_expires_at = self.coordinator.heartbeat(
|
self._lease_expires_at = self.coordinator.heartbeat(
|
||||||
self.task, self.config.worker_id
|
self.task, self.config.worker_id
|
||||||
)
|
)
|
||||||
|
delay = self._next_delay()
|
||||||
except Exception as error: # Surface the lease loss in the main state machine.
|
except Exception as error: # Surface the lease loss in the main state machine.
|
||||||
self._error = error
|
self._error = error
|
||||||
return
|
return
|
||||||
delay = min(
|
|
||||||
self.config.heartbeat_interval, self._seconds_until_expiry() / 2
|
def _next_delay(self) -> float:
|
||||||
)
|
return min(self.config.heartbeat_interval, self._seconds_until_expiry() / 2)
|
||||||
|
|
||||||
def _seconds_until_expiry(self) -> float:
|
def _seconds_until_expiry(self) -> float:
|
||||||
try:
|
try:
|
||||||
@@ -72,12 +75,16 @@ class LeaseHeartbeat:
|
|||||||
class WorkerDaemon:
|
class WorkerDaemon:
|
||||||
def __init__(self, config: WorkerConfig, coordinator: CoordinatorClient, artifacts: ArtifactClient, runner: Runner) -> None:
|
def __init__(self, config: WorkerConfig, coordinator: CoordinatorClient, artifacts: ArtifactClient, runner: Runner) -> None:
|
||||||
self.config, self.coordinator, self.artifacts, self.runner = config, coordinator, artifacts, runner
|
self.config, self.coordinator, self.artifacts, self.runner = config, coordinator, artifacts, runner
|
||||||
|
self.worker_id = config.worker_id
|
||||||
|
self._registered = False
|
||||||
self.log = logging.getLogger("scimesh.worker")
|
self.log = logging.getLogger("scimesh.worker")
|
||||||
|
|
||||||
def run_forever(self) -> None:
|
def run_forever(self) -> None:
|
||||||
failures = 0
|
failures = 0
|
||||||
while True:
|
while True:
|
||||||
try:
|
try:
|
||||||
|
if not self._registered:
|
||||||
|
self._register_worker()
|
||||||
self._cleanup_expired_directories()
|
self._cleanup_expired_directories()
|
||||||
claimed = self.run_once()
|
claimed = self.run_once()
|
||||||
failures = 0
|
failures = 0
|
||||||
@@ -89,16 +96,17 @@ class WorkerDaemon:
|
|||||||
self._sleep(min(self.config.poll_interval * 2 ** min(failures, 6), 60.0))
|
self._sleep(min(self.config.poll_interval * 2 ** min(failures, 6), 60.0))
|
||||||
|
|
||||||
def run_once(self) -> bool:
|
def run_once(self) -> bool:
|
||||||
|
worker_id = self._worker_id()
|
||||||
self._log("claiming")
|
self._log("claiming")
|
||||||
task = self.coordinator.claim(self.config.worker_id, self.config.capabilities)
|
task = self.coordinator.claim(worker_id, self.config.capabilities)
|
||||||
if task is None:
|
if task is None:
|
||||||
self._log("idle")
|
self._log("idle")
|
||||||
return False
|
return False
|
||||||
started = time.monotonic()
|
started = time.monotonic()
|
||||||
task_dir = self.config.work_dir / task.task_id / str(task.attempt)
|
task_dir = self.config.work_dir / task.task_id / str(task.attempt)
|
||||||
task_dir.mkdir(parents=True, exist_ok=False)
|
|
||||||
heartbeat = LeaseHeartbeat(task, self.coordinator, self.config)
|
heartbeat = LeaseHeartbeat(task, self.coordinator, self.config)
|
||||||
try:
|
try:
|
||||||
|
task_dir.mkdir(parents=True, exist_ok=False)
|
||||||
heartbeat.start()
|
heartbeat.start()
|
||||||
self._log("downloading", task)
|
self._log("downloading", task)
|
||||||
input_path = task_dir / "input"
|
input_path = task_dir / "input"
|
||||||
@@ -108,20 +116,28 @@ class WorkerDaemon:
|
|||||||
self._log("running", task)
|
self._log("running", task)
|
||||||
result = self.runner.run(task, task_dir)
|
result = self.runner.run(task, task_dir)
|
||||||
heartbeat.raise_if_failed()
|
heartbeat.raise_if_failed()
|
||||||
manifests = [
|
if len(result.artifacts) != 1:
|
||||||
{
|
raise ValueError("runner must produce exactly one result artifact")
|
||||||
"uri": self.artifacts.upload(task, self.config.worker_id, artifact),
|
artifact = result.artifacts[0]
|
||||||
"sha256": sha256_file(artifact.path),
|
uploaded = self.artifacts.upload(task, worker_id, artifact)
|
||||||
"content_type": artifact.content_type,
|
manifest = self._result_manifest(uploaded)
|
||||||
}
|
|
||||||
for artifact in result.artifacts
|
|
||||||
]
|
|
||||||
if not manifests:
|
|
||||||
raise ValueError("runner produced no artifacts")
|
|
||||||
self._log("submitting", task)
|
self._log("submitting", task)
|
||||||
heartbeat.raise_if_failed()
|
heartbeat.raise_if_failed()
|
||||||
self.coordinator.submit(task, {"worker_id": self.config.worker_id, "attempt": task.attempt, "status": "completed", "result": manifests[0], "artifacts": manifests, "metrics": {**result.metrics, "elapsed_seconds": round(time.monotonic() - started, 3)}})
|
self.coordinator.submit(
|
||||||
|
task,
|
||||||
|
{
|
||||||
|
"worker_id": worker_id,
|
||||||
|
"attempt": task.attempt,
|
||||||
|
"result": manifest,
|
||||||
|
"metrics": {
|
||||||
|
**result.metrics,
|
||||||
|
"elapsed_seconds": round(time.monotonic() - started, 3),
|
||||||
|
},
|
||||||
|
},
|
||||||
|
)
|
||||||
self._log("idle", task, elapsed_seconds=round(time.monotonic() - started, 3))
|
self._log("idle", task, elapsed_seconds=round(time.monotonic() - started, 3))
|
||||||
|
except CoordinatorConflictError as error:
|
||||||
|
self._log("lease_lost", task, error_type=type(error).__name__)
|
||||||
except Exception as error:
|
except Exception as error:
|
||||||
self._log("failed", task, error_type=type(error).__name__)
|
self._log("failed", task, error_type=type(error).__name__)
|
||||||
self._report_failure(task, error)
|
self._report_failure(task, error)
|
||||||
@@ -132,12 +148,38 @@ class WorkerDaemon:
|
|||||||
def _report_failure(self, task: ClaimedTask, error: Exception) -> None:
|
def _report_failure(self, task: ClaimedTask, error: Exception) -> None:
|
||||||
message = str(error).replace(str(self.config.work_dir), "<worker-dir>")[:300]
|
message = str(error).replace(str(self.config.work_dir), "<worker-dir>")[:300]
|
||||||
try:
|
try:
|
||||||
self.coordinator.fail(task, {"worker_id": self.config.worker_id, "attempt": task.attempt, "error_code": type(error).__name__, "error_message": message})
|
self.coordinator.fail(task, {"worker_id": self._worker_id(), "attempt": task.attempt, "error_code": type(error).__name__, "error_message": message})
|
||||||
except CoordinatorTransientError:
|
except CoordinatorTransientError:
|
||||||
raise
|
raise
|
||||||
except Exception:
|
except Exception:
|
||||||
self._log("failed", task, error_type="FailureReportError")
|
self._log("failed", task, error_type="FailureReportError")
|
||||||
|
|
||||||
|
def _register_worker(self) -> None:
|
||||||
|
registered = self.coordinator.register(
|
||||||
|
self.config.worker_name,
|
||||||
|
self.config.capabilities,
|
||||||
|
self.config.cpu_count,
|
||||||
|
self.config.memory_mb,
|
||||||
|
)
|
||||||
|
self.worker_id = registered.worker_id
|
||||||
|
self.config = replace(
|
||||||
|
self.config,
|
||||||
|
worker_id=registered.worker_id,
|
||||||
|
heartbeat_interval=registered.heartbeat_interval_seconds,
|
||||||
|
)
|
||||||
|
self._registered = True
|
||||||
|
self._log("registered")
|
||||||
|
|
||||||
|
def _worker_id(self) -> str:
|
||||||
|
if not self.worker_id:
|
||||||
|
raise ValueError("worker is not registered")
|
||||||
|
return self.worker_id
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _result_manifest(uploaded: UploadedArtifact) -> dict[str, object]:
|
||||||
|
"""Keep completion payload exact: coordinator owns all artifact metadata."""
|
||||||
|
return {"artifact_id": uploaded.artifact_id}
|
||||||
|
|
||||||
def _log(self, state: str, task: ClaimedTask | None = None, **extra: object) -> None:
|
def _log(self, state: str, task: ClaimedTask | None = None, **extra: object) -> None:
|
||||||
fields = {"worker_id": self.config.worker_id, "task_id": task.task_id if task else None, "attempt": task.attempt if task else None, "state": state, **extra}
|
fields = {"worker_id": self.config.worker_id, "task_id": task.task_id if task else None, "attempt": task.attempt if task else None, "state": state, **extra}
|
||||||
self.log.info("worker_event %s", fields)
|
self.log.info("worker_event %s", fields)
|
||||||
|
|||||||
+108
-6
@@ -3,8 +3,40 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
|
from datetime import datetime
|
||||||
|
from math import isfinite
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
from urllib.parse import urlsplit
|
||||||
|
from uuid import UUID
|
||||||
|
|
||||||
|
|
||||||
|
def _required_string(value: object, field: str) -> str:
|
||||||
|
if not isinstance(value, str) or not value.strip():
|
||||||
|
raise ValueError(f"{field} must be a non-empty string")
|
||||||
|
return value
|
||||||
|
|
||||||
|
|
||||||
|
def _coordinator_uri(value: object, field: str) -> str:
|
||||||
|
uri = _required_string(value, field)
|
||||||
|
parsed = urlsplit(uri)
|
||||||
|
if uri.startswith("/"):
|
||||||
|
# ``//host/path`` is a network-path reference: urljoin would resolve
|
||||||
|
# it to another origin. Dot segments are rejected for the same reason
|
||||||
|
# we reject unsafe local task identifiers.
|
||||||
|
if parsed.netloc or any(segment == ".." for segment in parsed.path.split("/")):
|
||||||
|
raise ValueError(f"{field} must be a safe coordinator path")
|
||||||
|
return uri
|
||||||
|
if parsed.scheme not in {"http", "https"} or not parsed.hostname:
|
||||||
|
raise ValueError(f"{field} must be an absolute HTTP(S) URL or coordinator path")
|
||||||
|
return uri
|
||||||
|
|
||||||
|
|
||||||
|
def _sha256(value: object, field: str) -> str:
|
||||||
|
digest = _required_string(value, field).lower()
|
||||||
|
if len(digest) != 64 or any(character not in "0123456789abcdef" for character in digest):
|
||||||
|
raise ValueError(f"{field} must be a SHA-256 hex digest")
|
||||||
|
return digest
|
||||||
|
|
||||||
|
|
||||||
@dataclass(frozen=True)
|
@dataclass(frozen=True)
|
||||||
@@ -26,13 +58,28 @@ class ClaimedTask:
|
|||||||
def from_json(cls, data: dict[str, Any]) -> "ClaimedTask":
|
def from_json(cls, data: dict[str, Any]) -> "ClaimedTask":
|
||||||
try:
|
try:
|
||||||
input_data = data["input"]
|
input_data = data["input"]
|
||||||
|
if not isinstance(input_data, dict):
|
||||||
|
raise ValueError("input must be an object")
|
||||||
|
raw_attempt = data["attempt"]
|
||||||
|
if isinstance(raw_attempt, bool) or not isinstance(raw_attempt, int) or raw_attempt < 1:
|
||||||
|
raise ValueError("attempt must be a positive integer")
|
||||||
|
task_id = str(UUID(_required_string(data["task_id"], "task_id")))
|
||||||
|
lease_expires_at = _required_string(data["lease_expires_at"], "lease_expires_at")
|
||||||
|
if datetime.fromisoformat(lease_expires_at.replace("Z", "+00:00")).tzinfo is None:
|
||||||
|
raise ValueError("lease_expires_at must include a timezone")
|
||||||
|
parameters = data.get("parameters", {})
|
||||||
|
if not isinstance(parameters, dict):
|
||||||
|
raise ValueError("parameters must be an object")
|
||||||
return cls(
|
return cls(
|
||||||
task_id=str(data["task_id"]),
|
task_id=task_id,
|
||||||
attempt=int(data["attempt"]),
|
attempt=raw_attempt,
|
||||||
lease_expires_at=str(data["lease_expires_at"]),
|
lease_expires_at=lease_expires_at,
|
||||||
workload=str(data["workload"]),
|
workload=_required_string(data["workload"], "workload"),
|
||||||
input=InputArtifact(uri=str(input_data["uri"]), sha256=str(input_data["sha256"])),
|
input=InputArtifact(
|
||||||
parameters=dict(data.get("parameters", {})),
|
uri=_coordinator_uri(input_data["uri"], "input.uri"),
|
||||||
|
sha256=_sha256(input_data["sha256"], "input.sha256"),
|
||||||
|
),
|
||||||
|
parameters=parameters,
|
||||||
)
|
)
|
||||||
except (KeyError, TypeError, ValueError) as error:
|
except (KeyError, TypeError, ValueError) as error:
|
||||||
raise ValueError("invalid claimed-task response") from error
|
raise ValueError("invalid claimed-task response") from error
|
||||||
@@ -44,6 +91,61 @@ class ProducedArtifact:
|
|||||||
content_type: str
|
content_type: str
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class UploadedArtifact:
|
||||||
|
"""Coordinator-owned artifact metadata returned after a successful upload."""
|
||||||
|
|
||||||
|
artifact_id: str
|
||||||
|
uri: str
|
||||||
|
sha256: str
|
||||||
|
size_bytes: int
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_json(cls, data: object) -> "UploadedArtifact":
|
||||||
|
if not isinstance(data, dict):
|
||||||
|
raise ValueError("artifact upload response must be an object")
|
||||||
|
raw_size = data.get("size_bytes")
|
||||||
|
if isinstance(raw_size, bool) or not isinstance(raw_size, int) or raw_size < 0:
|
||||||
|
raise ValueError("artifact size_bytes must be a non-negative integer")
|
||||||
|
try:
|
||||||
|
return cls(
|
||||||
|
artifact_id=str(UUID(_required_string(data.get("artifact_id"), "artifact_id"))),
|
||||||
|
uri=_coordinator_uri(data.get("uri"), "uri"),
|
||||||
|
sha256=_sha256(data.get("sha256"), "sha256"),
|
||||||
|
size_bytes=raw_size,
|
||||||
|
)
|
||||||
|
except ValueError as error:
|
||||||
|
raise ValueError("invalid artifact upload response") from error
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class RegisteredWorker:
|
||||||
|
"""Identity and heartbeat policy returned by worker registration."""
|
||||||
|
|
||||||
|
worker_id: str
|
||||||
|
heartbeat_interval_seconds: float
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_json(cls, data: object) -> "RegisteredWorker":
|
||||||
|
if not isinstance(data, dict):
|
||||||
|
raise ValueError("worker registration response must be an object")
|
||||||
|
raw_interval = data.get("heartbeat_interval_seconds")
|
||||||
|
if (
|
||||||
|
isinstance(raw_interval, bool)
|
||||||
|
or not isinstance(raw_interval, (int, float))
|
||||||
|
or not isfinite(raw_interval)
|
||||||
|
or raw_interval <= 0
|
||||||
|
):
|
||||||
|
raise ValueError("heartbeat_interval_seconds must be positive")
|
||||||
|
try:
|
||||||
|
return cls(
|
||||||
|
worker_id=str(UUID(_required_string(data.get("worker_id"), "worker_id"))),
|
||||||
|
heartbeat_interval_seconds=float(raw_interval),
|
||||||
|
)
|
||||||
|
except ValueError as error:
|
||||||
|
raise ValueError("invalid worker registration response") from error
|
||||||
|
|
||||||
|
|
||||||
@dataclass(frozen=True)
|
@dataclass(frozen=True)
|
||||||
class RunResult:
|
class RunResult:
|
||||||
artifacts: tuple[ProducedArtifact, ...]
|
artifacts: tuple[ProducedArtifact, ...]
|
||||||
|
|||||||
@@ -20,9 +20,13 @@ class SciMeshRunner:
|
|||||||
def run(self, task: ClaimedTask, task_dir: Path) -> RunResult:
|
def run(self, task: ClaimedTask, task_dir: Path) -> RunResult:
|
||||||
input_path = task_dir / "input"
|
input_path = task_dir / "input"
|
||||||
output_path = task_dir / "result.csv"
|
output_path = task_dir / "result.csv"
|
||||||
command = [sys.executable, "-m", "scimesh.cli", task.workload, str(input_path)]
|
# The coordinator contract historically used underscores while the
|
||||||
|
# public SciMesh CLI uses hyphens. Accept both spellings at this narrow
|
||||||
|
# boundary so an API job cannot turn into an opaque worker failure.
|
||||||
|
workload = task.workload.replace("_", "-")
|
||||||
|
command = [sys.executable, "-m", "scimesh.cli", workload, str(input_path)]
|
||||||
params = task.parameters
|
params = task.parameters
|
||||||
if task.workload == "similarity-search":
|
if workload == "similarity-search":
|
||||||
self._reject_unknown(params, {"query_id", "query_smiles", "top_k", "threshold", "threshold_direction", "max_rows", "progress_every"})
|
self._reject_unknown(params, {"query_id", "query_smiles", "top_k", "threshold", "threshold_direction", "max_rows", "progress_every"})
|
||||||
query_id, query_smiles = params.get("query_id"), params.get("query_smiles")
|
query_id, query_smiles = params.get("query_id"), params.get("query_smiles")
|
||||||
if (query_id is None) == (query_smiles is None):
|
if (query_id is None) == (query_smiles is None):
|
||||||
@@ -31,7 +35,7 @@ class SciMeshRunner:
|
|||||||
command += ["--query-id", self._string(params, "query_id")] if query_id is not None else ["--query-smiles", self._string(params, "query_smiles")]
|
command += ["--query-id", self._string(params, "query_id")] if query_id is not None else ["--query-smiles", self._string(params, "query_smiles")]
|
||||||
command += ["--top-k", str(top_k)]
|
command += ["--top-k", str(top_k)]
|
||||||
self._append_common_options(command, params)
|
self._append_common_options(command, params)
|
||||||
elif task.workload == "similarity-graph":
|
elif workload == "similarity-graph":
|
||||||
self._reject_unknown(params, {"threshold", "threshold_direction", "block_size", "max_rows", "progress_every"})
|
self._reject_unknown(params, {"threshold", "threshold_direction", "block_size", "max_rows", "progress_every"})
|
||||||
threshold = self._number(params, "threshold")
|
threshold = self._number(params, "threshold")
|
||||||
command += ["--threshold", str(threshold)]
|
command += ["--threshold", str(threshold)]
|
||||||
|
|||||||
@@ -0,0 +1,51 @@
|
|||||||
|
"""Small HTTP transport helpers shared by coordinator and artifact clients."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from urllib.request import HTTPRedirectHandler, Request
|
||||||
|
from urllib.parse import urlsplit
|
||||||
|
|
||||||
|
|
||||||
|
def origin(uri: str) -> tuple[str, str, int | None]:
|
||||||
|
"""Return a normalized HTTP origin for authorization decisions."""
|
||||||
|
parsed = urlsplit(uri)
|
||||||
|
scheme = parsed.scheme.lower()
|
||||||
|
default_port = {"http": 80, "https": 443}.get(scheme)
|
||||||
|
return scheme, (parsed.hostname or "").lower(), parsed.port or default_port
|
||||||
|
|
||||||
|
|
||||||
|
class SameOriginAuthRedirectHandler(HTTPRedirectHandler):
|
||||||
|
"""Strip coordinator authorization when an artifact redirect changes origin."""
|
||||||
|
|
||||||
|
def __init__(self, coordinator_origin: tuple[str, str, int | None]) -> None:
|
||||||
|
super().__init__()
|
||||||
|
self.coordinator_origin = coordinator_origin
|
||||||
|
|
||||||
|
def redirect_request(
|
||||||
|
self,
|
||||||
|
req: Request,
|
||||||
|
fp: object,
|
||||||
|
code: int,
|
||||||
|
msg: str,
|
||||||
|
headers: object,
|
||||||
|
newurl: str,
|
||||||
|
) -> Request | None:
|
||||||
|
redirected = super().redirect_request(req, fp, code, msg, headers, newurl)
|
||||||
|
if redirected and origin(newurl) != self.coordinator_origin:
|
||||||
|
redirected.remove_header("Authorization")
|
||||||
|
return redirected
|
||||||
|
|
||||||
|
|
||||||
|
class NoRedirectHandler(HTTPRedirectHandler):
|
||||||
|
"""Reject redirects for mutating coordinator API calls."""
|
||||||
|
|
||||||
|
def redirect_request(
|
||||||
|
self,
|
||||||
|
req: Request,
|
||||||
|
fp: object,
|
||||||
|
code: int,
|
||||||
|
msg: str,
|
||||||
|
headers: object,
|
||||||
|
newurl: str,
|
||||||
|
) -> Request | None:
|
||||||
|
return None
|
||||||
@@ -47,10 +47,15 @@ def _fingerprinted_molecules(
|
|||||||
tsv_path: Path, max_rows: int | None
|
tsv_path: Path, max_rows: int | None
|
||||||
) -> tuple[list[GraphMolecule], DatasetStats]:
|
) -> tuple[list[GraphMolecule], DatasetStats]:
|
||||||
stats = DatasetStats()
|
stats = DatasetStats()
|
||||||
molecules = [
|
molecules: list[GraphMolecule] = []
|
||||||
GraphMolecule(record.molecule_id, fingerprint(record.molecule))
|
seen_ids: set[str] = set()
|
||||||
for record in iter_valid_molecules(tsv_path, stats, max_rows=max_rows)
|
for record in iter_valid_molecules(tsv_path, stats, max_rows=max_rows):
|
||||||
]
|
if not record.molecule_id:
|
||||||
|
raise ValueError("Dataset contains an empty chembl_id")
|
||||||
|
if record.molecule_id in seen_ids:
|
||||||
|
raise ValueError(f"Dataset contains a duplicate chembl_id: {record.molecule_id}")
|
||||||
|
seen_ids.add(record.molecule_id)
|
||||||
|
molecules.append(GraphMolecule(record.molecule_id, fingerprint(record.molecule)))
|
||||||
return molecules, stats
|
return molecules, stats
|
||||||
|
|
||||||
|
|
||||||
@@ -119,6 +124,7 @@ def build_similarity_graph(
|
|||||||
|
|
||||||
def write_graph_edges(output_path: Path, edges: list[SimilarityEdge]) -> None:
|
def write_graph_edges(output_path: Path, edges: list[SimilarityEdge]) -> None:
|
||||||
"""Write a deterministic sparse edge list CSV."""
|
"""Write a deterministic sparse edge list CSV."""
|
||||||
|
output_path.parent.mkdir(parents=True, exist_ok=True)
|
||||||
with output_path.open("w", encoding="utf-8", newline="") as destination:
|
with output_path.open("w", encoding="utf-8", newline="") as destination:
|
||||||
writer = csv.DictWriter(destination, fieldnames=["source_id", "target_id", "similarity"])
|
writer = csv.DictWriter(destination, fieldnames=["source_id", "target_id", "similarity"])
|
||||||
writer.writeheader()
|
writer.writeheader()
|
||||||
|
|||||||
@@ -135,6 +135,7 @@ def search_similar(
|
|||||||
|
|
||||||
def write_search_results(output_path: Path, matches: list[SimilarityMatch]) -> None:
|
def write_search_results(output_path: Path, matches: list[SimilarityMatch]) -> None:
|
||||||
"""Write ranked matches to a deterministic CSV file."""
|
"""Write ranked matches to a deterministic CSV file."""
|
||||||
|
output_path.parent.mkdir(parents=True, exist_ok=True)
|
||||||
with output_path.open("w", encoding="utf-8", newline="") as destination:
|
with output_path.open("w", encoding="utf-8", newline="") as destination:
|
||||||
writer = csv.DictWriter(
|
writer = csv.DictWriter(
|
||||||
destination,
|
destination,
|
||||||
|
|||||||
@@ -2,6 +2,7 @@ from __future__ import annotations
|
|||||||
|
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
|
import pytest
|
||||||
from rdkit import DataStructs
|
from rdkit import DataStructs
|
||||||
|
|
||||||
from scimesh.chemistry.dataset import DatasetStats, iter_valid_molecules
|
from scimesh.chemistry.dataset import DatasetStats, iter_valid_molecules
|
||||||
@@ -61,3 +62,19 @@ def test_graph_supports_less_than_threshold_direction(small_dataset: Path) -> No
|
|||||||
)
|
)
|
||||||
|
|
||||||
assert all(edge.similarity <= 0.15 for edge in result.edges)
|
assert all(edge.similarity <= 0.15 for edge in result.edges)
|
||||||
|
|
||||||
|
|
||||||
|
def test_graph_rejects_duplicate_identifiers(tmp_path: Path) -> None:
|
||||||
|
dataset = tmp_path / "duplicate_ids.tsv"
|
||||||
|
dataset.write_text(
|
||||||
|
"chembl_id\tcanonical_smiles\nDUP\tCCO\nDUP\tCCC\n", encoding="utf-8"
|
||||||
|
)
|
||||||
|
|
||||||
|
with pytest.raises(ValueError, match="duplicate chembl_id"):
|
||||||
|
build_similarity_graph(dataset, threshold=0.1, block_size=1)
|
||||||
|
|
||||||
|
|
||||||
|
def test_graph_writer_creates_missing_output_directory(tmp_path: Path) -> None:
|
||||||
|
output = tmp_path / "nested" / "edges.csv"
|
||||||
|
write_graph_edges(output, [])
|
||||||
|
assert output.read_text(encoding="utf-8").startswith("source_id,target_id")
|
||||||
|
|||||||
@@ -6,7 +6,11 @@ from rdkit import Chem, DataStructs
|
|||||||
|
|
||||||
from scimesh.chemistry.dataset import DatasetStats, find_molecule_by_id, iter_valid_molecules
|
from scimesh.chemistry.dataset import DatasetStats, find_molecule_by_id, iter_valid_molecules
|
||||||
from scimesh.chemistry.fingerprints import fingerprint
|
from scimesh.chemistry.fingerprints import fingerprint
|
||||||
from scimesh.workloads.similarity_search import SimilarityMatch, search_similar
|
from scimesh.workloads.similarity_search import (
|
||||||
|
SimilarityMatch,
|
||||||
|
search_similar,
|
||||||
|
write_search_results,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def test_search_matches_full_sorting_and_skips_query_and_invalid(
|
def test_search_matches_full_sorting_and_skips_query_and_invalid(
|
||||||
@@ -56,3 +60,9 @@ def test_search_can_rank_and_filter_least_similar_molecules(
|
|||||||
assert result.matches == sorted(
|
assert result.matches == sorted(
|
||||||
result.matches, key=lambda match: match.sort_key("less")
|
result.matches, key=lambda match: match.sort_key("less")
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_search_writer_creates_missing_output_directory(tmp_path: Path) -> None:
|
||||||
|
output = tmp_path / "nested" / "results.csv"
|
||||||
|
write_search_results(output, [])
|
||||||
|
assert output.read_text(encoding="utf-8").startswith("rank,chembl_id")
|
||||||
|
|||||||
+141
-6
@@ -11,9 +11,17 @@ import pytest
|
|||||||
from scimesh.worker.config import WorkerConfig
|
from scimesh.worker.config import WorkerConfig
|
||||||
from scimesh.worker.coordinator import CoordinatorTransientError
|
from scimesh.worker.coordinator import CoordinatorTransientError
|
||||||
from scimesh.worker.daemon import LeaseHeartbeat, WorkerDaemon
|
from scimesh.worker.daemon import LeaseHeartbeat, WorkerDaemon
|
||||||
from scimesh.worker.models import ClaimedTask, InputArtifact, ProducedArtifact, RunResult
|
from scimesh.worker.models import (
|
||||||
|
ClaimedTask,
|
||||||
|
InputArtifact,
|
||||||
|
ProducedArtifact,
|
||||||
|
RegisteredWorker,
|
||||||
|
RunResult,
|
||||||
|
UploadedArtifact,
|
||||||
|
)
|
||||||
from scimesh.worker.artifacts import HttpArtifactClient, _SameOriginAuthRedirectHandler, _origin
|
from scimesh.worker.artifacts import HttpArtifactClient, _SameOriginAuthRedirectHandler, _origin
|
||||||
from scimesh.worker.runners import SciMeshRunner
|
from scimesh.worker.runners import SciMeshRunner
|
||||||
|
from scimesh.worker.transport import NoRedirectHandler
|
||||||
|
|
||||||
|
|
||||||
class FakeCoordinator:
|
class FakeCoordinator:
|
||||||
@@ -24,6 +32,11 @@ class FakeCoordinator:
|
|||||||
task, self.task = self.task, None
|
task, self.task = self.task, None
|
||||||
return task
|
return task
|
||||||
|
|
||||||
|
def register(
|
||||||
|
self, name: str, capabilities: tuple[str, ...], cpu_count: int, memory_mb: int | None
|
||||||
|
) -> RegisteredWorker:
|
||||||
|
return RegisteredWorker("11111111-1111-4111-8111-111111111111", 15)
|
||||||
|
|
||||||
def submit(self, task: ClaimedTask, payload: dict) -> None:
|
def submit(self, task: ClaimedTask, payload: dict) -> None:
|
||||||
self.submissions.append(payload)
|
self.submissions.append(payload)
|
||||||
|
|
||||||
@@ -42,9 +55,17 @@ class FakeArtifacts:
|
|||||||
def download(self, uri: str, destination: Path) -> None:
|
def download(self, uri: str, destination: Path) -> None:
|
||||||
destination.write_bytes(self.content)
|
destination.write_bytes(self.content)
|
||||||
|
|
||||||
def upload(self, task: ClaimedTask, worker_id: str, artifact: ProducedArtifact) -> str:
|
def upload(
|
||||||
|
self, task: ClaimedTask, worker_id: str, artifact: ProducedArtifact
|
||||||
|
) -> UploadedArtifact:
|
||||||
self.uploaded.append((task.task_id, worker_id, artifact.path))
|
self.uploaded.append((task.task_id, worker_id, artifact.path))
|
||||||
return f"https://example.test/tasks/{task.task_id}/artifacts/{artifact.path.name}"
|
content = artifact.path.read_bytes()
|
||||||
|
return UploadedArtifact(
|
||||||
|
"22222222-2222-4222-8222-222222222222",
|
||||||
|
f"https://example.test/tasks/{task.task_id}/artifacts/{artifact.path.name}",
|
||||||
|
hashlib.sha256(content).hexdigest(),
|
||||||
|
len(content),
|
||||||
|
)
|
||||||
|
|
||||||
class FakeRunner:
|
class FakeRunner:
|
||||||
def __init__(self) -> None:
|
def __init__(self) -> None:
|
||||||
@@ -75,9 +96,10 @@ def test_claims_runs_uploads_and_submits_csv(tmp_path: Path) -> None:
|
|||||||
assert runner.calls == 1
|
assert runner.calls == 1
|
||||||
assert len(artifacts.uploaded) == 1
|
assert len(artifacts.uploaded) == 1
|
||||||
assert coordinator.heartbeats == [("task-1", 1, "worker-1")]
|
assert coordinator.heartbeats == [("task-1", 1, "worker-1")]
|
||||||
assert coordinator.submissions[0]["status"] == "completed"
|
assert "status" not in coordinator.submissions[0]
|
||||||
assert coordinator.submissions[0]["result"]["content_type"] == "text/csv"
|
assert coordinator.submissions[0]["result"] == {
|
||||||
assert coordinator.submissions[0]["result"]["uri"].startswith("https://example.test/tasks/task-1/artifacts/")
|
"artifact_id": "22222222-2222-4222-8222-222222222222"
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
def test_no_task_does_not_create_directory(tmp_path: Path) -> None:
|
def test_no_task_does_not_create_directory(tmp_path: Path) -> None:
|
||||||
@@ -95,6 +117,14 @@ def test_bad_checksum_reports_failure_without_running(tmp_path: Path) -> None:
|
|||||||
assert not coordinator.submissions
|
assert not coordinator.submissions
|
||||||
|
|
||||||
|
|
||||||
|
def test_directory_creation_failure_is_reported(tmp_path: Path) -> None:
|
||||||
|
content = b"input fixture"
|
||||||
|
worker, coordinator, _, _, config = daemon(tmp_path, make_task(content), content)
|
||||||
|
(config.work_dir / "task-1" / "1").mkdir(parents=True)
|
||||||
|
assert worker.run_once() is True
|
||||||
|
assert coordinator.failures[0]["error_code"] == "FileExistsError"
|
||||||
|
|
||||||
|
|
||||||
def test_transient_claim_error_is_propagated_for_bounded_backoff(tmp_path: Path) -> None:
|
def test_transient_claim_error_is_propagated_for_bounded_backoff(tmp_path: Path) -> None:
|
||||||
class UnavailableCoordinator(FakeCoordinator):
|
class UnavailableCoordinator(FakeCoordinator):
|
||||||
def claim(self, worker_id: str, capabilities: tuple[str, ...]) -> ClaimedTask | None:
|
def claim(self, worker_id: str, capabilities: tuple[str, ...]) -> ClaimedTask | None:
|
||||||
@@ -123,6 +153,15 @@ def test_input_token_is_sent_only_to_the_coordinator_origin() -> None:
|
|||||||
assert client._auth_headers_for("https://bucket.example/presigned") == {}
|
assert client._auth_headers_for("https://bucket.example/presigned") == {}
|
||||||
|
|
||||||
|
|
||||||
|
def test_relative_input_uri_is_resolved_against_the_coordinator() -> None:
|
||||||
|
client = HttpArtifactClient("https://coordinator.example/api", 10, "secret")
|
||||||
|
assert client._auth_headers_for("https://coordinator.example/tasks/1/input") == {
|
||||||
|
"Authorization": "Bearer secret"
|
||||||
|
}
|
||||||
|
# The coordinator's contract returns root-relative artifact paths.
|
||||||
|
assert client.coordinator_url == "https://coordinator.example/api"
|
||||||
|
|
||||||
|
|
||||||
def test_redirect_to_external_storage_strips_authorization() -> None:
|
def test_redirect_to_external_storage_strips_authorization() -> None:
|
||||||
handler = _SameOriginAuthRedirectHandler(_origin("https://coordinator.example"))
|
handler = _SameOriginAuthRedirectHandler(_origin("https://coordinator.example"))
|
||||||
source = Request(
|
source = Request(
|
||||||
@@ -133,6 +172,12 @@ def test_redirect_to_external_storage_strips_authorization() -> None:
|
|||||||
assert redirected.get_header("Authorization") is None
|
assert redirected.get_header("Authorization") is None
|
||||||
|
|
||||||
|
|
||||||
|
def test_api_requests_never_follow_redirects() -> None:
|
||||||
|
handler = NoRedirectHandler()
|
||||||
|
request = Request("https://coordinator.example/tasks/claim", headers={"Authorization": "Bearer secret"})
|
||||||
|
assert handler.redirect_request(request, None, 302, "Found", {}, "https://other.example") is None
|
||||||
|
|
||||||
|
|
||||||
def test_lease_is_renewed_while_a_runner_is_still_working(tmp_path: Path) -> None:
|
def test_lease_is_renewed_while_a_runner_is_still_working(tmp_path: Path) -> None:
|
||||||
content = b"input fixture"
|
content = b"input fixture"
|
||||||
worker, coordinator, _, _, config = daemon(tmp_path, make_task(content), content)
|
worker, coordinator, _, _, config = daemon(tmp_path, make_task(content), content)
|
||||||
@@ -184,3 +229,93 @@ def test_runner_maps_graph_and_smiles_search_parameters(tmp_path: Path, monkeypa
|
|||||||
assert "--block-size" in commands[0] and "42" in commands[0]
|
assert "--block-size" in commands[0] and "42" in commands[0]
|
||||||
assert "--max-rows" in commands[0] and "7" in commands[0]
|
assert "--max-rows" in commands[0] and "7" in commands[0]
|
||||||
assert "--query-smiles" in commands[1] and "CCO" in commands[1]
|
assert "--query-smiles" in commands[1] and "CCO" in commands[1]
|
||||||
|
|
||||||
|
|
||||||
|
def test_runner_accepts_coordinator_workload_names(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||||
|
commands: list[list[str]] = []
|
||||||
|
|
||||||
|
def fake_run(command: list[str], **_: object) -> None:
|
||||||
|
commands.append(command)
|
||||||
|
output = Path(command[command.index("--output") + 1])
|
||||||
|
output.parent.mkdir(parents=True, exist_ok=True)
|
||||||
|
output.write_text("a,b\\n", encoding="utf-8")
|
||||||
|
|
||||||
|
monkeypatch.setattr("scimesh.worker.runners.subprocess.run", fake_run)
|
||||||
|
task = ClaimedTask(
|
||||||
|
"search", 1, "2026-07-30T00:00:00Z", "similarity_search",
|
||||||
|
InputArtifact("https://example/input", "a" * 64), {"query_smiles": "CCO"},
|
||||||
|
)
|
||||||
|
SciMeshRunner().run(task, tmp_path / "search")
|
||||||
|
assert commands[0][3] == "similarity-search"
|
||||||
|
|
||||||
|
|
||||||
|
def test_claimed_task_rejects_path_traversal_and_invalid_metadata() -> None:
|
||||||
|
payload = {
|
||||||
|
"task_id": "../outside",
|
||||||
|
"attempt": 1,
|
||||||
|
"lease_expires_at": "2026-07-30T00:00:00Z",
|
||||||
|
"workload": "similarity-search",
|
||||||
|
"input": {"uri": "https://example.test/input", "sha256": "a" * 64},
|
||||||
|
"parameters": {},
|
||||||
|
}
|
||||||
|
with pytest.raises(ValueError, match="invalid claimed-task response"):
|
||||||
|
ClaimedTask.from_json(payload)
|
||||||
|
|
||||||
|
payload["task_id"] = "11111111-1111-4111-8111-111111111111"
|
||||||
|
payload["input"] = {"uri": "//outside.example/input", "sha256": "a" * 64}
|
||||||
|
with pytest.raises(ValueError, match="invalid claimed-task response"):
|
||||||
|
ClaimedTask.from_json(payload)
|
||||||
|
|
||||||
|
payload["input"] = {"uri": "/tasks/../outside/input", "sha256": "a" * 64}
|
||||||
|
with pytest.raises(ValueError, match="invalid claimed-task response"):
|
||||||
|
ClaimedTask.from_json(payload)
|
||||||
|
|
||||||
|
|
||||||
|
def test_claimed_task_accepts_a_coordinator_relative_input_path() -> None:
|
||||||
|
task = ClaimedTask.from_json(
|
||||||
|
{
|
||||||
|
"task_id": "11111111-1111-4111-8111-111111111111",
|
||||||
|
"attempt": 1,
|
||||||
|
"lease_expires_at": "2026-07-30T00:00:00Z",
|
||||||
|
"workload": "similarity_search",
|
||||||
|
"input": {"uri": "/tasks/11111111-1111-4111-8111-111111111111/input", "sha256": "a" * 64},
|
||||||
|
"parameters": {},
|
||||||
|
}
|
||||||
|
)
|
||||||
|
assert task.input.uri.startswith("/tasks/")
|
||||||
|
|
||||||
|
|
||||||
|
def test_uploaded_artifact_requires_complete_durable_metadata() -> None:
|
||||||
|
artifact = UploadedArtifact.from_json(
|
||||||
|
{
|
||||||
|
"artifact_id": "22222222-2222-4222-8222-222222222222",
|
||||||
|
"uri": "https://coordinator.example/artifacts/222/download",
|
||||||
|
"sha256": "a" * 64,
|
||||||
|
"size_bytes": 12,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
assert artifact.size_bytes == 12
|
||||||
|
with pytest.raises(ValueError, match="artifact size_bytes"):
|
||||||
|
UploadedArtifact.from_json({"artifact_id": "missing"})
|
||||||
|
|
||||||
|
|
||||||
|
def test_environment_overrides_allow_cli_only_configuration(monkeypatch: pytest.MonkeyPatch, tmp_path: Path) -> None:
|
||||||
|
monkeypatch.delenv("SCIMESH_COORDINATOR_URL", raising=False)
|
||||||
|
config = WorkerConfig.from_environment(
|
||||||
|
{
|
||||||
|
"coordinator_url": "https://coordinator.example",
|
||||||
|
"work_dir": tmp_path,
|
||||||
|
"worker_name": "test-worker",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
assert config.coordinator_url == "https://coordinator.example"
|
||||||
|
assert config.worker_id is None
|
||||||
|
assert "similarity-search" in config.capabilities
|
||||||
|
assert "similarity_search" in config.capabilities
|
||||||
|
|
||||||
|
|
||||||
|
def test_worker_registration_sets_returned_identity(tmp_path: Path) -> None:
|
||||||
|
worker, _, _, _, _ = daemon(tmp_path, None, b"")
|
||||||
|
worker._register_worker()
|
||||||
|
assert worker.worker_id == "11111111-1111-4111-8111-111111111111"
|
||||||
|
assert worker.config.heartbeat_interval == 15
|
||||||
|
|||||||
Reference in New Issue
Block a user