diff --git a/arkruntime/client.go b/arkruntime/client.go index 9c8fb83..59bd145 100644 --- a/arkruntime/client.go +++ b/arkruntime/client.go @@ -41,6 +41,7 @@ type Client struct { mandatoryRefreshTimeout int batchHTTPClient *http.Client + sessionStreamClient *http.Client modelBreakerProvider *utils.ModelBreakerProvider } @@ -62,6 +63,15 @@ func (c *Client) HTTPClient() *http.Client { return c.config.HTTPClient } +func newSessionStreamClient(client *http.Client) *http.Client { + if client == nil { + client = http.DefaultClient + } + streamClient := *client + streamClient.Timeout = 0 + return &streamClient +} + // NewVolcClient constructs a client targeting the Volcengine cloud // (ark.cn-beijing.volces.com). Reads ARK_API_KEY for api-key auth and // VOLC_ACCESSKEY/VOLC_SECRETKEY for AK/SK auth, in that preference order. @@ -131,6 +141,7 @@ func newClientWithConfig(config clientConfig) *Client { advisoryRefreshTimeout: model.DefaultAdvisoryRefreshTimeout, mandatoryRefreshTimeout: model.DefaultMandatoryRefreshTimeout, batchHTTPClient: newBatchHTTPClient(config.batchMaxParallel), + sessionStreamClient: newSessionStreamClient(config.HTTPClient), modelBreakerProvider: utils.NewModelBreakerProvider(), } } diff --git a/arkruntime/lib/environments/integration_test.go b/arkruntime/lib/environments/integration_test.go index cb68359..9fd8777 100644 --- a/arkruntime/lib/environments/integration_test.go +++ b/arkruntime/lib/environments/integration_test.go @@ -4,6 +4,7 @@ package environments import ( "context" + "encoding/json" "io" "net/http" "strings" @@ -12,6 +13,8 @@ import ( "time" 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" ) type fakeEnvironmentWorkerAPI struct { @@ -29,6 +32,15 @@ type fakeEnvironmentWorkerAPI struct { onStop func() } +type blockingWorkerTool struct{} + +func (*blockingWorkerTool) Name() string { return "blocking" } + +func (*blockingWorkerTool) Execute(ctx context.Context, _ json.RawMessage) toolset.Result { + <-ctx.Done() + return toolset.ErrorResult(ctx.Err().Error()) +} + func (f *fakeEnvironmentWorkerAPI) PollWork(context.Context, selfhosted.PollWorkRequest) (*selfhosted.WorkItem, error) { f.mu.Lock() defer f.mu.Unlock() @@ -131,15 +143,12 @@ func TestEnvironmentWorkerRunHandlesPolledWorkInProcess(t *testing.T) { api.mu.Lock() defer api.mu.Unlock() - if api.ackCount != 1 || api.stopCount != 2 { + if api.ackCount != 1 || api.stopCount != 1 { t.Fatalf("ack_count=%d stop_count=%d", api.ackCount, api.stopCount) } if force, ok := api.stops[0].Force.Get(); !ok || !force { t.Fatalf("worker stop should be force=true: %+v", api.stops[0]) } - if _, ok := api.stops[1].Force.Get(); ok { - t.Fatalf("poller release stop should be force=false: %+v", api.stops[1]) - } if len(api.sent) != 1 { t.Fatalf("sent events = %+v", api.sent) } @@ -155,6 +164,105 @@ func TestEnvironmentWorkerRunHandlesPolledWorkInProcess(t *testing.T) { } } +func TestEnvironmentWorkerForwardsToolTimeout(t *testing.T) { + maxIdle := 10 * time.Millisecond + toolTimeout := 20 * time.Millisecond + tool := &blockingWorkerTool{} + api := &fakeEnvironmentWorkerAPI{ + events: []selfhosted.Event{ + { + ID: "toolu_timeout", + Type: selfhosted.EventTypeAgentCustomToolUse, + Name: tool.Name(), + ToolUseID: "call_timeout", + Input: selfhosted.RawJSON(`{}`), + SessionThreadID: "thread_timeout", + }, + { + ID: "evt_idle", + Type: selfhosted.EventTypeSessionStatusIdle, + StopReason: &selfhosted.SessionStopReason{Type: selfhosted.SessionStopReasonEndTurn}, + }, + }, + } + worker := NewEnvironmentWorker(api, EnvironmentWorkerOptions{ + EnvironmentID: "env_timeout", + Workdir: t.TempDir(), + MaxIdle: &maxIdle, + ToolTimeout: toolTimeout, + CustomTools: map[string]toolset.Tool{tool.Name(): tool}, + }) + if err := worker.HandleItem(context.Background(), HandleItemOptions{ + WorkID: "work_timeout", + EnvironmentID: "env_timeout", + SessionID: "sess_timeout", + }); err != nil { + t.Fatal(err) + } + + api.mu.Lock() + defer api.mu.Unlock() + if len(api.sent) != 1 { + t.Fatalf("sent events = %+v", api.sent) + } + result := api.sent[0] + if result.IsError == nil || !*result.IsError { + t.Fatalf("result is_error = %v", result.IsError) + } + if len(result.Content) != 1 || !strings.Contains(result.Content[0].Text, "tool execution timed out after 20ms") { + t.Fatalf("result content = %+v", result.Content) + } +} + +func TestEnvironmentWorkerForwardsToolTimeoutToDefaultTools(t *testing.T) { + maxIdle := 10 * time.Millisecond + api := &fakeEnvironmentWorkerAPI{ + events: []selfhosted.Event{ + { + ID: "toolu_default_timeout", + Type: selfhosted.EventTypeAgentToolUse, + Name: "bash", + Input: selfhosted.RawJSON(`{"command":"sleep 0.1; printf timeout-forwarded"}`), + EvaluatedPermission: selfhosted.PermissionAllow, + }, + { + ID: "evt_idle", + Type: selfhosted.EventTypeSessionStatusIdle, + StopReason: &selfhosted.SessionStopReason{Type: selfhosted.SessionStopReasonEndTurn}, + }, + }, + } + worker := NewEnvironmentWorker(api, EnvironmentWorkerOptions{ + EnvironmentID: "env_default_timeout", + Workdir: t.TempDir(), + MaxIdle: &maxIdle, + ToolTimeout: 2 * time.Second, + ToolContext: &agenttoolset.AgentToolContext{ + ToolTimeout: 20 * time.Millisecond, + }, + }) + if err := worker.HandleItem(context.Background(), HandleItemOptions{ + WorkID: "work_default_timeout", + EnvironmentID: "env_default_timeout", + SessionID: "sess_default_timeout", + }); err != nil { + t.Fatal(err) + } + + api.mu.Lock() + defer api.mu.Unlock() + if len(api.sent) != 1 { + t.Fatalf("sent events = %+v", api.sent) + } + result := api.sent[0] + if result.IsError == nil || *result.IsError { + t.Fatalf("result is_error = %v", result.IsError) + } + if len(result.Content) != 1 || !strings.Contains(result.Content[0].Text, "timeout-forwarded") { + t.Fatalf("result content = %+v", result.Content) + } +} + func TestEnvironmentWorkerStopsWorkOnSessionIdleEvent(t *testing.T) { root := t.TempDir() maxIdle := 10 * time.Millisecond diff --git a/arkruntime/lib/environments/poller.go b/arkruntime/lib/environments/poller.go index 74a003b..d37ab4b 100644 --- a/arkruntime/lib/environments/poller.go +++ b/arkruntime/lib/environments/poller.go @@ -34,10 +34,14 @@ type WorkPollerOptions struct { BlockMS int ReclaimOlderThanMS int Drain bool - Logger *log.Logger + // AutoStop 控制 Poller 是否在下一次 Next 或 Close 时停止上一条 work。 + // 未设置时默认开启;当调用方自行管理 work 生命周期时应显式关闭。 + AutoStop environment.OptBool + Logger *log.Logger } // WorkPoller 负责从 environment work queue 里 poll 并 ack work。 +// WorkPoller 不支持并发调用,所有方法必须由同一个 goroutine 顺序执行。 type WorkPoller struct { ctx context.Context api selfhosted.API @@ -49,6 +53,7 @@ type WorkPoller struct { failures int discards int closed bool + autoStop bool } // NewWorkPoller 创建 work poller。 @@ -58,10 +63,11 @@ func NewWorkPoller(ctx context.Context, api selfhosted.API, opts WorkPollerOptio opts.WorkerID = defaultWorkerID() } p := &WorkPoller{ - ctx: ctx, - api: api, - opts: opts, - logger: logger.With("component", "work-poller", "environment_id", opts.EnvironmentID), + ctx: ctx, + api: api, + opts: opts, + logger: logger.With("component", "work-poller", "environment_id", opts.EnvironmentID), + autoStop: opts.AutoStop.Or(true), } if opts.EnvironmentID == "" { p.err = errors.New("environments: EnvironmentID is required") @@ -148,7 +154,9 @@ func (p *WorkPoller) Next() bool { continue } p.current = item - p.pendingStop = p.makeStopClosure(*item) + if p.autoStop { + p.pendingStop = p.makeStopClosure(*item) + } p.discards = 0 p.logger.Info("claimed work", "work_id", item.ID, "session_id", item.SessionIDValue()) return true @@ -165,7 +173,7 @@ func (p *WorkPoller) Err() error { return p.err } -// Close 停止 poller,并对最后 yield 的 work 发送 StopWork。 +// Close 停止 poller;AutoStop 开启时对最后 yield 的 work 发送 StopWork。 func (p *WorkPoller) Close() error { if p.closed { return nil diff --git a/arkruntime/lib/environments/poller_test.go b/arkruntime/lib/environments/poller_test.go index 2e6627a..1f3546c 100644 --- a/arkruntime/lib/environments/poller_test.go +++ b/arkruntime/lib/environments/poller_test.go @@ -108,6 +108,30 @@ func TestWorkPollerAcksAndStopsOnClose(t *testing.T) { } } +func TestWorkPollerAutoStopDisabledDoesNotStop(t *testing.T) { + api := &fakePollerAPI{ + pollItem: newTestWorkItem(testWorkID, "env_1", testSessionID), + } + poller := NewWorkPoller(context.Background(), api, WorkPollerOptions{ + EnvironmentID: "env_1", + WorkerID: "worker_1", + Drain: true, + AutoStop: environment.NewOptBool(false), + }) + if !poller.Next() { + t.Fatalf("Next returned false: %v", poller.Err()) + } + if poller.Next() { + t.Fatal("Next should return false after the queue is drained") + } + if err := poller.Close(); err != nil { + t.Fatal(err) + } + if api.ackCount != 1 || api.stopCount != 0 { + t.Fatalf("ack_count=%d stop_count=%d", api.ackCount, api.stopCount) + } +} + func TestWorkPollerStopsPreviousBeforeNext(t *testing.T) { api := &fakePollerAPI{ pollItem: newTestWorkItem(testWorkID, "env_1", testSessionID), diff --git a/arkruntime/lib/environments/worker.go b/arkruntime/lib/environments/worker.go index e4b91ff..4070cdc 100644 --- a/arkruntime/lib/environments/worker.go +++ b/arkruntime/lib/environments/worker.go @@ -39,6 +39,8 @@ type EnvironmentWorkerOptions struct { ToolsFunc func(env *agenttoolset.AgentToolContext) (*agenttoolset.Set, error) // MaxIdle 是 session 在 end_turn idle 后继续等待事件的时间,nil 使用默认值。 MaxIdle *time.Duration + // ToolTimeout 限制单次工具执行时间,非正值使用 SessionToolRunner 默认值。 + ToolTimeout time.Duration // CustomTools 是按名称注册的自定义工具。 CustomTools map[string]agenttoolset.Tool // Logger 接收 worker 生命周期日志,nil 使用 log.Default。 @@ -89,6 +91,7 @@ func (w *EnvironmentWorker) Run(ctx context.Context) error { poller := NewWorkPoller(ctx, w.api, WorkPollerOptions{ EnvironmentID: environmentID, WorkerID: w.opts.WorkerID, + AutoStop: environment.NewOptBool(false), Logger: w.opts.Logger, }) defer func() { _ = poller.Close() }() @@ -213,6 +216,7 @@ func (w *EnvironmentWorker) handleItem(ctx context.Context, work claimedWork, us CustomTools: w.opts.CustomTools, ResultStore: store, MaxIdle: w.opts.MaxIdle, + ToolTimeout: w.effectiveToolTimeout(), Logger: w.opts.Logger, }) defer func() { _ = runner.Close() }() @@ -260,6 +264,9 @@ func (w *EnvironmentWorker) toolContext(workdir string) *agenttoolset.AgentToolC if w.opts.UnrestrictedPaths { base.UnrestrictedPaths = true } + if timeout := w.effectiveToolTimeout(); timeout > 0 { + base.ToolTimeout = timeout + } if base.Env != nil { env := make(map[string]string, len(base.Env)) for k, v := range base.Env { @@ -270,6 +277,16 @@ func (w *EnvironmentWorker) toolContext(workdir string) *agenttoolset.AgentToolC return &base } +func (w *EnvironmentWorker) effectiveToolTimeout() time.Duration { + if w.opts.ToolTimeout > 0 { + return w.opts.ToolTimeout + } + if w.opts.ToolContext != nil && w.opts.ToolContext.ToolTimeout > 0 { + return w.opts.ToolContext.ToolTimeout + } + return 0 +} + func (w *EnvironmentWorker) workdirFor(sessionID string, useWorkdirAsSession bool) (string, error) { root := w.opts.Workdir if root == "" { diff --git a/arkruntime/self_hosted_client_test.go b/arkruntime/self_hosted_client_test.go index a9523b3..d3def25 100644 --- a/arkruntime/self_hosted_client_test.go +++ b/arkruntime/self_hosted_client_test.go @@ -11,6 +11,7 @@ import ( "net/http/httptest" "strings" "testing" + "time" "github.com/volcengine/ark-runtime-go/arkruntime/model/environment" "github.com/volcengine/ark-runtime-go/arkruntime/model/session" @@ -141,6 +142,59 @@ func TestEnvironmentWorkRequests(t *testing.T) { } } +func TestSessionStreamsIgnoreHTTPClientTimeout(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "text/event-stream") + w.WriteHeader(http.StatusOK) + if flusher, ok := w.(http.Flusher); ok { + flusher.Flush() + } + time.Sleep(80 * time.Millisecond) + _, _ = io.WriteString(w, "data: {\"type\":\"session.status_idle\"}\n\n") + })) + defer server.Close() + + client := NewClientWithApiKey( + "test-api-key", + WithBaseUrl(server.URL), + WithHTTPClient(&http.Client{Timeout: 20 * time.Millisecond}), + ) + tests := []struct { + name string + open func(context.Context) (*session.StreamDecoder, error) + }{ + { + name: "session", + open: func(ctx context.Context) (*session.StreamDecoder, error) { + return client.StreamSessionEvents(ctx, "sess-1") + }, + }, + { + name: "thread", + open: func(ctx context.Context) (*session.StreamDecoder, error) { + return client.StreamSessionThreadEvents(ctx, "sess-1", "thread-1") + }, + }, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + decoder, err := test.open(ctx) + if err != nil { + t.Fatal(err) + } + defer decoder.Close() //nolint:errcheck // test cleanup + if !decoder.Next() { + t.Fatalf("stream ended before delayed event: %v", decoder.Err()) + } + if decoder.Event().Type != "session.status_idle" { + t.Fatalf("event type = %q", decoder.Event().Type) + } + }) + } +} + func TestEnvironmentPollEmptyQueueReturnsNil(t *testing.T) { server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if r.Method != http.MethodGet || r.URL.Path != "/environments/env-empty/work/poll" { diff --git a/arkruntime/sessions.go b/arkruntime/sessions.go index 0daedb4..f6cba20 100644 --- a/arkruntime/sessions.go +++ b/arkruntime/sessions.go @@ -225,7 +225,7 @@ func (c *Client) StreamSessionEventsWithParams( if reqErr != nil { return nil, reqErr } - resp, err := c.config.HTTPClient.Do(req) + resp, err := c.sessionStreamClient.Do(req) if err != nil { return nil, err } @@ -403,7 +403,7 @@ func (c *Client) StreamSessionThreadEventsWithParams( if reqErr != nil { return nil, reqErr } - resp, err := c.config.HTTPClient.Do(req) + resp, err := c.sessionStreamClient.Do(req) if err != nil { return nil, err }