diff --git a/spindle/agentproto/protocol.go b/spindle/agentproto/protocol.go index ac12e9094..4d4a904c8 100644 --- a/spindle/agentproto/protocol.go +++ b/spindle/agentproto/protocol.go @@ -20,6 +20,13 @@ const ( type Message = agentv1.Message +func ValidateProtocolVersion(got uint32) error { + if got != ProtocolVersion { + return fmt.Errorf("agent protocol version mismatch: got %d, want %d", got, ProtocolVersion) + } + return nil +} + var validator protovalidate.Validator func init() { diff --git a/spindle/agentproto/protocol_test.go b/spindle/agentproto/protocol_test.go index ccc7dc946..0dbe80d85 100644 --- a/spindle/agentproto/protocol_test.go +++ b/spindle/agentproto/protocol_test.go @@ -55,3 +55,12 @@ func TestValidation(t *testing.T) { t.Fatal("expected message with multiple payloads to fail validation") } } + +func TestValidateProtocolVersion(t *testing.T) { + if err := ValidateProtocolVersion(ProtocolVersion); err != nil { + t.Fatalf("current protocol rejected: %v", err) + } + if err := ValidateProtocolVersion(ProtocolVersion + 1); err == nil { + t.Fatal("mismatched protocol accepted") + } +} diff --git a/spindle/db/events.go b/spindle/db/events.go index 2166c8cf9..c89b62208 100644 --- a/spindle/db/events.go +++ b/spindle/db/events.go @@ -75,22 +75,6 @@ func (d *DB) createStatusEvent( return d.insertEvent(event, n) } -// stamps the mill's own clock so it orders against the cursor like a local write -func (d *DB) InsertEventStatus( - pipelineAtUri string, - workflow string, - status string, - workflowError *string, - exitCode *int64, - n *notifier.Notifier, -) error { - event, err := statusEvent(pipelineAtUri, workflow, status, workflowError, exitCode) - if err != nil { - return err - } - return d.insertEvent(event, n) -} - // deleting the lease in the same transaction prevents the terminal event // from replaying func (d *DB) CompleteMillLease( diff --git a/spindle/engine/engine.go b/spindle/engine/engine.go index 87095c2b5..e25b0704a 100644 --- a/spindle/engine/engine.go +++ b/spindle/engine/engine.go @@ -205,25 +205,6 @@ func workflowResult(ctx context.Context, err error) string { return "success" } -func reportWorkflowStatusError(l *slog.Logger, database *db.DB, n *notifier.Notifier, wid models.WorkflowId, err error) { - if errors.Is(err, ErrTimedOut) { - dbErr := database.StatusTimeout(wid, n) - if dbErr != nil { - l.Error("failed to set workflow status to timeout", "wid", wid, "err", dbErr) - } - } else if errors.Is(err, ErrWorkflowCanceled) { - dbErr := database.StatusCancelled(wid, err.Error(), -1, n) - if dbErr != nil { - l.Error("failed to set workflow status to cancelled", "wid", wid, "err", dbErr) - } - } else { - dbErr := database.StatusFailed(wid, err.Error(), -1, n) - if dbErr != nil { - l.Error("failed to set workflow status to failed", "wid", wid, "err", dbErr) - } - } -} - func StartWorkflows(l *slog.Logger, vault secrets.Manager, cfg *config.Config, qm *quota.Manager, stores *artifactstore.Stores, db *db.DB, n *notifier.Notifier, ctx context.Context, pipeline *models.Pipeline, pipelineId models.PipelineId) { var allSecrets []secrets.UnlockedSecret diff --git a/spindle/engines/microvm/agent.go b/spindle/engines/microvm/agent.go index 23db3e51c..4a12e7742 100644 --- a/spindle/engines/microvm/agent.go +++ b/spindle/engines/microvm/agent.go @@ -159,6 +159,9 @@ func (s *AgentSession) Init(ctx context.Context, init *agentv1.Init) error { if helloPayload == nil { return fmt.Errorf("expected agent hello, got nil") } + if err := agentproto.ValidateProtocolVersion(helloPayload.ProtocolVersion); err != nil { + return err + } s.l.Info("agent connected", "protocol", helloPayload.ProtocolVersion, "version", helloPayload.AgentVersion, "boot", helloPayload.BootId, "nix", helloPayload.NixVersion) if err := s.enc.Encode(&agentproto.Message{ diff --git a/spindle/engines/microvm/agent_test.go b/spindle/engines/microvm/agent_test.go new file mode 100644 index 000000000..b37025794 --- /dev/null +++ b/spindle/engines/microvm/agent_test.go @@ -0,0 +1,64 @@ +//go:build linux + +package microvm + +import ( + "context" + "io" + "log/slog" + "net" + "strings" + "testing" + "time" + + "tangled.org/core/spindle/agentproto" + agentv1 "tangled.org/core/spindle/agentproto/gen" +) + +func TestAgentSessionInitProtocolVersion(t *testing.T) { + for _, test := range []struct { + name string + version uint32 + wantErr bool + }{ + {"current", agentproto.ProtocolVersion, false}, + {"stale", agentproto.ProtocolVersion + 1, true}, + } { + t.Run(test.name, func(t *testing.T) { + host, guest := net.Pipe() + defer host.Close() + defer guest.Close() + session := NewAgentSession(host, slog.New(slog.NewTextHandler(io.Discard, nil))) + guestErr := make(chan error, 1) + go func() { + encoder := agentproto.NewEncoder(guest) + if err := encoder.Encode(&agentproto.Message{Id: "hello", Hello: &agentv1.Hello{ProtocolVersion: test.version}}); err != nil { + guestErr <- err + return + } + if test.wantErr { + guestErr <- nil + return + } + message, err := agentproto.NewDecoder(guest).Decode() + if err == nil && message.Init == nil { + err = io.ErrUnexpectedEOF + } + guestErr <- err + }() + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + err := session.Init(ctx, &agentv1.Init{JobId: "job"}) + if test.wantErr { + if err == nil || !strings.Contains(err.Error(), "protocol version mismatch") { + t.Fatalf("Init error = %v", err) + } + } else if err != nil { + t.Fatalf("Init failed: %v", err) + } + if err := <-guestErr; err != nil { + t.Fatalf("guest side failed: %v", err) + } + }) + } +} diff --git a/spindle/mill/executor/executor.go b/spindle/mill/executor/executor.go index e96804394..df229876f 100644 --- a/spindle/mill/executor/executor.go +++ b/spindle/mill/executor/executor.go @@ -181,6 +181,7 @@ func (e *Executor) Connect(ctx context.Context) { defer e.n.Unsubscribe(sub) backoff := dialBackoffMin +reconnect: for { if ctx.Err() != nil { break @@ -192,7 +193,7 @@ func (e *Executor) Connect(ctx context.Context) { e.l.Warn("mill session ended; reconnecting", "err", err, "backoff", backoff) select { case <-ctx.Done(): - break + break reconnect case <-time.After(backoff): } backoff = min(backoff*2, dialBackoffMax) diff --git a/spindle/mill/integration_test.go b/spindle/mill/integration_test.go index 375fc506c..50d36fa51 100644 --- a/spindle/mill/integration_test.go +++ b/spindle/mill/integration_test.go @@ -374,19 +374,6 @@ func waitForStatus(t *testing.T, d *db.DB, wid models.WorkflowId, want string) b return false } -func waitForLogFileContent(t *testing.T, path, want string) bool { - t.Helper() - deadline := time.Now().Add(5 * time.Second) - for time.Now().Before(deadline) { - data, err := os.ReadFile(path) - if err == nil && strings.Contains(string(data), want) { - return true - } - time.Sleep(20 * time.Millisecond) - } - return false -} - func waitForFileRemoval(t *testing.T, path string) bool { t.Helper() deadline := time.Now().Add(5 * time.Second) diff --git a/spindle/mill/mill.go b/spindle/mill/mill.go index ea5d36b3b..9d5b3e311 100644 --- a/spindle/mill/mill.go +++ b/spindle/mill/mill.go @@ -670,8 +670,8 @@ func (m *Mill) bid(ctx context.Context, engineName string, wid models.WorkflowId TargetEngine: engineName, RawPipelineJson: rawPipeline, RawWorkflowJson: rawWorkflow, - Knot: wid.Knot, - Rkey: wid.Rkey, + Knot: wid.PipelineId.Knot, + Rkey: wid.PipelineId.Rkey, TtlSeconds: uint32(m.cfg.ReconnectGrace / time.Second), Traceparent: traceparent, Tracestate: tracestate, diff --git a/spindle/mill/restore.go b/spindle/mill/restore.go index d29892cd6..5a4f12740 100644 --- a/spindle/mill/restore.go +++ b/spindle/mill/restore.go @@ -30,8 +30,8 @@ func (m *Mill) persistLease(lease *RemoteLease, state string) error { NodeID: lease.nodeID, Epoch: lease.epoch, Engine: lease.engine, - Knot: lease.wid.Knot, - Rkey: lease.wid.Rkey, + Knot: lease.wid.PipelineId.Knot, + Rkey: lease.wid.PipelineId.Rkey, Workflow: lease.wid.Name, State: state, QuotaReservationID: quotaID, diff --git a/spindle/models/models.go b/spindle/models/models.go index 586d557db..77bbbb025 100644 --- a/spindle/models/models.go +++ b/spindle/models/models.go @@ -30,7 +30,7 @@ type WorkflowId struct { } func (wid WorkflowId) String() string { - return fmt.Sprintf("%s-%s-%s", normalize(wid.Knot), wid.Rkey, normalize(wid.Name)) + return fmt.Sprintf("%s-%s-%s", normalize(wid.PipelineId.Knot), wid.PipelineId.Rkey, normalize(wid.Name)) } func normalize(name string) string { diff --git a/spindle/xrpc/ci_pipeline_subscribe_logs.go b/spindle/xrpc/ci_pipeline_subscribe_logs.go index 5adf44353..4a6d763fa 100644 --- a/spindle/xrpc/ci_pipeline_subscribe_logs.go +++ b/spindle/xrpc/ci_pipeline_subscribe_logs.go @@ -303,8 +303,6 @@ func (x *Xrpc) handleSubscribeLogs(w http.ResponseWriter, r *http.Request, pipel } } -func strptr(s string) *string { return &s } - func strptrOrNil(s string) *string { if s == "" { return nil diff --git a/spindle/xrpc/validation_test.go b/spindle/xrpc/validation_test.go new file mode 100644 index 000000000..1450b9a8b --- /dev/null +++ b/spindle/xrpc/validation_test.go @@ -0,0 +1,23 @@ +package xrpc + +import ( + "strings" + "testing" +) + +func TestRequireSha(t *testing.T) { + if err := requireSha(strings.Repeat("a", 40)); err != nil { + t.Fatalf("valid SHA rejected: %v", err) + } + for name, sha := range map[string]string{ + "short": strings.Repeat("a", 39), + "long": strings.Repeat("a", 41), + "non-hex": strings.Repeat("z", 40), + } { + t.Run(name, func(t *testing.T) { + if err := requireSha(sha); err == nil { + t.Fatalf("invalid SHA %q accepted", sha) + } + }) + } +} diff --git a/spindle/xrpc/xrpc.go b/spindle/xrpc/xrpc.go index c88566843..b39de81c8 100644 --- a/spindle/xrpc/xrpc.go +++ b/spindle/xrpc/xrpc.go @@ -3,6 +3,7 @@ package xrpc import ( "context" _ "embed" + "encoding/hex" "encoding/json" "errors" "fmt" @@ -35,6 +36,9 @@ func requireSha(sha string) error { if len(sha) != 40 { return fmt.Errorf("sha must be a 40-character commit hash") } + if _, err := hex.DecodeString(sha); err != nil { + return fmt.Errorf("sha must be hexadecimal: %w", err) + } return nil }