Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
19 changes: 16 additions & 3 deletions arkruntime/environment_work.go
Original file line number Diff line number Diff line change
Expand Up @@ -138,6 +138,10 @@ func (c *Client) StopWork(
if body.WorkID == "" {
return errors.New("missing required work_id")
}
force, forceSet := body.Force.Get()
if _, reasonSet := body.Reason.Get(); reasonSet && (!forceSet || !force) {
return errors.New("reason requires force=true")
}
u := c.fullURL(fmt.Sprintf("%s/%s/work/%s/stop",
environmentsPrefix,
environment.PathEscape(body.EnvironmentID),
Expand All @@ -148,15 +152,24 @@ func (c *Client) StopWork(
}

type stopWorkRequestBody struct {
Force *bool `json:"force,omitempty"`
Force *bool `json:"force,omitempty"`
Reason *environment.WorkStopReason `json:"reason,omitempty"`
}

func stopWorkBody(body *environment.StopWorkRequest) any {
request := stopWorkRequestBody{}
force, ok := body.Force.Get()
if !ok {
if ok {
request.Force = &force
}
reason, ok := body.Reason.Get()
if ok {
request.Reason = &reason
}
if request.Force == nil && request.Reason == nil {
return nil
}
return stopWorkRequestBody{Force: &force}
return request
}

func (c *Client) doControlPlaneRequest(
Expand Down
89 changes: 88 additions & 1 deletion arkruntime/lib/environments/integration_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -5,18 +5,45 @@ package environments
import (
"context"
"encoding/json"
"errors"
"io"
"net/http"
"strings"
"sync"
"testing"
"time"

"github.com/volcengine/ark-runtime-go/arkruntime/model/environment"
selfhosted "github.com/volcengine/ark-runtime-go/arkruntime/selfhosted"
"github.com/volcengine/ark-runtime-go/arkruntime/tools/agenttoolset"
"github.com/volcengine/ark-runtime-go/arkruntime/toolset"
)

func TestStopReasonForExit(t *testing.T) {
tests := []struct {
name string
err error
cause heartbeatStopCause
want environment.WorkStopReason
wantValue bool
}{
{name: "completed", want: environment.WorkStopReasonCompleted, wantValue: true},
{name: "idle timeout", err: selfhosted.ErrIdleTimeout, want: environment.WorkStopReasonCompleted, wantValue: true},
{name: "session terminated", err: selfhosted.ErrSessionTerminated, want: environment.WorkStopReasonCompleted, wantValue: true},
{name: "external cancellation", err: context.Canceled, want: environment.WorkStopReasonWorkerAbnormal, wantValue: true},
{name: "worker abnormal", err: errors.New("worker failed"), want: environment.WorkStopReasonWorkerAbnormal, wantValue: true},
{name: "platform stop", err: context.Canceled, cause: heartbeatStopCauseStopRequested},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got, ok := stopReasonForExit(tt.err, tt.cause).Get()
if ok != tt.wantValue || got != tt.want {
t.Fatalf("stop reason = %q, set=%v", got, ok)
}
})
}
}

type fakeEnvironmentWorkerAPI struct {
mu sync.Mutex
events []selfhosted.Event
Expand All @@ -29,6 +56,7 @@ type fakeEnvironmentWorkerAPI struct {

heartbeat func(context.Context, selfhosted.HeartbeatWorkRequest) (*selfhosted.HeartbeatResponse, error)
getSession func(context.Context, selfhosted.GetSessionRequest) (*selfhosted.Session, error)
listEvents func(context.Context, selfhosted.ListEventsRequest) (*selfhosted.ListEventsResponse, error)
onStop func()
}

Expand Down Expand Up @@ -86,7 +114,10 @@ func (f *fakeEnvironmentWorkerAPI) GetSession(ctx context.Context, req selfhoste
return &selfhosted.Session{ID: req.SessionID}, nil
}

func (f *fakeEnvironmentWorkerAPI) ListEvents(context.Context, selfhosted.ListEventsRequest) (*selfhosted.ListEventsResponse, error) {
func (f *fakeEnvironmentWorkerAPI) ListEvents(ctx context.Context, req selfhosted.ListEventsRequest) (*selfhosted.ListEventsResponse, error) {
if f.listEvents != nil {
return f.listEvents(ctx, req)
}
f.mu.Lock()
defer f.mu.Unlock()
if len(f.events) == 0 {
Expand All @@ -97,6 +128,57 @@ func (f *fakeEnvironmentWorkerAPI) ListEvents(context.Context, selfhosted.ListEv
return &selfhosted.ListEventsResponse{Events: events}, nil
}

func TestEnvironmentWorkerExternalCancellationReportsWorkerAbnormal(t *testing.T) {
listStarted := make(chan struct{})
api := &fakeEnvironmentWorkerAPI{
listEvents: func(ctx context.Context, _ selfhosted.ListEventsRequest) (*selfhosted.ListEventsResponse, error) {
close(listStarted)
<-ctx.Done()
return nil, ctx.Err()
},
}
worker := NewEnvironmentWorker(api, EnvironmentWorkerOptions{
EnvironmentID: "env_local",
WorkerID: "worker_local",
Workdir: t.TempDir(),
})
ctx, cancel := context.WithCancel(context.Background())
t.Cleanup(cancel)
done := make(chan error, 1)
go func() {
done <- worker.HandleItem(ctx, HandleItemOptions{
WorkID: "work_local",
EnvironmentID: "env_local",
SessionID: "sess_local",
})
}()

select {
case <-listStarted:
case <-time.After(time.Second):
t.Fatal("worker did not start event polling")
}
cancel()
select {
case err := <-done:
if err != nil {
t.Fatalf("HandleItem err = %v", err)
}
case <-time.After(time.Second):
t.Fatal("worker did not stop after context cancellation")
}

api.mu.Lock()
defer api.mu.Unlock()
if len(api.stops) != 1 {
t.Fatalf("stop count = %d, want 1", len(api.stops))
}
reason, ok := api.stops[0].Reason.Get()
if !ok || reason != environment.WorkStopReasonWorkerAbnormal {
t.Fatalf("stop reason = %q, set=%v, want worker_abnormal", reason, ok)
}
}

func (f *fakeEnvironmentWorkerAPI) SendEvent(_ context.Context, req selfhosted.SendEventRequest) error {
f.mu.Lock()
defer f.mu.Unlock()
Expand Down Expand Up @@ -156,6 +238,9 @@ func TestEnvironmentWorkerRunHandlesPolledWorkInProcess(t *testing.T) {
if force, ok := api.stops[0].Force.Get(); !ok || !force {
t.Fatalf("worker stop should be force=true: %+v", api.stops[0])
}
if reason, ok := api.stops[0].Reason.Get(); !ok || reason != environment.WorkStopReasonCompleted {
t.Fatalf("worker stop reason = %q, want completed", reason)
}
if len(api.sent) != 1 {
t.Fatalf("sent events = %+v", api.sent)
}
Expand Down Expand Up @@ -305,6 +390,8 @@ func TestEnvironmentWorkerStopsWorkOnSessionIdleEvent(t *testing.T) {
t.Fatalf("worker stop = %+v", got)
} else if force, ok := got.Force.Get(); !ok || !force {
t.Fatalf("worker stop should be force=true: %+v", got)
} else if reason, ok := got.Reason.Get(); !ok || reason != environment.WorkStopReasonCompleted {
t.Fatalf("worker stop reason = %q, want completed", reason)
}
if len(api.sent) != 0 {
t.Fatalf("sent events = %+v", api.sent)
Expand Down
1 change: 1 addition & 0 deletions arkruntime/lib/environments/poller.go
Original file line number Diff line number Diff line change
Expand Up @@ -224,6 +224,7 @@ func (p *WorkPoller) discardInvalidWork(item selfhosted.WorkItem, _ string) {
EnvironmentID: item.EnvironmentID,
WorkID: item.ID,
Force: environment.NewOptBool(true),
Reason: environment.NewOptWorkStopReason(environment.WorkStopReasonOthers),
}); err != nil && !isResolvedStatus(err) {
p.logger.Warn("stop invalid work failed", "work_id", item.ID, "err", err)
}
Expand Down
22 changes: 22 additions & 0 deletions arkruntime/lib/environments/poller_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -207,6 +207,28 @@ func TestWorkPollerStopsPollingOnPermanentAckFailure(t *testing.T) {
}
}

func TestWorkPollerMarksInvalidWorkAsOther(t *testing.T) {
api := &fakePollerAPI{
pollItem: newTestWorkItem(testWorkID, "env_1", ""),
}
poller := NewWorkPoller(context.Background(), api, WorkPollerOptions{
EnvironmentID: "env_1",
WorkerID: "worker_1",
Drain: true,
})

if poller.Next() {
t.Fatal("Next should discard work without a session id")
}
if api.stopCount != 1 {
t.Fatalf("stop_count=%d", api.stopCount)
}
reason, ok := api.stops[0].Reason.Get()
if !ok || reason != environment.WorkStopReasonOthers {
t.Fatalf("stop reason = %q, set=%v", reason, ok)
}
}

func TestWorkPollerTreatsEmptyWorkIDAsEmptyPoll(t *testing.T) {
api := &fakePollerAPI{
pollItem: newTestWorkItem("", "env_1", testSessionID),
Expand Down
23 changes: 21 additions & 2 deletions arkruntime/lib/environments/worker.go
Original file line number Diff line number Diff line change
Expand Up @@ -172,7 +172,11 @@ func (w *EnvironmentWorker) handleItem(ctx context.Context, work claimedWork) (e
<-heartbeatDone
cause := loadHeartbeatStopCause(&heartbeatCause)
if shouldStopItem(cause) {
_ = w.stopItem(api, work)
exitErr := err
if exitErr == nil {
exitErr = ctx.Err()
}
_ = w.stopItem(api, work, stopReasonForExit(exitErr, cause))
} else {
logger.Info("skip stop work after heartbeat ownership became uncertain", "cause", cause)
}
Expand Down Expand Up @@ -301,13 +305,18 @@ func (w *EnvironmentWorker) workdir() (string, error) {
return filepath.Abs(root)
}

func (w *EnvironmentWorker) stopItem(api selfhosted.API, work claimedWork) error {
func (w *EnvironmentWorker) stopItem(
api selfhosted.API,
work claimedWork,
reason environment.OptWorkStopReason,
) error {
stopCtx, stopCancel := context.WithTimeout(context.Background(), stopTimeout)
defer stopCancel()
req := selfhosted.StopWorkRequest{
EnvironmentID: work.EnvironmentID,
WorkID: work.ID,
Force: environment.NewOptBool(true),
Reason: reason,
}
if err := api.StopWork(stopCtx, req); err != nil {
if selfhosted.IsStatus(err, 409) || selfhosted.IsStatus(err, 412) {
Expand All @@ -320,6 +329,16 @@ func (w *EnvironmentWorker) stopItem(api selfhosted.API, work claimedWork) error
return nil
}

func stopReasonForExit(err error, cause heartbeatStopCause) environment.OptWorkStopReason {
if cause == heartbeatStopCauseStopRequested {
return environment.OptWorkStopReason{}
}
if err == nil || errors.Is(err, selfhosted.ErrIdleTimeout) || errors.Is(err, selfhosted.ErrSessionTerminated) {
return environment.NewOptWorkStopReason(environment.WorkStopReasonCompleted)
}
return environment.NewOptWorkStopReason(environment.WorkStopReasonWorkerAbnormal)
}

func (w *EnvironmentWorker) logger() *selfhostedlog.Logger {
return selfhostedlog.New(w.opts.Logger)
}
Expand Down
Loading
Loading