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
11 changes: 11 additions & 0 deletions arkruntime/client.go
Original file line number Diff line number Diff line change
Expand Up @@ -41,6 +41,7 @@ type Client struct {
mandatoryRefreshTimeout int

batchHTTPClient *http.Client
sessionStreamClient *http.Client
modelBreakerProvider *utils.ModelBreakerProvider
}

Expand All @@ -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.
Expand Down Expand Up @@ -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(),
}
}
Expand Down
116 changes: 112 additions & 4 deletions arkruntime/lib/environments/integration_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@ package environments

import (
"context"
"encoding/json"
"io"
"net/http"
"strings"
Expand All @@ -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 {
Expand All @@ -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()
Expand Down Expand Up @@ -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)
}
Expand All @@ -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
Expand Down
22 changes: 15 additions & 7 deletions arkruntime/lib/environments/poller.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -49,6 +53,7 @@ type WorkPoller struct {
failures int
discards int
closed bool
autoStop bool
}

// NewWorkPoller 创建 work poller。
Expand All @@ -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")
Expand Down Expand Up @@ -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
Expand All @@ -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
Expand Down
24 changes: 24 additions & 0 deletions arkruntime/lib/environments/poller_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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),
Expand Down
17 changes: 17 additions & 0 deletions arkruntime/lib/environments/worker.go
Original file line number Diff line number Diff line change
Expand Up @@ -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。
Expand Down Expand Up @@ -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() }()
Expand Down Expand Up @@ -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() }()
Expand Down Expand Up @@ -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 {
Expand All @@ -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 == "" {
Expand Down
54 changes: 54 additions & 0 deletions arkruntime/self_hosted_client_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down Expand Up @@ -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" {
Expand Down
Loading
Loading