From 57ef3fb36e5a820c0da2d81313a3c6c6dd8ca5b1 Mon Sep 17 00:00:00 2001 From: "ark-hand[bot]" <315378070+ark-hand[bot]@users.noreply.github.com> Date: Wed, 26 Aug 2026 07:02:23 +0000 Subject: [PATCH] feat(selfHosted): add self-hosted worker runtime MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ## 简述 为 Ark Runtime Go SDK 增加 Managed Agents Self-Hosted worker 能力,使客户可以在自有运行环境中 poll work、维护 lease、执行 session tool call 并安装 skill。 ## 改动 - 增加 self-hosted work API、WorkPoller、EnvironmentWorker 和 HandleItem。 - 增加 SessionToolRunner、bash/file/search 工具和 skill 初始化。 - 对齐 MA poll/ack/heartbeat/stop 与 event 契约。 - 加固 ACK 所有权、heartbeat lease、idle 退出、SSE 恢复、进程取消、路径隔离和压缩包安全。 - 增加 THIRD_PARTY_NOTICES.md,保留所参考 Anthropic MIT 实现的归属声明,项目主许可证仍为 Apache-2.0。 ## 验证 - test/run.sh --all 通过。 - test/run.sh --contract 通过。 - test/run.sh --docker 通过,覆盖宿主 poll/ack + Docker HandleItem、kill/SIGTERM、idle、stop_requested、skill 和凭证隔离。 - Go/Python/Java STG 命令执行与 SkillHub/私有 skill E2E 通过。 See merge request: !84 Sync-Source-Commit: 74391730e0fdc8a21a6fc0519c5fe372d3d6fbb1 Ark-APIs-Commit: 644ae285bd16f965668c33400bf6dcce54261f62 Hand-Written-Reason: Manual self-hosted worker implementation; no ark-apis regeneration involved. Release-Version: 0.4.0 --- README.md | 5 + THIRD_PARTY_NOTICES.md | 39 + arkruntime/client.go | 8 + arkruntime/config_test.go | 33 + arkruntime/environment_work.go | 212 + arkruntime/internal/selfhostedlog/logger.go | 86 + arkruntime/lib/environments/heartbeat.go | 133 + .../lib/environments/integration_test.go | 440 ++ arkruntime/lib/environments/poller.go | 290 ++ arkruntime/lib/environments/poller_test.go | 225 ++ arkruntime/lib/environments/workdir.go | 52 + arkruntime/lib/environments/worker.go | 383 ++ arkruntime/lib/environments/worker_test.go | 50 + arkruntime/model/agent/oas_json_gen.go | 835 +++- arkruntime/model/agent/oas_parameters_gen.go | 4 + arkruntime/model/agent/oas_schemas_gen.go | 459 ++- arkruntime/model/agent/oas_validators_gen.go | 93 + arkruntime/model/environment/oas_json_gen.go | 1117 +++++- .../model/environment/oas_parameters_gen.go | 35 + .../model/environment/oas_schemas_gen.go | 639 +++ .../model/environment/oas_validators_gen.go | 103 + arkruntime/model/environment/work_shim.go | 108 + arkruntime/model/session/oas_json_gen.go | 3566 +++++++++++++---- arkruntime/model/session/oas_schemas_gen.go | 1155 +++++- .../model/session/oas_validators_gen.go | 165 +- .../model/session/session_stream_shim.go | 21 +- .../model/session/session_thread_shim.go | 16 + arkruntime/model/skill/oas_json_gen.go | 141 +- arkruntime/model/skill/oas_parameters_gen.go | 6 + arkruntime/model/skill/oas_schemas_gen.go | 129 +- arkruntime/model/skill/skill_shim.go | 20 +- arkruntime/self_hosted_client_test.go | 412 ++ arkruntime/selfhosted/client_api.go | 481 +++ arkruntime/selfhosted/client_api_test.go | 234 ++ arkruntime/selfhosted/defaults.go | 12 + arkruntime/selfhosted/doc.go | 4 + arkruntime/selfhosted/envinit/initializer.go | 394 ++ .../selfhosted/envinit/initializer_test.go | 315 ++ arkruntime/selfhosted/errors.go | 58 + arkruntime/selfhosted/session_event.go | 43 + arkruntime/selfhosted/session_tool_runner.go | 1110 +++++ .../selfhosted/session_tool_runner_test.go | 259 ++ arkruntime/selfhosted/skill_hub.go | 130 + arkruntime/selfhosted/tool_result_store.go | 223 ++ .../selfhosted/tool_result_store_test.go | 59 + arkruntime/selfhosted/types.go | 614 +++ arkruntime/session_events_raw.go | 49 + arkruntime/sessions.go | 36 + arkruntime/skill_content.go | 75 + arkruntime/skills.go | 26 +- arkruntime/tools/agenttoolset/agenttoolset.go | 91 + arkruntime/toolset/bash.go | 475 +++ arkruntime/toolset/file.go | 268 ++ arkruntime/toolset/path.go | 145 + arkruntime/toolset/process_other.go | 15 + arkruntime/toolset/process_unix.go | 21 + arkruntime/toolset/search.go | 339 ++ arkruntime/toolset/toolset_test.go | 410 ++ arkruntime/toolset/types.go | 167 + examples/README.md | 4 +- examples/byteplus/sessions_loop/main.go | 2 +- examples/self_hosted_worker/main.go | 52 + examples/volc/sessions_loop/main.go | 2 +- 63 files changed, 16075 insertions(+), 1018 deletions(-) create mode 100644 THIRD_PARTY_NOTICES.md create mode 100644 arkruntime/config_test.go create mode 100644 arkruntime/environment_work.go create mode 100644 arkruntime/internal/selfhostedlog/logger.go create mode 100644 arkruntime/lib/environments/heartbeat.go create mode 100644 arkruntime/lib/environments/integration_test.go create mode 100644 arkruntime/lib/environments/poller.go create mode 100644 arkruntime/lib/environments/poller_test.go create mode 100644 arkruntime/lib/environments/workdir.go create mode 100644 arkruntime/lib/environments/worker.go create mode 100644 arkruntime/lib/environments/worker_test.go create mode 100644 arkruntime/model/environment/work_shim.go create mode 100644 arkruntime/self_hosted_client_test.go create mode 100644 arkruntime/selfhosted/client_api.go create mode 100644 arkruntime/selfhosted/client_api_test.go create mode 100644 arkruntime/selfhosted/defaults.go create mode 100644 arkruntime/selfhosted/doc.go create mode 100644 arkruntime/selfhosted/envinit/initializer.go create mode 100644 arkruntime/selfhosted/envinit/initializer_test.go create mode 100644 arkruntime/selfhosted/errors.go create mode 100644 arkruntime/selfhosted/session_event.go create mode 100644 arkruntime/selfhosted/session_tool_runner.go create mode 100644 arkruntime/selfhosted/session_tool_runner_test.go create mode 100644 arkruntime/selfhosted/skill_hub.go create mode 100644 arkruntime/selfhosted/tool_result_store.go create mode 100644 arkruntime/selfhosted/tool_result_store_test.go create mode 100644 arkruntime/selfhosted/types.go create mode 100644 arkruntime/session_events_raw.go create mode 100644 arkruntime/skill_content.go create mode 100644 arkruntime/tools/agenttoolset/agenttoolset.go create mode 100644 arkruntime/toolset/bash.go create mode 100644 arkruntime/toolset/file.go create mode 100644 arkruntime/toolset/path.go create mode 100644 arkruntime/toolset/process_other.go create mode 100644 arkruntime/toolset/process_unix.go create mode 100644 arkruntime/toolset/search.go create mode 100644 arkruntime/toolset/toolset_test.go create mode 100644 arkruntime/toolset/types.go create mode 100644 examples/self_hosted_worker/main.go diff --git a/README.md b/README.md index 2df984f..51cc0b7 100644 --- a/README.md +++ b/README.md @@ -294,3 +294,8 @@ ARK_API_KEY=your-key go run examples/volc/responses/basic/main.go - Go 1.20 or later - A Volcengine or BytePlus ModelArk API key + +## Third-party notices + +See [THIRD_PARTY_NOTICES.md](./THIRD_PARTY_NOTICES.md) for third-party +attribution notices. diff --git a/THIRD_PARTY_NOTICES.md b/THIRD_PARTY_NOTICES.md new file mode 100644 index 0000000..8cfd333 --- /dev/null +++ b/THIRD_PARTY_NOTICES.md @@ -0,0 +1,39 @@ +# Third-Party Notices + +This repository contains code that is derived from or structurally adapted +from third-party open-source projects. + +## Anthropic self-hosted worker SDK + +Portions of the self-hosted worker lifecycle and local agent tool +implementations under +`arkruntime/selfhosted`, `arkruntime/lib/environments`, `arkruntime/toolset`, +and `arkruntime/tools/agenttoolset` are structurally adapted from Anthropic's +self-hosted worker SDK implementation: + +https://github.com/anthropics/anthropic-sdk-go + +The upstream project is licensed under the MIT License. The MIT copyright and +permission notice is preserved below as required by that license. + +```text +Copyright 2023 Anthropic, PBC. + +Permission is hereby granted, free of charge, to any person obtaining a copy +of this software and associated documentation files (the "Software"), to deal +in the Software without restriction, including without limitation the rights +to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +copies of the Software, and to permit persons to whom the Software is +furnished to do so, subject to the following conditions: + +The above copyright notice and this permission notice shall be included in all +copies or substantial portions of the Software. + +THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +SOFTWARE. +``` diff --git a/arkruntime/client.go b/arkruntime/client.go index cf4ea68..9c8fb83 100644 --- a/arkruntime/client.go +++ b/arkruntime/client.go @@ -54,6 +54,14 @@ func NewClientWithAkSk(ak, sk string, setters ...ConfigOption) *Client { return newClientWithConfig(config) } +// HTTPClient returns the configured HTTP client. +func (c *Client) HTTPClient() *http.Client { + if c == nil || c.config.HTTPClient == nil { + return http.DefaultClient + } + return c.config.HTTPClient +} + // 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. diff --git a/arkruntime/config_test.go b/arkruntime/config_test.go new file mode 100644 index 0000000..8812538 --- /dev/null +++ b/arkruntime/config_test.go @@ -0,0 +1,33 @@ +// Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +// SPDX-License-Identifier: Apache-2.0 + +package arkruntime + +import "testing" + +func TestNewClientConfigBaseURL(t *testing.T) { + tests := []struct { + name string + options []ConfigOption + want string + }{ + { + name: "default", + want: "https://ark.cn-beijing.volces.com/api/v3", + }, + { + name: "override", + options: []ConfigOption{WithBaseUrl("https://example.com/api/v3/")}, + want: "https://example.com/api/v3", + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + config := NewClientConfig("test-api-key", "", "", test.options...) + if config.BaseURL != test.want { + t.Fatalf("BaseURL = %q, want %q", config.BaseURL, test.want) + } + }) + } +} diff --git a/arkruntime/environment_work.go b/arkruntime/environment_work.go new file mode 100644 index 0000000..962d341 --- /dev/null +++ b/arkruntime/environment_work.go @@ -0,0 +1,212 @@ +// Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +// SPDX-License-Identifier: Apache-2.0 + +package arkruntime + +import ( + "context" + "errors" + "fmt" + "io" + "net/http" + "net/url" + "strconv" + + "github.com/volcengine/ark-runtime-go/arkruntime/model" + "github.com/volcengine/ark-runtime-go/arkruntime/model/environment" + "github.com/volcengine/ark-runtime-go/arkruntime/utils" +) + +const ( + environmentWorkWorkerIDHeader = "Ark-Worker-ID" +) + +// PollWork polls one work item from an Environment work queue. +func (c *Client) PollWork( + ctx context.Context, + body *environment.PollWorkRequest, + setters ...requestOption, +) (*environment.WorkItem, error) { + if body == nil { + return nil, errors.New("missing required request body") + } + if body.EnvironmentID == "" { + return nil, errors.New("missing required environment_id") + } + q := url.Values{} + if body.WorkerID != "" { + setters = append(setters, WithCustomHeader(environmentWorkWorkerIDHeader, body.WorkerID)) + } + if body.BlockMS > 0 { + q.Set("block_ms", strconv.Itoa(body.BlockMS)) + } + if body.ReclaimOlderThanMS > 0 { + q.Set("reclaim_older_than_ms", strconv.Itoa(body.ReclaimOlderThanMS)) + } + u := c.fullURL(fmt.Sprintf("%s/%s/work/poll", environmentsPrefix, environment.PathEscape(body.EnvironmentID))) + if encoded := q.Encode(); encoded != "" { + u += "?" + encoded + } + + opts := append(setters, withBody(nil)) + wrap := &environment.WorkItemResponse{} + if err := c.doControlPlaneRequest(ctx, http.MethodGet, u, wrap, opts...); err != nil { + return nil, err + } + if wrap.ID == "" { + return nil, nil + } + return &wrap.WorkItem, nil +} + +// AckWork acknowledges one claimed work item. +func (c *Client) AckWork( + ctx context.Context, + body *environment.AckWorkRequest, + setters ...requestOption, +) error { + if body == nil { + return errors.New("missing required request body") + } + if body.EnvironmentID == "" { + return errors.New("missing required environment_id") + } + if body.WorkID == "" { + return errors.New("missing required work_id") + } + u := c.fullURL(fmt.Sprintf("%s/%s/work/%s/ack", + environmentsPrefix, + environment.PathEscape(body.EnvironmentID), + environment.PathEscape(body.WorkID), + )) + if workerID, ok := body.WorkerID.Get(); ok { + setters = append(setters, WithCustomHeader(environmentWorkWorkerIDHeader, workerID)) + } + wrap := &environment.WorkItemResponse{} + return c.doControlPlaneRequest(ctx, http.MethodPost, u, wrap, append(setters, withBody(nil))...) +} + +// HeartbeatWork refreshes a claimed work lease. +func (c *Client) HeartbeatWork( + ctx context.Context, + body *environment.HeartbeatWorkRequest, + setters ...requestOption, +) (*environment.HeartbeatWorkResponse, error) { + if body == nil { + return nil, errors.New("missing required request body") + } + if body.EnvironmentID == "" { + return nil, errors.New("missing required environment_id") + } + if body.WorkID == "" { + return nil, errors.New("missing required work_id") + } + q := url.Values{} + if desiredTTLSeconds, ok := body.DesiredTTLSeconds.Get(); ok && desiredTTLSeconds > 0 { + q.Set("desired_ttl_seconds", strconv.FormatInt(desiredTTLSeconds, 10)) + } + if expectedLastHeartbeat, ok := body.ExpectedLastHeartbeat.Get(); ok && expectedLastHeartbeat != "" { + q.Set("expected_last_heartbeat", expectedLastHeartbeat) + } + u := c.fullURL(fmt.Sprintf("%s/%s/work/%s/heartbeat", + environmentsPrefix, + environment.PathEscape(body.EnvironmentID), + environment.PathEscape(body.WorkID), + )) + if encoded := q.Encode(); encoded != "" { + u += "?" + encoded + } + wrap := &environment.HeartbeatWorkResponseWrapper{} + if err := c.doControlPlaneRequest(ctx, http.MethodPost, u, wrap, append(setters, withBody(nil))...); err != nil { + return nil, err + } + return &wrap.HeartbeatWorkResponse, nil +} + +// StopWork releases or stops one claimed work item. +func (c *Client) StopWork( + ctx context.Context, + body *environment.StopWorkRequest, + setters ...requestOption, +) error { + if body == nil { + return errors.New("missing required request body") + } + if body.EnvironmentID == "" { + return errors.New("missing required environment_id") + } + if body.WorkID == "" { + return errors.New("missing required work_id") + } + u := c.fullURL(fmt.Sprintf("%s/%s/work/%s/stop", + environmentsPrefix, + environment.PathEscape(body.EnvironmentID), + environment.PathEscape(body.WorkID), + )) + wrap := &environment.WorkItemResponse{} + return c.doControlPlaneRequest(ctx, http.MethodPost, u, wrap, append(setters, withBody(stopWorkBody(body)))...) +} + +type stopWorkRequestBody struct { + Force *bool `json:"force,omitempty"` +} + +func stopWorkBody(body *environment.StopWorkRequest) any { + force, ok := body.Force.Get() + if !ok { + return nil + } + return stopWorkRequestBody{Force: &force} +} + +func (c *Client) doControlPlaneRequest( + ctx context.Context, + method string, + u string, + v model.Response, + setters ...requestOption, +) error { + return utils.Retry( + ctx, + utils.RetryPolicy{ + MaxAttempts: c.config.RetryTimes, + InitialBackoff: model.ErrorRetryBaseDelay, + MaxBackoff: model.ErrorRetryMaxDelay, + }, + func() bool { return true }, + func() error { + req, reqErr := c.newRequest(ctx, method, u, "", "", setters...) + if reqErr != nil { + return reqErr + } + return c.sendControlPlaneRequest(req, v) + }, + nil, + needRetryError, + ) +} + +func (c *Client) sendControlPlaneRequest(req *http.Request, v model.Response) error { + requestID := req.Header.Get(model.ClientRequestHeader) + req.Header.Set("Accept", "application/json") + if req.Header.Get("Content-Type") == "" { + req.Header.Set("Content-Type", "application/json") + } + + res, err := c.config.HTTPClient.Do(req) + if err != nil { + return model.NewRequestError(http.StatusInternalServerError, err, requestID) + } + defer res.Body.Close() //nolint:errcheck // response body close errors are non-actionable + + if v != nil { + v.SetHeader(res.Header) + } + if isFailureStatusCode(res) { + return c.handleErrorResp(res) + } + if err := decodeResponse(res.Body, v); err != nil && !errors.Is(err, io.EOF) { + return model.NewRequestError(res.StatusCode, err, requestID) + } + return nil +} diff --git a/arkruntime/internal/selfhostedlog/logger.go b/arkruntime/internal/selfhostedlog/logger.go new file mode 100644 index 0000000..0af2ecb --- /dev/null +++ b/arkruntime/internal/selfhostedlog/logger.go @@ -0,0 +1,86 @@ +// Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +// SPDX-License-Identifier: Apache-2.0 + +// Package selfhostedlog 为 self-hosted worker 提供兼容 Go 1.20 的结构化日志适配。 +package selfhostedlog + +import ( + "fmt" + "log" + "strconv" + "strings" +) + +// Logger 在标准库 log.Logger 上保留 self-hosted worker 使用的键值日志接口。 +type Logger struct { + base *log.Logger + attrs []any +} + +// New 创建日志适配器,base 为空时使用 log.Default。 +func New(base *log.Logger) *Logger { + if base == nil { + base = log.Default() + } + return &Logger{base: base} +} + +// With 返回附带固定字段的新日志适配器。 +func (l *Logger) With(args ...any) *Logger { + if l == nil { + l = New(nil) + } + attrs := make([]any, 0, len(l.attrs)+len(args)) + attrs = append(attrs, l.attrs...) + attrs = append(attrs, args...) + return &Logger{base: l.base, attrs: attrs} +} + +// Debug 记录调试日志。 +func (l *Logger) Debug(message string, args ...any) { l.output("DEBUG", message, args...) } + +// Info 记录信息日志。 +func (l *Logger) Info(message string, args ...any) { l.output("INFO", message, args...) } + +// Warn 记录警告日志。 +func (l *Logger) Warn(message string, args ...any) { l.output("WARN", message, args...) } + +// Error 记录错误日志。 +func (l *Logger) Error(message string, args ...any) { l.output("ERROR", message, args...) } + +func (l *Logger) output(level, message string, args ...any) { + if l == nil { + l = New(nil) + } + all := make([]any, 0, len(l.attrs)+len(args)) + all = append(all, l.attrs...) + all = append(all, args...) + l.base.Printf("level=%s msg=%s%s", level, quoteValue(message), formatAttrs(all)) +} + +func formatAttrs(attrs []any) string { + if len(attrs) == 0 { + return "" + } + var builder strings.Builder + for index := 0; index < len(attrs); index += 2 { + key := fmt.Sprint(attrs[index]) + value := any("") + if index+1 < len(attrs) { + value = attrs[index+1] + } + builder.WriteByte(' ') + builder.WriteString(key) + builder.WriteByte('=') + builder.WriteString(quoteValue(value)) + } + return builder.String() +} + +func quoteValue(value any) string { + text := fmt.Sprint(value) + if text == "" || strings.ContainsAny(text, " \t\r\n\"=") { + return strconv.Quote(text) + } + return text +} diff --git a/arkruntime/lib/environments/heartbeat.go b/arkruntime/lib/environments/heartbeat.go new file mode 100644 index 0000000..88fa267 --- /dev/null +++ b/arkruntime/lib/environments/heartbeat.go @@ -0,0 +1,133 @@ +// Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +// SPDX-License-Identifier: Apache-2.0 +package environments + +import ( + "context" + "net/http" + "time" + + selfhosted "github.com/volcengine/ark-runtime-go/arkruntime/selfhosted" +) + +const ( + heartbeatDefault = 30 * time.Second + heartbeatFloor = time.Second +) + +type heartbeatStopCause string + +const ( + heartbeatStopCauseLeaseLost heartbeatStopCause = "lease_lost" + heartbeatStopCauseLeaseNotExtended heartbeatStopCause = "lease_not_extended" + heartbeatStopCauseStopRequested heartbeatStopCause = "stop_requested" + heartbeatStopCauseHeartbeatLost heartbeatStopCause = "heartbeat_lost" + heartbeatStopCausePermanentFailure heartbeatStopCause = "heartbeat_permanent_failure" +) + +func (w *EnvironmentWorker) heartbeatLoop(ctx context.Context, item selfhosted.WorkItem, api selfhosted.API, cancel context.CancelFunc, markStopped func(heartbeatStopCause)) { + interval := clampHeartbeatInterval(heartbeatDefault/2, heartbeatDefault) + ttl := heartbeatDefault + logger := w.logger().With("component", "environment-worker", "work_id", item.ID, "session_id", item.SessionIDValue()) + last := item.LatestHeartbeatValue() + if last == "" { + last = selfhosted.ExpectedLastHeartbeatNoHeartbeat + } + lastSuccess := time.Now() + beat := func() bool { + beatCtx, beatCancel := context.WithTimeout(ctx, interval) + defer beatCancel() + resp, err := api.HeartbeatWork(beatCtx, selfhosted.HeartbeatWorkRequest{ + EnvironmentID: item.EnvironmentID, + WorkID: item.ID, + ExpectedLastHeartbeat: last, + DesiredTTLSeconds: int(ttl / time.Second), + }) + if err != nil { + if selfhosted.IsStatus(err, http.StatusPreconditionFailed) { + logger.Warn("heartbeat lease lost", "err", err) + if markStopped != nil { + markStopped(heartbeatStopCauseLeaseLost) + } + cancel() + return false + } + if ctx.Err() != nil { + return false + } + if isFatal4xx(err) { + logger.Error("heartbeat permanent failure", "err", err) + if markStopped != nil { + markStopped(heartbeatStopCausePermanentFailure) + } + cancel() + return false + } + if stale := time.Since(lastSuccess); stale > ttl { + logger.Error("heartbeat staleness exceeded", "since_last_success", stale, "ttl", ttl, "err", err) + if markStopped != nil { + markStopped(heartbeatStopCauseHeartbeatLost) + } + cancel() + return false + } + logger.Warn("heartbeat failed", "since_last_success", time.Since(lastSuccess), "ttl", ttl, "err", err) + return true + } + if resp == nil { + logger.Warn("heartbeat empty response") + return true + } + lastSuccess = time.Now() + if resp.LastHeartbeat != "" { + last = resp.LastHeartbeat + } + if resp.TTLSeconds > 0 { + ttl = time.Duration(resp.TTLSeconds) * time.Second + interval = clampHeartbeatInterval(ttl/2, heartbeatDefault) + } + state := resp.State + if state == selfhosted.WorkStateStopping || state == selfhosted.WorkStateStopped { + logger.Info("heartbeat stop requested", "state", state) + if markStopped != nil { + markStopped(heartbeatStopCauseStopRequested) + } + cancel() + return false + } + if resp.LeaseExtended != nil && !*resp.LeaseExtended { + logger.Warn("heartbeat lease not extended", "state", state) + if markStopped != nil { + markStopped(heartbeatStopCauseLeaseNotExtended) + } + cancel() + return false + } + return true + } + if !beat() { + return + } + for { + timer := time.NewTimer(interval) + select { + case <-ctx.Done(): + timer.Stop() + return + case <-timer.C: + } + if !beat() { + return + } + } +} + +func clampHeartbeatInterval(interval, maxInterval time.Duration) time.Duration { + if interval < heartbeatFloor { + return heartbeatFloor + } + if maxInterval > 0 && interval > maxInterval { + return maxInterval + } + return interval +} diff --git a/arkruntime/lib/environments/integration_test.go b/arkruntime/lib/environments/integration_test.go new file mode 100644 index 0000000..44ff59f --- /dev/null +++ b/arkruntime/lib/environments/integration_test.go @@ -0,0 +1,440 @@ +// Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +// SPDX-License-Identifier: Apache-2.0 +package environments + +import ( + "context" + "io" + "net/http" + "strings" + "sync" + "testing" + "time" + + selfhosted "github.com/volcengine/ark-runtime-go/arkruntime/selfhosted" +) + +type fakeEnvironmentWorkerAPI struct { + mu sync.Mutex + events []selfhosted.Event + sent []selfhosted.Event + pollItem *selfhosted.WorkItem + ackCount int + stopCount int + stops []selfhosted.StopWorkRequest + stopErr error + + heartbeat func(context.Context, selfhosted.HeartbeatWorkRequest) (*selfhosted.HeartbeatResponse, error) + getSession func(context.Context, selfhosted.GetSessionRequest) (*selfhosted.Session, error) + onStop func() +} + +func (f *fakeEnvironmentWorkerAPI) PollWork(context.Context, selfhosted.PollWorkRequest) (*selfhosted.WorkItem, error) { + f.mu.Lock() + defer f.mu.Unlock() + item := f.pollItem + f.pollItem = nil + return item, nil +} + +func (f *fakeEnvironmentWorkerAPI) AckWork(context.Context, selfhosted.AckWorkRequest) error { + f.mu.Lock() + defer f.mu.Unlock() + f.ackCount++ + return nil +} + +func (f *fakeEnvironmentWorkerAPI) HeartbeatWork(ctx context.Context, req selfhosted.HeartbeatWorkRequest) (*selfhosted.HeartbeatResponse, error) { + if f.heartbeat != nil { + return f.heartbeat(ctx, req) + } + return &selfhosted.HeartbeatResponse{ + LastHeartbeat: time.Now().UTC().Format(time.RFC3339Nano), + LeaseExtended: selfhosted.BoolPtr(true), + State: selfhosted.WorkStateActive, + }, nil +} + +func (f *fakeEnvironmentWorkerAPI) StopWork(_ context.Context, req selfhosted.StopWorkRequest) error { + f.mu.Lock() + f.stopCount++ + f.stops = append(f.stops, req) + onStop := f.onStop + f.mu.Unlock() + if onStop != nil { + onStop() + } + return f.stopErr +} + +func (f *fakeEnvironmentWorkerAPI) GetSession(ctx context.Context, req selfhosted.GetSessionRequest) (*selfhosted.Session, error) { + if f.getSession != nil { + return f.getSession(ctx, req) + } + return &selfhosted.Session{ID: req.SessionID}, nil +} + +func (f *fakeEnvironmentWorkerAPI) ListEvents(context.Context, selfhosted.ListEventsRequest) (*selfhosted.ListEventsResponse, error) { + f.mu.Lock() + defer f.mu.Unlock() + if len(f.events) == 0 { + return &selfhosted.ListEventsResponse{}, nil + } + events := append([]selfhosted.Event(nil), f.events...) + f.events = nil + return &selfhosted.ListEventsResponse{Events: events}, nil +} + +func (f *fakeEnvironmentWorkerAPI) SendEvent(_ context.Context, req selfhosted.SendEventRequest) error { + f.mu.Lock() + defer f.mu.Unlock() + f.sent = append(f.sent, req.Event) + return nil +} + +func (f *fakeEnvironmentWorkerAPI) OpenSkill(context.Context, selfhosted.OpenSkillRequest) (*selfhosted.SkillContent, error) { + return &selfhosted.SkillContent{Body: io.NopCloser(strings.NewReader(""))}, nil +} + +func TestEnvironmentWorkerRunHandlesPolledWorkInProcess(t *testing.T) { + root := t.TempDir() + maxIdle := 30 * time.Millisecond + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + api := &fakeEnvironmentWorkerAPI{ + pollItem: &selfhosted.WorkItem{ + ID: "work_local", + EnvironmentID: "env_local", + SessionID: "sess_local", + LeaseID: "lease_local", + }, + events: []selfhosted.Event{ + { + ID: "toolu_local", + Type: selfhosted.EventTypeAgentToolUse, + Name: "bash", + Input: selfhosted.RawJSON(`{"command":"echo local-e2e"}`), + EvaluatedPermission: selfhosted.PermissionAllow, + }, + { + ID: "evt_idle", + Type: selfhosted.EventTypeSessionStatusIdle, + StopReason: &selfhosted.SessionStopReason{Type: selfhosted.SessionStopReasonEndTurn}, + }, + }, + onStop: cancel, + } + w := NewEnvironmentWorker(api, EnvironmentWorkerOptions{ + EnvironmentID: "env_local", + WorkerID: "worker_local", + Workdir: root, + MaxIdle: &maxIdle, + }) + if err := w.Run(ctx); err != nil { + t.Fatalf("Run err = %v", err) + } + + api.mu.Lock() + defer api.mu.Unlock() + if api.ackCount != 1 || api.stopCount != 2 { + t.Fatalf("ack_count=%d stop_count=%d", api.ackCount, api.stopCount) + } + if !api.stops[0].Force { + t.Fatalf("worker stop should be force=true: %+v", api.stops[0]) + } + if api.stops[1].Force { + 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) + } + result := api.sent[0] + if result.Type != selfhosted.EventTypeUserToolResult || result.ToolUseID != "toolu_local" { + t.Fatalf("result event = %+v", result) + } + 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, "local-e2e") { + t.Fatalf("result content = %+v", result.Content) + } +} + +func TestEnvironmentWorkerStopsWorkOnSessionIdleEvent(t *testing.T) { + root := t.TempDir() + maxIdle := 10 * time.Millisecond + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + api := &fakeEnvironmentWorkerAPI{ + pollItem: &selfhosted.WorkItem{ + ID: "work_idle", + EnvironmentID: "env_local", + SessionID: "sess_idle", + LeaseID: "lease_idle", + }, + events: []selfhosted.Event{{ + ID: "evt_idle", + Type: selfhosted.EventTypeSessionStatusIdle, + StopReason: &selfhosted.SessionStopReason{Type: selfhosted.SessionStopReasonEndTurn}, + }}, + onStop: cancel, + } + w := NewEnvironmentWorker(api, EnvironmentWorkerOptions{ + EnvironmentID: "env_local", + WorkerID: "worker_local", + Workdir: root, + MaxIdle: &maxIdle, + }) + if err := w.Run(ctx); err != nil { + t.Fatalf("Run err = %v", err) + } + api.mu.Lock() + defer api.mu.Unlock() + if api.ackCount != 1 || api.stopCount == 0 { + t.Fatalf("ack_count=%d stop_count=%d", api.ackCount, api.stopCount) + } + if got := api.stops[0]; !got.Force || got.WorkID != "work_idle" { + t.Fatalf("worker stop = %+v", got) + } + if len(api.sent) != 0 { + t.Fatalf("sent events = %+v", api.sent) + } +} + +func TestEnvironmentWorkerHandleItemTreatsAlreadyStoppedAsResolved(t *testing.T) { + workdir := t.TempDir() + maxIdle := 10 * time.Millisecond + api := &fakeEnvironmentWorkerAPI{ + stopErr: &selfhosted.APIError{StatusCode: http.StatusConflict, Message: "already stopped"}, + events: []selfhosted.Event{{ + ID: "evt_idle", + Type: selfhosted.EventTypeSessionStatusIdle, + StopReason: &selfhosted.SessionStopReason{Type: selfhosted.SessionStopReasonEndTurn}, + }}, + } + w := NewEnvironmentWorker(api, EnvironmentWorkerOptions{ + EnvironmentID: "env_local", + WorkerID: "worker_local", + Workdir: workdir, + MaxIdle: &maxIdle, + }) + err := w.HandleItem(context.Background(), HandleItemOptions{ + WorkID: "work_local", + EnvironmentID: "env_local", + SessionID: "sess_local", + }) + if err != nil { + t.Fatal(err) + } +} + +func TestEnvironmentWorkerHeartbeatLeaseLostOnPreconditionFailed(t *testing.T) { + var got selfhosted.HeartbeatWorkRequest + api := &fakeEnvironmentWorkerAPI{ + heartbeat: func(_ context.Context, req selfhosted.HeartbeatWorkRequest) (*selfhosted.HeartbeatResponse, error) { + got = req + return nil, &selfhosted.APIError{StatusCode: http.StatusPreconditionFailed, Message: "lease lost"} + }, + } + w := NewEnvironmentWorker(api, EnvironmentWorkerOptions{ + EnvironmentID: "env_local", + WorkerID: "worker_local", + Workdir: t.TempDir(), + }) + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + var cause heartbeatStopCause + w.heartbeatLoop(ctx, selfhosted.WorkItem{ + ID: "work_local", + EnvironmentID: "env_local", + SessionID: "sess_local", + TTLSeconds: 1, + }, api, cancel, func(c heartbeatStopCause) { + cause = c + }) + if cause != heartbeatStopCauseLeaseLost { + t.Fatalf("heartbeat cause = %q", cause) + } + if got.ExpectedLastHeartbeat != selfhosted.ExpectedLastHeartbeatNoHeartbeat { + t.Fatalf("expected_last_heartbeat=%q", got.ExpectedLastHeartbeat) + } + if got.DesiredTTLSeconds != int(heartbeatDefault/time.Second) { + t.Fatalf("desired_ttl_seconds=%d", got.DesiredTTLSeconds) + } +} + +func TestEnvironmentWorkerSkipsStopAfterLeaseOwnershipLoss(t *testing.T) { + api := &fakeEnvironmentWorkerAPI{ + heartbeat: func(context.Context, selfhosted.HeartbeatWorkRequest) (*selfhosted.HeartbeatResponse, error) { + return nil, &selfhosted.APIError{StatusCode: http.StatusPreconditionFailed, Message: "lease lost"} + }, + getSession: func(ctx context.Context, _ selfhosted.GetSessionRequest) (*selfhosted.Session, error) { + <-ctx.Done() + return nil, ctx.Err() + }, + } + w := NewEnvironmentWorker(api, EnvironmentWorkerOptions{ + EnvironmentID: "env_local", + WorkerID: "worker_local", + Workdir: t.TempDir(), + }) + if err := w.HandleItem(context.Background(), HandleItemOptions{ + WorkID: "work_local", + EnvironmentID: "env_local", + SessionID: "sess_local", + }); err != nil { + t.Fatal(err) + } + api.mu.Lock() + defer api.mu.Unlock() + if api.stopCount != 0 { + t.Fatalf("stop_count=%d stops=%+v", api.stopCount, api.stops) + } +} + +func TestEnvironmentWorkerWaitsForHeartbeatCauseBeforeStopping(t *testing.T) { + heartbeatStarted := make(chan struct{}) + api := &fakeEnvironmentWorkerAPI{ + heartbeat: func(ctx context.Context, _ selfhosted.HeartbeatWorkRequest) (*selfhosted.HeartbeatResponse, error) { + close(heartbeatStarted) + <-ctx.Done() + return nil, &selfhosted.APIError{StatusCode: http.StatusPreconditionFailed, Message: "lease lost"} + }, + getSession: func(context.Context, selfhosted.GetSessionRequest) (*selfhosted.Session, error) { + <-heartbeatStarted + return nil, &selfhosted.APIError{StatusCode: http.StatusInternalServerError, Message: "session lookup failed"} + }, + } + w := NewEnvironmentWorker(api, EnvironmentWorkerOptions{ + EnvironmentID: "env_local", + WorkerID: "worker_local", + Workdir: t.TempDir(), + }) + if err := w.HandleItem(context.Background(), HandleItemOptions{ + WorkID: "work_local", + EnvironmentID: "env_local", + SessionID: "sess_local", + }); err == nil { + t.Fatal("expected session lookup failure") + } + api.mu.Lock() + defer api.mu.Unlock() + if api.stopCount != 0 { + t.Fatalf("stop_count=%d stops=%+v", api.stopCount, api.stops) + } +} + +func TestShouldStopItemOnlyWhileOwnershipIsKnown(t *testing.T) { + tests := []struct { + cause heartbeatStopCause + want bool + }{ + {cause: "", want: true}, + {cause: heartbeatStopCauseStopRequested, want: true}, + {cause: heartbeatStopCauseLeaseLost, want: false}, + {cause: heartbeatStopCauseLeaseNotExtended, want: false}, + {cause: heartbeatStopCauseHeartbeatLost, want: false}, + {cause: heartbeatStopCausePermanentFailure, want: false}, + } + for _, test := range tests { + if got := shouldStopItem(test.cause); got != test.want { + t.Fatalf("shouldStopItem(%q)=%v want=%v", test.cause, got, test.want) + } + } +} + +func TestEnvironmentWorkerHeartbeatStopsOnPermanent4xx(t *testing.T) { + api := &fakeEnvironmentWorkerAPI{ + heartbeat: func(context.Context, selfhosted.HeartbeatWorkRequest) (*selfhosted.HeartbeatResponse, error) { + return nil, &selfhosted.APIError{StatusCode: http.StatusUnauthorized, Message: "bad key"} + }, + } + w := NewEnvironmentWorker(api, EnvironmentWorkerOptions{ + EnvironmentID: "env_local", + WorkerID: "worker_local", + Workdir: t.TempDir(), + }) + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + var cause heartbeatStopCause + w.heartbeatLoop(ctx, selfhosted.WorkItem{ + ID: "work_local", + EnvironmentID: "env_local", + SessionID: "sess_local", + }, api, cancel, func(c heartbeatStopCause) { + cause = c + }) + if cause != heartbeatStopCausePermanentFailure { + t.Fatalf("heartbeat cause = %q", cause) + } + if ctx.Err() == nil { + t.Fatal("heartbeat should cancel session context") + } +} + +func TestEnvironmentWorkerHeartbeatPrefersStopRequestedStateOverLeaseNotExtended(t *testing.T) { + api := &fakeEnvironmentWorkerAPI{ + heartbeat: func(context.Context, selfhosted.HeartbeatWorkRequest) (*selfhosted.HeartbeatResponse, error) { + return &selfhosted.HeartbeatResponse{ + LastHeartbeat: time.Now().UTC().Format(time.RFC3339Nano), + LeaseExtended: selfhosted.BoolPtr(false), + State: selfhosted.WorkStateStopping, + }, nil + }, + } + w := NewEnvironmentWorker(api, EnvironmentWorkerOptions{ + EnvironmentID: "env_local", + WorkerID: "worker_local", + Workdir: t.TempDir(), + }) + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + var cause heartbeatStopCause + w.heartbeatLoop(ctx, selfhosted.WorkItem{ + ID: "work_local", + EnvironmentID: "env_local", + SessionID: "sess_local", + }, api, cancel, func(c heartbeatStopCause) { + cause = c + }) + if cause != heartbeatStopCauseStopRequested { + t.Fatalf("heartbeat cause = %q", cause) + } + if ctx.Err() == nil { + t.Fatal("heartbeat should cancel session context") + } +} + +func TestEnvironmentWorkerHeartbeatUsesAnthropicDefaultTTL(t *testing.T) { + var got selfhosted.HeartbeatWorkRequest + var requestTimeout time.Duration + api := &fakeEnvironmentWorkerAPI{ + heartbeat: func(ctx context.Context, req selfhosted.HeartbeatWorkRequest) (*selfhosted.HeartbeatResponse, error) { + got = req + deadline, ok := ctx.Deadline() + if !ok { + t.Fatal("heartbeat request has no deadline") + } + requestTimeout = time.Until(deadline) + return nil, &selfhosted.APIError{StatusCode: http.StatusPreconditionFailed, Message: "lease lost"} + }, + } + w := NewEnvironmentWorker(api, EnvironmentWorkerOptions{ + EnvironmentID: "env_local", + WorkerID: "worker_local", + Workdir: t.TempDir(), + }) + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + w.heartbeatLoop(ctx, selfhosted.WorkItem{ + ID: "work_local", + EnvironmentID: "env_local", + SessionID: "sess_local", + }, api, cancel, nil) + if got.DesiredTTLSeconds != int(heartbeatDefault/time.Second) { + t.Fatalf("desired_ttl_seconds=%d", got.DesiredTTLSeconds) + } + if requestTimeout < heartbeatDefault/2-time.Second || requestTimeout > heartbeatDefault/2 { + t.Fatalf("heartbeat request timeout=%s", requestTimeout) + } +} diff --git a/arkruntime/lib/environments/poller.go b/arkruntime/lib/environments/poller.go new file mode 100644 index 0000000..c7c8380 --- /dev/null +++ b/arkruntime/lib/environments/poller.go @@ -0,0 +1,290 @@ +// Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +// SPDX-License-Identifier: Apache-2.0 +package environments + +import ( + "context" + crand "crypto/rand" + "encoding/hex" + "errors" + "fmt" + "io" + "log" + "math/rand" + "net/http" + "os" + "time" + + "github.com/volcengine/ark-runtime-go/arkruntime" + "github.com/volcengine/ark-runtime-go/arkruntime/internal/selfhostedlog" + selfhosted "github.com/volcengine/ark-runtime-go/arkruntime/selfhosted" +) + +const ( + defaultPollBlockMS = 999 + pollBackoffCap = 60 * time.Second + stopTimeout = 10 * time.Second +) + +// WorkPollerOptions 配置 WorkPoller。 +type WorkPollerOptions struct { + EnvironmentID string + WorkerID string + BlockMS int + ReclaimOlderThanMS int + Drain bool + Logger *log.Logger +} + +// WorkPoller 负责从 environment work queue 里 poll 并 ack work。 +type WorkPoller struct { + ctx context.Context + api selfhosted.API + opts WorkPollerOptions + logger *selfhostedlog.Logger + current *selfhosted.WorkItem + err error + pendingStop func() + failures int + discards int + closed bool +} + +// NewWorkPoller 创建 work poller。 +func NewWorkPoller(ctx context.Context, api selfhosted.API, opts WorkPollerOptions) *WorkPoller { + logger := selfhostedlog.New(opts.Logger) + if opts.WorkerID == "" { + opts.WorkerID = defaultWorkerID() + } + p := &WorkPoller{ + ctx: ctx, + api: api, + opts: opts, + logger: logger.With("component", "work-poller", "environment_id", opts.EnvironmentID), + } + if opts.EnvironmentID == "" { + p.err = errors.New("environments: EnvironmentID is required") + } + if api == nil { + p.err = errors.New("environments: API is required") + } + return p +} + +// NewWorkPollerForClient 创建绑定到 arkruntime.Client 的 work poller。 +func NewWorkPollerForClient(ctx context.Context, client *arkruntime.Client, opts WorkPollerOptions) *WorkPoller { + return NewWorkPoller(ctx, selfhosted.NewClientAPI(client), opts) +} + +// Next 推进到下一条已 ack 的 work。 +func (p *WorkPoller) Next() bool { + p.runPendingStop() + if p.err != nil || p.closed { + return false + } + for { + if p.ctx.Err() != nil { + return false + } + req := selfhosted.PollWorkRequest{ + EnvironmentID: p.opts.EnvironmentID, + WorkerID: p.opts.WorkerID, + MaxItems: 1, + ReclaimOlderThanMS: p.opts.ReclaimOlderThanMS, + } + if p.opts.BlockMS >= 0 { + req.BlockMS = defaultPollBlockMS + if p.opts.BlockMS > 0 { + req.BlockMS = p.opts.BlockMS + } + } + item, err := p.api.PollWork(p.ctx, req) + if err != nil { + if p.ctx.Err() != nil { + return false + } + if isFatal4xx(err) { + p.err = fmt.Errorf("environments: poll work: %w", err) + return false + } + p.failures++ + sleepFor := jitter(backoff(p.failures)/2, backoff(p.failures)) + p.logger.Warn("poll work failed", "err", err, "sleep", sleepFor) + sleep(p.ctx, sleepFor) + continue + } + p.failures = 0 + if item == nil || item.ID == "" { + if p.opts.Drain { + return false + } + sleep(p.ctx, jitter(time.Second, 3*time.Second)) + continue + } + if item.EnvironmentID == "" { + item.EnvironmentID = p.opts.EnvironmentID + } + if item.SessionIDValue() == "" { + p.discardInvalidWork(*item, "work item does not contain session id") + continue + } + if err := p.api.AckWork(p.ctx, selfhosted.AckWorkRequest{ + EnvironmentID: item.EnvironmentID, + WorkID: item.ID, + WorkerID: p.opts.WorkerID, + }); err != nil { + p.logger.Warn("ack work failed", "work_id", item.ID, "err", err) + // ACK 是 queued -> starting 的竞争操作。失败时无法证明当前 + // worker 拥有该 work,不能 force stop,否则可能终止竞争胜者。 + if isResolvedStatus(err) { + continue + } + if isFatal4xx(err) { + p.err = fmt.Errorf("environments: ack work: %w", err) + return false + } + p.backoffDiscard(*item) + continue + } + p.current = item + p.pendingStop = p.makeStopClosure(*item) + p.discards = 0 + p.logger.Info("claimed work", "work_id", item.ID, "session_id", item.SessionIDValue()) + return true + } +} + +// Current 返回最近一次 Next claim 到的 work。 +func (p *WorkPoller) Current() *selfhosted.WorkItem { + return p.current +} + +// Err 返回最近一次 Next 的错误;不可恢复错误会阻止后续 Next。 +func (p *WorkPoller) Err() error { + return p.err +} + +// Close 停止 poller,并对最后 yield 的 work 发送 StopWork。 +func (p *WorkPoller) Close() error { + if p.closed { + return nil + } + p.closed = true + p.runPendingStop() + return nil +} + +func (p *WorkPoller) runPendingStop() { + if p.pendingStop == nil { + return + } + stop := p.pendingStop + p.pendingStop = nil + p.current = nil + stop() +} + +func (p *WorkPoller) makeStopClosure(item selfhosted.WorkItem) func() { + return func() { + ctx, cancel := context.WithTimeout(context.Background(), stopTimeout) + defer cancel() + if err := p.api.StopWork(ctx, selfhosted.StopWorkRequest{ + EnvironmentID: item.EnvironmentID, + WorkID: item.ID, + }); err != nil && !isResolvedStatus(err) { + p.logger.Warn("stop work failed", "work_id", item.ID, "err", err) + } + } +} + +func (p *WorkPoller) discardInvalidWork(item selfhosted.WorkItem, _ string) { + if item.ID == "" { + return + } + if err := p.api.AckWork(p.ctx, selfhosted.AckWorkRequest{ + EnvironmentID: item.EnvironmentID, + WorkID: item.ID, + WorkerID: p.opts.WorkerID, + }); err != nil { + p.logger.Warn("ack invalid work failed", "work_id", item.ID, "err", err) + return + } + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + if err := p.api.StopWork(ctx, selfhosted.StopWorkRequest{ + EnvironmentID: item.EnvironmentID, + WorkID: item.ID, + Force: true, + }); err != nil && !isResolvedStatus(err) { + p.logger.Warn("stop invalid work failed", "work_id", item.ID, "err", err) + } + p.backoffDiscard(item) +} + +func (p *WorkPoller) backoffDiscard(item selfhosted.WorkItem) { + p.discards++ + sleepFor := jitter(backoff(p.discards)/2, backoff(p.discards)) + p.logger.Warn("backing off after unprocessable work item", "work_id", item.ID, "sleep", sleepFor) + sleep(p.ctx, sleepFor) +} + +func isResolvedStatus(err error) bool { + return selfhosted.IsStatus(err, http.StatusNotFound) || + selfhosted.IsStatus(err, http.StatusConflict) || + selfhosted.IsStatus(err, http.StatusPreconditionFailed) +} + +func sleep(ctx context.Context, d time.Duration) { + timer := time.NewTimer(d) + defer timer.Stop() + select { + case <-ctx.Done(): + case <-timer.C: + } +} + +func isFatal4xx(err error) bool { + var apiErr *selfhosted.APIError + if !errors.As(err, &apiErr) { + return false + } + return apiErr.StatusCode >= 400 && apiErr.StatusCode < 500 && + apiErr.StatusCode != http.StatusRequestTimeout && + apiErr.StatusCode != http.StatusConflict && + apiErr.StatusCode != http.StatusTooManyRequests +} + +func backoff(n int) time.Duration { + if n <= 0 { + return time.Second + } + if n > 6 { + return pollBackoffCap + } + if d := time.Duration(1<= 'a' && c <= 'z' || c >= 'A' && c <= 'Z' || c >= '0' && c <= '9' { + continue + } + if c == '.' || c == '_' || c == '-' { + continue + } + return false + } + return true +} diff --git a/arkruntime/lib/environments/worker.go b/arkruntime/lib/environments/worker.go new file mode 100644 index 0000000..47af04c --- /dev/null +++ b/arkruntime/lib/environments/worker.go @@ -0,0 +1,383 @@ +// Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +// SPDX-License-Identifier: Apache-2.0 +// Package environments 提供 self-hosted environment worker 的 SDK 风格组合层。 +package environments + +import ( + "context" + "errors" + "fmt" + "log" + "os" + "path/filepath" + "sync/atomic" + "time" + + "github.com/volcengine/ark-runtime-go/arkruntime" + "github.com/volcengine/ark-runtime-go/arkruntime/internal/selfhostedlog" + selfhosted "github.com/volcengine/ark-runtime-go/arkruntime/selfhosted" + "github.com/volcengine/ark-runtime-go/arkruntime/selfhosted/envinit" + "github.com/volcengine/ark-runtime-go/arkruntime/tools/agenttoolset" +) + +// EnvironmentWorkerOptions 配置 EnvironmentWorker。 +type EnvironmentWorkerOptions struct { + // EnvironmentID 是 worker 持续 poll work 的 self-hosted environment,Run 必填。 + EnvironmentID string + // WorkerID 是上报给控制面的 worker 标识,空值时自动生成。 + WorkerID string + // Workdir 是每个 session 工作目录的根目录,空值时使用进程当前目录。 + Workdir string + // UnrestrictedPaths 控制文件工具是否允许访问 Workdir 之外的路径。 + UnrestrictedPaths bool + // ToolContext 提供默认工具集使用的高级执行环境配置。 + ToolContext *agenttoolset.AgentToolContext + // Tools 是所有 session 复用的固定工具集合;设置 ToolsFunc 时忽略此字段。 + Tools *agenttoolset.Set + // ToolsFunc 为每个 session 创建绑定到其工作目录的工具集合。 + ToolsFunc func(env *agenttoolset.AgentToolContext) (*agenttoolset.Set, error) + // MaxIdle 是 session 在 end_turn idle 后继续等待事件的时间,nil 使用默认值。 + MaxIdle *time.Duration + // CustomTools 是按名称注册的自定义工具。 + CustomTools map[string]agenttoolset.Tool + // Logger 接收 worker 生命周期日志,nil 使用 log.Default。 + Logger *log.Logger +} + +// EnvironmentWorker 组合 work poll、session 初始化、tool event loop 和 heartbeat。 +type EnvironmentWorker struct { + api selfhosted.API + opts EnvironmentWorkerOptions +} + +// NewEnvironmentWorker 创建绑定到控制面 API 的 environment worker。 +func NewEnvironmentWorker(api selfhosted.API, opts EnvironmentWorkerOptions) *EnvironmentWorker { + if opts.Workdir == "" { + cwd, err := os.Getwd() + if err == nil { + opts.Workdir = cwd + } else { + opts.Workdir = "." + } + } + if opts.WorkerID == "" { + opts.WorkerID = defaultWorkerID() + } + return &EnvironmentWorker{api: api, opts: opts} +} + +// NewEnvironmentWorkerForClient 创建绑定到 arkruntime.Client 的 environment worker。 +func NewEnvironmentWorkerForClient(client *arkruntime.Client, opts EnvironmentWorkerOptions) *EnvironmentWorker { + return NewEnvironmentWorker(selfhosted.NewClientAPI(client), opts) +} + +// Run 持续 poll environment work queue 并处理 session work。 +func (w *EnvironmentWorker) Run(ctx context.Context) error { + if ctx == nil { + ctx = context.Background() + } + environmentID := w.opts.EnvironmentID + if environmentID == "" { + return errors.New("environments: EnvironmentID is required") + } + if w.api == nil { + return errors.New("environments: API is required") + } + logger := w.logger().With("component", "environment-worker", "environment_id", environmentID) + + poller := NewWorkPoller(ctx, w.api, WorkPollerOptions{ + EnvironmentID: environmentID, + WorkerID: w.opts.WorkerID, + Logger: w.opts.Logger, + }) + defer func() { _ = poller.Close() }() + + for poller.Next() { + work := poller.Current() + if work == nil { + continue + } + w.runClaimedWork(ctx, *work, logger) + } + if err := poller.Err(); err != nil && ctx.Err() == nil { + return err + } + return nil +} + +func (w *EnvironmentWorker) runClaimedWork(ctx context.Context, item selfhosted.WorkItem, logger *selfhostedlog.Logger) { + if err := w.handleItem(ctx, item, false); err != nil && + !isBenignWorkerExit(err) { + logger.Warn("handle work failed", "work_id", item.ID, "session_id", item.SessionIDValue(), "err", err) + } +} + +// HandleItemOptions 指定一个已经 claim 的 work item。 +type HandleItemOptions struct { + WorkID string + EnvironmentID string + SessionID string + LatestHeartbeatAt string +} + +// HandleItem 处理单个已 claim 的 session work,适合作为 sandbox/container 入口。 +func (w *EnvironmentWorker) HandleItem(ctx context.Context, opts HandleItemOptions) error { + if ctx == nil { + ctx = context.Background() + } + item, err := w.resolveHandleItem(opts) + if err != nil { + return err + } + if err := w.handleItem(ctx, item, true); err != nil && + !isBenignWorkerExit(err) { + return err + } + return nil +} + +func (w *EnvironmentWorker) handleItem(ctx context.Context, item selfhosted.WorkItem, useWorkdirAsSession bool) (err error) { + if w.api == nil { + return errors.New("environments: API is required") + } + if item.EnvironmentID == "" { + item.EnvironmentID = firstNonEmpty(w.opts.EnvironmentID, os.Getenv("MA_ENVIRONMENT_ID")) + } + if item.ID == "" { + return errors.New("work item id must not be empty") + } + sessionID := item.SessionIDValue() + if sessionID == "" { + return errors.New("work item does not contain session id") + } + workdir, err := w.workdirFor(sessionID, useWorkdirAsSession) + if err != nil { + return err + } + api := w.api + logger := w.logger().With("component", "environment-worker", "work_id", item.ID, "session_id", sessionID, "workdir", workdir) + + workCtx, cancel := context.WithCancel(ctx) + defer cancel() + var heartbeatCause atomic.Value + heartbeatDone := make(chan struct{}) + go func() { + defer close(heartbeatDone) + w.heartbeatLoop(workCtx, item, api, cancel, func(cause heartbeatStopCause) { + heartbeatCause.Store(string(cause)) + }) + }() + + defer func() { + cancel() + <-heartbeatDone + cause := loadHeartbeatStopCause(&heartbeatCause) + if shouldStopItem(cause) { + _ = w.stopItem(api, item) + } else { + logger.Info("skip stop work after heartbeat ownership became uncertain", "cause", cause) + } + if cause == heartbeatStopCauseStopRequested && errors.Is(err, context.Canceled) { + err = nil + } + }() + + session, err := api.GetSession(workCtx, selfhosted.GetSessionRequest{SessionID: sessionID}) + if err != nil { + return fmt.Errorf("get session: %w", err) + } + if session == nil { + err = errors.New("session response is empty") + return err + } + if session.ID == "" { + session.ID = sessionID + } + toolEnv := w.toolContext(workdir) + initOpts := toolEnv.InitOptions() + initOpts.Workdir = workdir + initOpts.Logger = w.opts.Logger + if err := envinit.New(api, initOpts).Setup(workCtx, session); err != nil { + return fmt.Errorf("setup session environment: %w", err) + } + tools, owned, err := w.toolsFor(toolEnv) + if err != nil { + return fmt.Errorf("create toolset: %w", err) + } + if owned { + defer func() { _ = tools.Close() }() + } + store, err := selfhosted.NewFileToolResultStore(workdir) + if err != nil { + return fmt.Errorf("create tool result store: %w", err) + } + runner := selfhosted.NewSessionToolRunner(workCtx, api, sessionID, selfhosted.SessionToolRunnerOptions{ + WorkID: item.ID, + Tools: tools, + CustomTools: w.opts.CustomTools, + ResultStore: store, + MaxIdle: w.opts.MaxIdle, + Logger: w.opts.Logger, + }) + defer func() { _ = runner.Close() }() + for runner.Next() { + result := runner.Current() + logger.Info("tool event handled", "tool_use_id", result.ToolUseID, "tool", result.Name, "custom", result.Custom, "posted", result.Posted) + } + if err := runner.Err(); err != nil && + !errors.Is(err, selfhosted.ErrIdleTimeout) && + !errors.Is(err, selfhosted.ErrSessionTerminated) && + !errors.Is(err, context.Canceled) && + !errors.Is(err, context.DeadlineExceeded) { + return err + } + return runner.Err() +} + +func (w *EnvironmentWorker) toolsFor(env *agenttoolset.AgentToolContext) (*agenttoolset.Set, bool, error) { + if w.opts.ToolsFunc != nil { + tools, err := w.opts.ToolsFunc(env) + if err != nil { + return nil, false, err + } + if tools == nil { + return nil, false, errors.New("tools func returned nil") + } + return tools, true, nil + } + if w.opts.Tools != nil { + return w.opts.Tools, false, nil + } + tools, err := agenttoolset.AgentToolset20260401(env) + if err != nil { + return nil, false, err + } + return tools, true, nil +} + +func (w *EnvironmentWorker) toolContext(workdir string) *agenttoolset.AgentToolContext { + base := agenttoolset.AgentToolContext{} + if w.opts.ToolContext != nil { + base = *w.opts.ToolContext + } + base.Workdir = workdir + if w.opts.UnrestrictedPaths { + base.UnrestrictedPaths = true + } + if base.Env != nil { + env := make(map[string]string, len(base.Env)) + for k, v := range base.Env { + env[k] = v + } + base.Env = env + } + return &base +} + +func (w *EnvironmentWorker) workdirFor(sessionID string, useWorkdirAsSession bool) (string, error) { + root := w.opts.Workdir + if root == "" { + root = "." + } + absRoot, err := filepath.Abs(root) + if err != nil { + return "", err + } + if useWorkdirAsSession { + return absRoot, nil + } + return sessionWorkdir(absRoot, sessionID) +} + +func (w *EnvironmentWorker) resolveHandleItem(opts HandleItemOptions) (selfhosted.WorkItem, error) { + item, err := workItemFromOptions(opts) + if err != nil { + return item, err + } + return item, nil +} + +func (w *EnvironmentWorker) stopItem(api selfhosted.API, item selfhosted.WorkItem) error { + stopCtx, stopCancel := context.WithTimeout(context.Background(), stopTimeout) + defer stopCancel() + req := selfhosted.StopWorkRequest{ + EnvironmentID: item.EnvironmentID, + WorkID: item.ID, + Force: true, + } + if err := api.StopWork(stopCtx, req); err != nil { + if selfhosted.IsStatus(err, 409) || selfhosted.IsStatus(err, 412) { + w.logger().Info("stop work already resolved", "work_id", item.ID, "err", err) + return nil + } + w.logger().Warn("stop work failed", "work_id", item.ID, "err", err) + return err + } + return nil +} + +func (w *EnvironmentWorker) logger() *selfhostedlog.Logger { + return selfhostedlog.New(w.opts.Logger) +} + +func loadHeartbeatStopCause(value *atomic.Value) heartbeatStopCause { + raw := value.Load() + if raw == nil { + return "" + } + cause, _ := raw.(string) + return heartbeatStopCause(cause) +} + +func shouldStopItem(cause heartbeatStopCause) bool { + switch cause { + case heartbeatStopCauseLeaseLost, + heartbeatStopCauseLeaseNotExtended, + heartbeatStopCauseHeartbeatLost, + heartbeatStopCausePermanentFailure: + return false + default: + return true + } +} + +func isBenignWorkerExit(err error) bool { + return errors.Is(err, context.Canceled) || + errors.Is(err, context.DeadlineExceeded) || + errors.Is(err, selfhosted.ErrIdleTimeout) || + errors.Is(err, selfhosted.ErrSessionTerminated) +} + +func workItemFromOptions(opts HandleItemOptions) (selfhosted.WorkItem, error) { + workID := firstNonEmpty(opts.WorkID, os.Getenv("MA_WORK_ID")) + environmentID := firstNonEmpty(opts.EnvironmentID, os.Getenv("MA_ENVIRONMENT_ID")) + sessionID := firstNonEmpty(opts.SessionID, os.Getenv("MA_SESSION_ID")) + latestHeartbeatAt := firstNonEmpty(opts.LatestHeartbeatAt, os.Getenv("MA_LATEST_HEARTBEAT_AT")) + if workID == "" { + return selfhosted.WorkItem{}, errors.New("environments: work id is required") + } + if environmentID == "" { + return selfhosted.WorkItem{}, errors.New("environments: environment id is required") + } + if sessionID == "" { + return selfhosted.WorkItem{}, errors.New("environments: session id is required") + } + return selfhosted.WorkItem{ + ID: workID, + EnvironmentID: environmentID, + LatestHeartbeatAt: latestHeartbeatAt, + Data: selfhosted.WorkData{ + Type: "session", + ID: sessionID, + }, + }, nil +} + +func firstNonEmpty(values ...string) string { + for _, value := range values { + if value != "" { + return value + } + } + return "" +} diff --git a/arkruntime/lib/environments/worker_test.go b/arkruntime/lib/environments/worker_test.go new file mode 100644 index 0000000..c03bc4e --- /dev/null +++ b/arkruntime/lib/environments/worker_test.go @@ -0,0 +1,50 @@ +// Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +// SPDX-License-Identifier: Apache-2.0 +package environments + +import "testing" + +func TestHandleItemOptionsBuildsWorkItem(t *testing.T) { + got, err := workItemFromOptions(HandleItemOptions{ + WorkID: testWorkID, + EnvironmentID: "env_1", + SessionID: testSessionID, + LatestHeartbeatAt: "2026-08-11T00:00:00Z", + }) + if err != nil { + t.Fatal(err) + } + if got.ID != testWorkID || got.EnvironmentID != "env_1" || got.SessionIDValue() != testSessionID { + t.Fatalf("work item = %+v", got) + } + if got.LatestHeartbeatValue() != "2026-08-11T00:00:00Z" { + t.Fatalf("work heartbeat = %+v", got) + } + if got.Data.Type != "session" || got.Data.ID != testSessionID { + t.Fatalf("work data = %+v", got.Data) + } +} + +func TestHandleItemOptionsUsesEnvironmentFallbacks(t *testing.T) { + t.Setenv("MA_WORK_ID", "work_env") + t.Setenv("MA_ENVIRONMENT_ID", "env_env") + t.Setenv("MA_SESSION_ID", "sess_env") + t.Setenv("MA_LATEST_HEARTBEAT_AT", "2026-08-11T00:00:00Z") + + got, err := workItemFromOptions(HandleItemOptions{}) + if err != nil { + t.Fatal(err) + } + if got.ID != "work_env" || got.EnvironmentID != "env_env" || got.SessionIDValue() != "sess_env" { + t.Fatalf("work item = %+v", got) + } + if got.LatestHeartbeatValue() != "2026-08-11T00:00:00Z" { + t.Fatalf("work heartbeat = %+v", got) + } +} + +func TestHandleItemOptionsRequiresWorkID(t *testing.T) { + if _, err := workItemFromOptions(HandleItemOptions{EnvironmentID: "env_1", SessionID: "sess_1"}); err == nil { + t.Fatal("expected work id error") + } +} diff --git a/arkruntime/model/agent/oas_json_gen.go b/arkruntime/model/agent/oas_json_gen.go index e7f8327..cb0fb0c 100644 --- a/arkruntime/model/agent/oas_json_gen.go +++ b/arkruntime/model/agent/oas_json_gen.go @@ -109,6 +109,12 @@ func (s *Agent) encodeFields(e *jx.Encoder) { e.ArrEnd() } } + { + if s.DisplayName.Set { + e.FieldStart("display_name") + s.DisplayName.Encode(e) + } + } { e.FieldStart("created_at") e.Str(s.CreatedAt) @@ -119,7 +125,7 @@ func (s *Agent) encodeFields(e *jx.Encoder) { } } -var jsonFieldsNameOfAgent = [15]string{ +var jsonFieldsNameOfAgent = [16]string{ 0: "id", 1: "type", 2: "name", @@ -133,8 +139,9 @@ var jsonFieldsNameOfAgent = [15]string{ 10: "multiagent", 11: "metadata", 12: "tags", - 13: "created_at", - 14: "updated_at", + 13: "display_name", + 14: "created_at", + 15: "updated_at", } // Decode decodes Agent from json. @@ -310,8 +317,18 @@ func (s *Agent) Decode(d *jx.Decoder) error { }(); err != nil { return errors.Wrap(err, "decode field \"tags\"") } + case "display_name": + if err := func() error { + s.DisplayName.Reset() + if err := s.DisplayName.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"display_name\"") + } case "created_at": - requiredBitSet[1] |= 1 << 5 + requiredBitSet[1] |= 1 << 6 if err := func() error { v, err := d.Str() s.CreatedAt = string(v) @@ -323,7 +340,7 @@ func (s *Agent) Decode(d *jx.Decoder) error { return errors.Wrap(err, "decode field \"created_at\"") } case "updated_at": - requiredBitSet[1] |= 1 << 6 + requiredBitSet[1] |= 1 << 7 if err := func() error { v, err := d.Str() s.UpdatedAt = string(v) @@ -345,7 +362,7 @@ func (s *Agent) Decode(d *jx.Decoder) error { var failures []validate.FieldError for i, mask := range [2]uint8{ 0b00010111, - 0b01100000, + 0b11000000, } { if result := (requiredBitSet[i] & mask) ^ mask; result != 0 { // Mask only required fields and check equality to mask using XOR. @@ -472,12 +489,80 @@ func (s *AgentRef) encodeFields(e *jx.Encoder) { s.Version.Encode(e) } } + { + if s.Name.Set { + e.FieldStart("name") + s.Name.Encode(e) + } + } + { + if s.Description.Set { + e.FieldStart("description") + s.Description.Encode(e) + } + } + { + if s.Model.Set { + e.FieldStart("model") + s.Model.Encode(e) + } + } + { + if s.System.Set { + e.FieldStart("system") + s.System.Encode(e) + } + } + { + if s.Tools != nil { + e.FieldStart("tools") + e.ArrStart() + for _, elem := range s.Tools { + elem.Encode(e) + } + e.ArrEnd() + } + } + { + if s.McpServers != nil { + e.FieldStart("mcp_servers") + e.ArrStart() + for _, elem := range s.McpServers { + elem.Encode(e) + } + e.ArrEnd() + } + } + { + if s.Skills != nil { + e.FieldStart("skills") + e.ArrStart() + for _, elem := range s.Skills { + elem.Encode(e) + } + e.ArrEnd() + } + } + { + if s.DisplayName.Set { + e.FieldStart("display_name") + s.DisplayName.Encode(e) + } + } } -var jsonFieldsNameOfAgentRef = [3]string{ - 0: "type", - 1: "id", - 2: "version", +var jsonFieldsNameOfAgentRef = [11]string{ + 0: "type", + 1: "id", + 2: "version", + 3: "name", + 4: "description", + 5: "model", + 6: "system", + 7: "tools", + 8: "mcp_servers", + 9: "skills", + 10: "display_name", } // Decode decodes AgentRef from json. @@ -485,7 +570,7 @@ func (s *AgentRef) Decode(d *jx.Decoder) error { if s == nil { return errors.New("invalid: unable to decode AgentRef to nil") } - var requiredBitSet [1]uint8 + var requiredBitSet [2]uint8 if err := d.ObjBytes(func(d *jx.Decoder, k []byte) error { switch string(k) { @@ -519,6 +604,107 @@ func (s *AgentRef) Decode(d *jx.Decoder) error { }(); err != nil { return errors.Wrap(err, "decode field \"version\"") } + case "name": + if err := func() error { + s.Name.Reset() + if err := s.Name.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"name\"") + } + case "description": + if err := func() error { + s.Description.Reset() + if err := s.Description.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"description\"") + } + case "model": + if err := func() error { + s.Model.Reset() + if err := s.Model.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"model\"") + } + case "system": + if err := func() error { + s.System.Reset() + if err := s.System.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"system\"") + } + case "tools": + if err := func() error { + s.Tools = make([]ToolItem, 0) + if err := d.Arr(func(d *jx.Decoder) error { + var elem ToolItem + if err := elem.Decode(d); err != nil { + return err + } + s.Tools = append(s.Tools, elem) + return nil + }); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"tools\"") + } + case "mcp_servers": + if err := func() error { + s.McpServers = make([]MCPServer, 0) + if err := d.Arr(func(d *jx.Decoder) error { + var elem MCPServer + if err := elem.Decode(d); err != nil { + return err + } + s.McpServers = append(s.McpServers, elem) + return nil + }); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"mcp_servers\"") + } + case "skills": + if err := func() error { + s.Skills = make([]AgentSkillRef, 0) + if err := d.Arr(func(d *jx.Decoder) error { + var elem AgentSkillRef + if err := elem.Decode(d); err != nil { + return err + } + s.Skills = append(s.Skills, elem) + return nil + }); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"skills\"") + } + case "display_name": + if err := func() error { + s.DisplayName.Reset() + if err := s.DisplayName.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"display_name\"") + } default: return d.Skip() } @@ -528,8 +714,9 @@ func (s *AgentRef) Decode(d *jx.Decoder) error { } // Validate required fields. var failures []validate.FieldError - for i, mask := range [1]uint8{ + for i, mask := range [2]uint8{ 0b00000001, + 0b00000000, } { if result := (requiredBitSet[i] & mask) ^ mask; result != 0 { // Mask only required fields and check equality to mask using XOR. @@ -640,12 +827,19 @@ func (s *AgentSkillRef) encodeFields(e *jx.Encoder) { s.Version.Encode(e) } } + { + if s.UseLatest.Set { + e.FieldStart("use_latest") + s.UseLatest.Encode(e) + } + } } -var jsonFieldsNameOfAgentSkillRef = [3]string{ +var jsonFieldsNameOfAgentSkillRef = [4]string{ 0: "type", 1: "skill_id", 2: "version", + 3: "use_latest", } // Decode decodes AgentSkillRef from json. @@ -687,6 +881,16 @@ func (s *AgentSkillRef) Decode(d *jx.Decoder) error { }(); err != nil { return errors.Wrap(err, "decode field \"version\"") } + case "use_latest": + if err := func() error { + s.UseLatest.Reset() + if err := s.UseLatest.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"use_latest\"") + } default: return d.Skip() } @@ -804,6 +1008,12 @@ func (s *CreateAgentRequest) encodeFields(e *jx.Encoder) { s.Description.Encode(e) } } + { + if s.DisplayName.Set { + e.FieldStart("display_name") + s.DisplayName.Encode(e) + } + } { if s.System.Set { e.FieldStart("system") @@ -864,17 +1074,18 @@ func (s *CreateAgentRequest) encodeFields(e *jx.Encoder) { } } -var jsonFieldsNameOfCreateAgentRequest = [10]string{ - 0: "name", - 1: "model", - 2: "description", - 3: "system", - 4: "mcp_servers", - 5: "tools", - 6: "skills", - 7: "multiagent", - 8: "metadata", - 9: "tags", +var jsonFieldsNameOfCreateAgentRequest = [11]string{ + 0: "name", + 1: "model", + 2: "description", + 3: "display_name", + 4: "system", + 5: "mcp_servers", + 6: "tools", + 7: "skills", + 8: "multiagent", + 9: "metadata", + 10: "tags", } // Decode decodes CreateAgentRequest from json. @@ -918,6 +1129,16 @@ func (s *CreateAgentRequest) Decode(d *jx.Decoder) error { }(); err != nil { return errors.Wrap(err, "decode field \"description\"") } + case "display_name": + if err := func() error { + s.DisplayName.Reset() + if err := s.DisplayName.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"display_name\"") + } case "system": if err := func() error { s.System.Reset() @@ -1052,54 +1273,222 @@ func (s *CreateAgentRequest) Decode(d *jx.Decoder) error { result &^= 1 << bitIdx } } - } - if len(failures) > 0 { - return &validate.Error{Fields: failures} + } + if len(failures) > 0 { + return &validate.Error{Fields: failures} + } + + return nil +} + +// MarshalJSON implements stdjson.Marshaler. +func (s *CreateAgentRequest) MarshalJSON() ([]byte, error) { + e := jx.Encoder{} + s.Encode(&e) + return e.Bytes(), nil +} + +// UnmarshalJSON implements stdjson.Unmarshaler. +func (s *CreateAgentRequest) UnmarshalJSON(data []byte) error { + d := jx.DecodeBytes(data) + return s.Decode(d) +} + +// Encode implements json.Marshaler. +func (s CreateAgentRequestMetadata) Encode(e *jx.Encoder) { + e.ObjStart() + s.encodeFields(e) + e.ObjEnd() +} + +// encodeFields implements json.Marshaler. +func (s CreateAgentRequestMetadata) encodeFields(e *jx.Encoder) { + for k, elem := range s { + e.FieldStart(k) + + e.Str(elem) + } +} + +// Decode decodes CreateAgentRequestMetadata from json. +func (s *CreateAgentRequestMetadata) Decode(d *jx.Decoder) error { + if s == nil { + return errors.New("invalid: unable to decode CreateAgentRequestMetadata to nil") + } + m := s.init() + if err := d.ObjBytes(func(d *jx.Decoder, k []byte) error { + var elem string + if err := func() error { + v, err := d.Str() + elem = string(v) + if err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrapf(err, "decode field %q", k) + } + m[string(k)] = elem + return nil + }); err != nil { + return errors.Wrap(err, "decode CreateAgentRequestMetadata") + } + + return nil +} + +// MarshalJSON implements stdjson.Marshaler. +func (s CreateAgentRequestMetadata) MarshalJSON() ([]byte, error) { + e := jx.Encoder{} + s.Encode(&e) + return e.Bytes(), nil +} + +// UnmarshalJSON implements stdjson.Unmarshaler. +func (s *CreateAgentRequestMetadata) UnmarshalJSON(data []byte) error { + d := jx.DecodeBytes(data) + return s.Decode(d) +} + +// Encode implements json.Marshaler. +func (s *CustomToolInputSchema) Encode(e *jx.Encoder) { + e.ObjStart() + s.encodeFields(e) + e.ObjEnd() +} + +// encodeFields encodes fields. +func (s *CustomToolInputSchema) encodeFields(e *jx.Encoder) { + { + if s.Type.Set { + e.FieldStart("type") + s.Type.Encode(e) + } + } + { + if s.Properties.Set { + e.FieldStart("properties") + s.Properties.Encode(e) + } + } + { + if s.Required != nil { + e.FieldStart("required") + e.ArrStart() + for _, elem := range s.Required { + e.Str(elem) + } + e.ArrEnd() + } + } +} + +var jsonFieldsNameOfCustomToolInputSchema = [3]string{ + 0: "type", + 1: "properties", + 2: "required", +} + +// Decode decodes CustomToolInputSchema from json. +func (s *CustomToolInputSchema) Decode(d *jx.Decoder) error { + if s == nil { + return errors.New("invalid: unable to decode CustomToolInputSchema to nil") + } + + if err := d.ObjBytes(func(d *jx.Decoder, k []byte) error { + switch string(k) { + case "type": + if err := func() error { + s.Type.Reset() + if err := s.Type.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"type\"") + } + case "properties": + if err := func() error { + s.Properties.Reset() + if err := s.Properties.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"properties\"") + } + case "required": + if err := func() error { + s.Required = make([]string, 0) + if err := d.Arr(func(d *jx.Decoder) error { + var elem string + v, err := d.Str() + elem = string(v) + if err != nil { + return err + } + s.Required = append(s.Required, elem) + return nil + }); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"required\"") + } + default: + return d.Skip() + } + return nil + }); err != nil { + return errors.Wrap(err, "decode CustomToolInputSchema") } return nil } // MarshalJSON implements stdjson.Marshaler. -func (s *CreateAgentRequest) MarshalJSON() ([]byte, error) { +func (s *CustomToolInputSchema) MarshalJSON() ([]byte, error) { e := jx.Encoder{} s.Encode(&e) return e.Bytes(), nil } // UnmarshalJSON implements stdjson.Unmarshaler. -func (s *CreateAgentRequest) UnmarshalJSON(data []byte) error { +func (s *CustomToolInputSchema) UnmarshalJSON(data []byte) error { d := jx.DecodeBytes(data) return s.Decode(d) } // Encode implements json.Marshaler. -func (s CreateAgentRequestMetadata) Encode(e *jx.Encoder) { +func (s CustomToolInputSchemaProperties) Encode(e *jx.Encoder) { e.ObjStart() s.encodeFields(e) e.ObjEnd() } // encodeFields implements json.Marshaler. -func (s CreateAgentRequestMetadata) encodeFields(e *jx.Encoder) { +func (s CustomToolInputSchemaProperties) encodeFields(e *jx.Encoder) { for k, elem := range s { e.FieldStart(k) - e.Str(elem) + if len(elem) != 0 { + e.Raw(elem) + } } } -// Decode decodes CreateAgentRequestMetadata from json. -func (s *CreateAgentRequestMetadata) Decode(d *jx.Decoder) error { +// Decode decodes CustomToolInputSchemaProperties from json. +func (s *CustomToolInputSchemaProperties) Decode(d *jx.Decoder) error { if s == nil { - return errors.New("invalid: unable to decode CreateAgentRequestMetadata to nil") + return errors.New("invalid: unable to decode CustomToolInputSchemaProperties to nil") } m := s.init() if err := d.ObjBytes(func(d *jx.Decoder, k []byte) error { - var elem string + var elem jx.Raw if err := func() error { - v, err := d.Str() - elem = string(v) + v, err := d.RawAppend(nil) + elem = jx.Raw(v) if err != nil { return err } @@ -1110,21 +1499,21 @@ func (s *CreateAgentRequestMetadata) Decode(d *jx.Decoder) error { m[string(k)] = elem return nil }); err != nil { - return errors.Wrap(err, "decode CreateAgentRequestMetadata") + return errors.Wrap(err, "decode CustomToolInputSchemaProperties") } return nil } // MarshalJSON implements stdjson.Marshaler. -func (s CreateAgentRequestMetadata) MarshalJSON() ([]byte, error) { +func (s CustomToolInputSchemaProperties) MarshalJSON() ([]byte, error) { e := jx.Encoder{} s.Encode(&e) return e.Bytes(), nil } // UnmarshalJSON implements stdjson.Unmarshaler. -func (s *CreateAgentRequestMetadata) UnmarshalJSON(data []byte) error { +func (s *CustomToolInputSchemaProperties) UnmarshalJSON(data []byte) error { d := jx.DecodeBytes(data) return s.Decode(d) } @@ -1550,11 +1939,50 @@ func (s *ModelConfig) encodeFields(e *jx.Encoder) { s.Speed.Encode(e) } } + { + if s.TokenLimits.Set { + e.FieldStart("token_limits") + s.TokenLimits.Encode(e) + } + } + { + if s.InputModalities != nil { + e.FieldStart("input_modalities") + e.ArrStart() + for _, elem := range s.InputModalities { + e.Str(elem) + } + e.ArrEnd() + } + } + { + if s.Provider.Set { + e.FieldStart("provider") + s.Provider.Encode(e) + } + } + { + if s.Thinking.Set { + e.FieldStart("thinking") + s.Thinking.Encode(e) + } + } + { + if s.ReasoningEffort.Set { + e.FieldStart("reasoning_effort") + s.ReasoningEffort.Encode(e) + } + } } -var jsonFieldsNameOfModelConfig = [2]string{ +var jsonFieldsNameOfModelConfig = [7]string{ 0: "id", 1: "speed", + 2: "token_limits", + 3: "input_modalities", + 4: "provider", + 5: "thinking", + 6: "reasoning_effort", } // Decode decodes ModelConfig from json. @@ -1588,6 +2016,65 @@ func (s *ModelConfig) Decode(d *jx.Decoder) error { }(); err != nil { return errors.Wrap(err, "decode field \"speed\"") } + case "token_limits": + if err := func() error { + s.TokenLimits.Reset() + if err := s.TokenLimits.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"token_limits\"") + } + case "input_modalities": + if err := func() error { + s.InputModalities = make([]string, 0) + if err := d.Arr(func(d *jx.Decoder) error { + var elem string + v, err := d.Str() + elem = string(v) + if err != nil { + return err + } + s.InputModalities = append(s.InputModalities, elem) + return nil + }); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"input_modalities\"") + } + case "provider": + if err := func() error { + s.Provider.Reset() + if err := s.Provider.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"provider\"") + } + case "thinking": + if err := func() error { + s.Thinking.Reset() + if err := s.Thinking.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"thinking\"") + } + case "reasoning_effort": + if err := func() error { + s.ReasoningEffort.Reset() + if err := s.ReasoningEffort.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"reasoning_effort\"") + } default: return d.Skip() } @@ -1946,6 +2433,73 @@ func (s *OptCreateAgentRequestMetadata) UnmarshalJSON(data []byte) error { return s.Decode(d) } +// Encode encodes CustomToolInputSchema as json. +func (o OptCustomToolInputSchema) Encode(e *jx.Encoder) { + if !o.Set { + return + } + o.Value.Encode(e) +} + +// Decode decodes CustomToolInputSchema from json. +func (o *OptCustomToolInputSchema) Decode(d *jx.Decoder) error { + if o == nil { + return errors.New("invalid: unable to decode OptCustomToolInputSchema to nil") + } + o.Set = true + if err := o.Value.Decode(d); err != nil { + return err + } + return nil +} + +// MarshalJSON implements stdjson.Marshaler. +func (s OptCustomToolInputSchema) MarshalJSON() ([]byte, error) { + e := jx.Encoder{} + s.Encode(&e) + return e.Bytes(), nil +} + +// UnmarshalJSON implements stdjson.Unmarshaler. +func (s *OptCustomToolInputSchema) UnmarshalJSON(data []byte) error { + d := jx.DecodeBytes(data) + return s.Decode(d) +} + +// Encode encodes CustomToolInputSchemaProperties as json. +func (o OptCustomToolInputSchemaProperties) Encode(e *jx.Encoder) { + if !o.Set { + return + } + o.Value.Encode(e) +} + +// Decode decodes CustomToolInputSchemaProperties from json. +func (o *OptCustomToolInputSchemaProperties) Decode(d *jx.Decoder) error { + if o == nil { + return errors.New("invalid: unable to decode OptCustomToolInputSchemaProperties to nil") + } + o.Set = true + o.Value = make(CustomToolInputSchemaProperties) + if err := o.Value.Decode(d); err != nil { + return err + } + return nil +} + +// MarshalJSON implements stdjson.Marshaler. +func (s OptCustomToolInputSchemaProperties) MarshalJSON() ([]byte, error) { + e := jx.Encoder{} + s.Encode(&e) + return e.Bytes(), nil +} + +// UnmarshalJSON implements stdjson.Unmarshaler. +func (s *OptCustomToolInputSchemaProperties) UnmarshalJSON(data []byte) error { + d := jx.DecodeBytes(data) + return s.Decode(d) +} + // Encode encodes int32 as json. func (o OptInt32) Encode(e *jx.Encoder) { if !o.Set { @@ -2183,6 +2737,39 @@ func (s *OptString) UnmarshalJSON(data []byte) error { return s.Decode(d) } +// Encode encodes TokenLimits as json. +func (o OptTokenLimits) Encode(e *jx.Encoder) { + if !o.Set { + return + } + o.Value.Encode(e) +} + +// Decode decodes TokenLimits from json. +func (o *OptTokenLimits) Decode(d *jx.Decoder) error { + if o == nil { + return errors.New("invalid: unable to decode OptTokenLimits to nil") + } + o.Set = true + if err := o.Value.Decode(d); err != nil { + return err + } + return nil +} + +// MarshalJSON implements stdjson.Marshaler. +func (s OptTokenLimits) MarshalJSON() ([]byte, error) { + e := jx.Encoder{} + s.Encode(&e) + return e.Bytes(), nil +} + +// UnmarshalJSON implements stdjson.Unmarshaler. +func (s *OptTokenLimits) UnmarshalJSON(data []byte) error { + d := jx.DecodeBytes(data) + return s.Decode(d) +} + // Encode encodes ToolDefaultConfig as json. func (o OptToolDefaultConfig) Encode(e *jx.Encoder) { if !o.Set { @@ -2539,6 +3126,103 @@ func (s *Tag) UnmarshalJSON(data []byte) error { return s.Decode(d) } +// Encode implements json.Marshaler. +func (s *TokenLimits) Encode(e *jx.Encoder) { + e.ObjStart() + s.encodeFields(e) + e.ObjEnd() +} + +// encodeFields encodes fields. +func (s *TokenLimits) encodeFields(e *jx.Encoder) { + { + if s.ContextWindow.Set { + e.FieldStart("context_window") + s.ContextWindow.Encode(e) + } + } + { + if s.MaxInputTokenLength.Set { + e.FieldStart("max_input_token_length") + s.MaxInputTokenLength.Encode(e) + } + } + { + if s.MaxOutputTokenLength.Set { + e.FieldStart("max_output_token_length") + s.MaxOutputTokenLength.Encode(e) + } + } +} + +var jsonFieldsNameOfTokenLimits = [3]string{ + 0: "context_window", + 1: "max_input_token_length", + 2: "max_output_token_length", +} + +// Decode decodes TokenLimits from json. +func (s *TokenLimits) Decode(d *jx.Decoder) error { + if s == nil { + return errors.New("invalid: unable to decode TokenLimits to nil") + } + + if err := d.ObjBytes(func(d *jx.Decoder, k []byte) error { + switch string(k) { + case "context_window": + if err := func() error { + s.ContextWindow.Reset() + if err := s.ContextWindow.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"context_window\"") + } + case "max_input_token_length": + if err := func() error { + s.MaxInputTokenLength.Reset() + if err := s.MaxInputTokenLength.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"max_input_token_length\"") + } + case "max_output_token_length": + if err := func() error { + s.MaxOutputTokenLength.Reset() + if err := s.MaxOutputTokenLength.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"max_output_token_length\"") + } + default: + return d.Skip() + } + return nil + }); err != nil { + return errors.Wrap(err, "decode TokenLimits") + } + + return nil +} + +// MarshalJSON implements stdjson.Marshaler. +func (s *TokenLimits) MarshalJSON() ([]byte, error) { + e := jx.Encoder{} + s.Encode(&e) + return e.Bytes(), nil +} + +// UnmarshalJSON implements stdjson.Unmarshaler. +func (s *TokenLimits) UnmarshalJSON(data []byte) error { + d := jx.DecodeBytes(data) + return s.Decode(d) +} + // Encode implements json.Marshaler. func (s *ToolConfig) Encode(e *jx.Encoder) { e.ObjStart() @@ -2977,6 +3661,12 @@ func (s *UpdateAgentRequest) encodeFields(e *jx.Encoder) { s.Name.Encode(e) } } + { + if s.DisplayName.Set { + e.FieldStart("display_name") + s.DisplayName.Encode(e) + } + } { if s.Model.Set { e.FieldStart("model") @@ -3037,19 +3727,31 @@ func (s *UpdateAgentRequest) encodeFields(e *jx.Encoder) { s.Metadata.Encode(e) } } + { + if s.Tags != nil { + e.FieldStart("tags") + e.ArrStart() + for _, elem := range s.Tags { + elem.Encode(e) + } + e.ArrEnd() + } + } } -var jsonFieldsNameOfUpdateAgentRequest = [10]string{ - 0: "version", - 1: "name", - 2: "model", - 3: "description", - 4: "system", - 5: "mcp_servers", - 6: "tools", - 7: "skills", - 8: "multiagent", - 9: "metadata", +var jsonFieldsNameOfUpdateAgentRequest = [12]string{ + 0: "version", + 1: "name", + 2: "display_name", + 3: "model", + 4: "description", + 5: "system", + 6: "mcp_servers", + 7: "tools", + 8: "skills", + 9: "multiagent", + 10: "metadata", + 11: "tags", } // Decode decodes UpdateAgentRequest from json. @@ -3083,6 +3785,16 @@ func (s *UpdateAgentRequest) Decode(d *jx.Decoder) error { }(); err != nil { return errors.Wrap(err, "decode field \"name\"") } + case "display_name": + if err := func() error { + s.DisplayName.Reset() + if err := s.DisplayName.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"display_name\"") + } case "model": if err := func() error { s.Model.Reset() @@ -3184,6 +3896,23 @@ func (s *UpdateAgentRequest) Decode(d *jx.Decoder) error { }(); err != nil { return errors.Wrap(err, "decode field \"metadata\"") } + case "tags": + if err := func() error { + s.Tags = make([]Tag, 0) + if err := d.Arr(func(d *jx.Decoder) error { + var elem Tag + if err := elem.Decode(d); err != nil { + return err + } + s.Tags = append(s.Tags, elem) + return nil + }); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"tags\"") + } default: return d.Skip() } diff --git a/arkruntime/model/agent/oas_parameters_gen.go b/arkruntime/model/agent/oas_parameters_gen.go index 2e9f8bb..a565fd2 100644 --- a/arkruntime/model/agent/oas_parameters_gen.go +++ b/arkruntime/model/agent/oas_parameters_gen.go @@ -11,6 +11,10 @@ type AgentsListParams struct { Limit OptInt32 `json:",omitempty,omitzero" query:"limit"` // 数字页码,从 1 开始;首页可留空。. Page OptString `json:",omitempty,omitzero" query:"page"` + // 按名称过滤。. + Name OptString `json:",omitempty,omitzero" query:"name"` + // 按展示名过滤。. + DisplayName OptString `json:",omitempty,omitzero" query:"display_name"` // 创建时间下界(RFC 3339)。. CreatedAtGte OptString `json:",omitempty,omitzero" query:"created_at_gte"` // 创建时间上界(RFC 3339)。. diff --git a/arkruntime/model/agent/oas_schemas_gen.go b/arkruntime/model/agent/oas_schemas_gen.go index 9f418ba..e1f373e 100644 --- a/arkruntime/model/agent/oas_schemas_gen.go +++ b/arkruntime/model/agent/oas_schemas_gen.go @@ -7,6 +7,7 @@ package agent import ( "github.com/go-faster/errors" + "github.com/go-faster/jx" ) // ARK Managed Agents 控制面核心资源 @@ -39,6 +40,8 @@ type Agent struct { Metadata OptAgentMetadata `json:"metadata"` // 资源标签。. Tags []Tag `json:"tags"` + // 展示名。. + DisplayName OptString `json:"display_name"` // RFC 3339 时间。. CreatedAt string `json:"created_at"` // RFC 3339 时间。. @@ -110,6 +113,11 @@ func (s *Agent) GetTags() []Tag { return s.Tags } +// GetDisplayName returns the value of DisplayName. +func (s *Agent) GetDisplayName() OptString { + return s.DisplayName +} + // GetCreatedAt returns the value of CreatedAt. func (s *Agent) GetCreatedAt() string { return s.CreatedAt @@ -185,6 +193,11 @@ func (s *Agent) SetTags(val []Tag) { s.Tags = val } +// SetDisplayName sets the value of DisplayName. +func (s *Agent) SetDisplayName(val OptString) { + s.DisplayName = val +} + // SetCreatedAt sets the value of CreatedAt. func (s *Agent) SetCreatedAt(val string) { s.CreatedAt = val @@ -216,6 +229,22 @@ type AgentRef struct { ID OptString `json:"id"` // 被引用 Agent 的版本号。. Version OptInt32 `json:"version"` + // Session 响应中冻结的成员 Agent 名称。. + Name OptString `json:"name"` + // Session 响应中冻结的成员 Agent 描述。. + Description OptString `json:"description"` + // Session 响应中冻结的成员 Agent 模型配置。. + Model OptModelConfig `json:"model"` + // Session 响应中冻结的成员 Agent system prompt。. + System OptString `json:"system"` + // Session 响应中冻结的成员 Agent 工具配置。. + Tools []ToolItem `json:"tools"` + // Session 响应中冻结的成员 Agent MCP servers。. + McpServers []MCPServer `json:"mcp_servers"` + // Session 响应中冻结的成员 Agent skills。. + Skills []AgentSkillRef `json:"skills"` + // Session 响应中冻结的成员 Agent 展示名。. + DisplayName OptString `json:"display_name"` } // GetType returns the value of Type. @@ -233,6 +262,46 @@ func (s *AgentRef) GetVersion() OptInt32 { return s.Version } +// GetName returns the value of Name. +func (s *AgentRef) GetName() OptString { + return s.Name +} + +// GetDescription returns the value of Description. +func (s *AgentRef) GetDescription() OptString { + return s.Description +} + +// GetModel returns the value of Model. +func (s *AgentRef) GetModel() OptModelConfig { + return s.Model +} + +// GetSystem returns the value of System. +func (s *AgentRef) GetSystem() OptString { + return s.System +} + +// GetTools returns the value of Tools. +func (s *AgentRef) GetTools() []ToolItem { + return s.Tools +} + +// GetMcpServers returns the value of McpServers. +func (s *AgentRef) GetMcpServers() []MCPServer { + return s.McpServers +} + +// GetSkills returns the value of Skills. +func (s *AgentRef) GetSkills() []AgentSkillRef { + return s.Skills +} + +// GetDisplayName returns the value of DisplayName. +func (s *AgentRef) GetDisplayName() OptString { + return s.DisplayName +} + // SetType sets the value of Type. func (s *AgentRef) SetType(val AgentRefType) { s.Type = val @@ -248,6 +317,46 @@ func (s *AgentRef) SetVersion(val OptInt32) { s.Version = val } +// SetName sets the value of Name. +func (s *AgentRef) SetName(val OptString) { + s.Name = val +} + +// SetDescription sets the value of Description. +func (s *AgentRef) SetDescription(val OptString) { + s.Description = val +} + +// SetModel sets the value of Model. +func (s *AgentRef) SetModel(val OptModelConfig) { + s.Model = val +} + +// SetSystem sets the value of System. +func (s *AgentRef) SetSystem(val OptString) { + s.System = val +} + +// SetTools sets the value of Tools. +func (s *AgentRef) SetTools(val []ToolItem) { + s.Tools = val +} + +// SetMcpServers sets the value of McpServers. +func (s *AgentRef) SetMcpServers(val []MCPServer) { + s.McpServers = val +} + +// SetSkills sets the value of Skills. +func (s *AgentRef) SetSkills(val []AgentSkillRef) { + s.Skills = val +} + +// SetDisplayName sets the value of DisplayName. +func (s *AgentRef) SetDisplayName(val OptString) { + s.DisplayName = val +} + // 多 Agent 引用类型。. // Ref: #/components/schemas/AgentRefType type AgentRefType string @@ -300,6 +409,8 @@ type AgentSkillRef struct { SkillID OptString `json:"skill_id"` // Skill 版本号,可选;不传走最新。. Version OptString `json:"version"` + // Session 快照中标识创建时用户选择的是使用最新版本。. + UseLatest OptBool `json:"use_latest"` } // GetType returns the value of Type. @@ -317,6 +428,11 @@ func (s *AgentSkillRef) GetVersion() OptString { return s.Version } +// GetUseLatest returns the value of UseLatest. +func (s *AgentSkillRef) GetUseLatest() OptBool { + return s.UseLatest +} + // SetType sets the value of Type. func (s *AgentSkillRef) SetType(val SkillRefType) { s.Type = val @@ -332,6 +448,11 @@ func (s *AgentSkillRef) SetVersion(val OptString) { s.Version = val } +// SetUseLatest sets the value of UseLatest. +func (s *AgentSkillRef) SetUseLatest(val OptBool) { + s.UseLatest = val +} + // 固定 `"agent"`。. type AgentType string @@ -376,6 +497,8 @@ type CreateAgentRequest struct { Model ModelConfig `json:"model"` // 描述信息。. Description OptString `json:"description"` + // 展示名。. + DisplayName OptString `json:"display_name"` // System prompt。. System OptString `json:"system"` // MCP 服务器列表。. @@ -407,6 +530,11 @@ func (s *CreateAgentRequest) GetDescription() OptString { return s.Description } +// GetDisplayName returns the value of DisplayName. +func (s *CreateAgentRequest) GetDisplayName() OptString { + return s.DisplayName +} + // GetSystem returns the value of System. func (s *CreateAgentRequest) GetSystem() OptString { return s.System @@ -457,6 +585,11 @@ func (s *CreateAgentRequest) SetDescription(val OptString) { s.Description = val } +// SetDisplayName sets the value of DisplayName. +func (s *CreateAgentRequest) SetDisplayName(val OptString) { + s.DisplayName = val +} + // SetSystem sets the value of System. func (s *CreateAgentRequest) SetSystem(val OptString) { s.System = val @@ -504,6 +637,59 @@ func (s *CreateAgentRequestMetadata) init() CreateAgentRequestMetadata { return m } +// Custom tool 的输入 JSON Schema。. +// Ref: #/components/schemas/CustomToolInputSchema +type CustomToolInputSchema struct { + // JSON Schema 顶层类型;缺省按 `"object"` 处理。. + Type OptString `json:"type"` + // JSON Schema properties 对象。. + Properties OptCustomToolInputSchemaProperties `json:"properties"` + // 必填属性名列表。. + Required []string `json:"required"` +} + +// GetType returns the value of Type. +func (s *CustomToolInputSchema) GetType() OptString { + return s.Type +} + +// GetProperties returns the value of Properties. +func (s *CustomToolInputSchema) GetProperties() OptCustomToolInputSchemaProperties { + return s.Properties +} + +// GetRequired returns the value of Required. +func (s *CustomToolInputSchema) GetRequired() []string { + return s.Required +} + +// SetType sets the value of Type. +func (s *CustomToolInputSchema) SetType(val OptString) { + s.Type = val +} + +// SetProperties sets the value of Properties. +func (s *CustomToolInputSchema) SetProperties(val OptCustomToolInputSchemaProperties) { + s.Properties = val +} + +// SetRequired sets the value of Required. +func (s *CustomToolInputSchema) SetRequired(val []string) { + s.Required = val +} + +// JSON Schema properties 对象。. +type CustomToolInputSchemaProperties map[string]jx.Raw + +func (s *CustomToolInputSchemaProperties) init() CustomToolInputSchemaProperties { + m := *s + if m == nil { + m = map[string]jx.Raw{} + *s = m + } + return m +} + // Delete Agent 响应体。. // Ref: #/components/schemas/DeleteAgentResponse type DeleteAgentResponse struct { @@ -646,6 +832,16 @@ type ModelConfig struct { ID string `json:"id"` // 速度档位。空字符串走默认(并非所有模型都支持 `fast`)。. Speed OptModelSpeed `json:"speed"` + // 模型 token 限制快照。. + TokenLimits OptTokenLimits `json:"token_limits"` + // 底模支持的输入模态列表。. + InputModalities []string `json:"input_modalities"` + // 模型提供方。. + Provider OptString `json:"provider"` + // Thinking 配置。. + Thinking OptString `json:"thinking"` + // 推理努力程度。. + ReasoningEffort OptString `json:"reasoning_effort"` } // GetID returns the value of ID. @@ -658,6 +854,31 @@ func (s *ModelConfig) GetSpeed() OptModelSpeed { return s.Speed } +// GetTokenLimits returns the value of TokenLimits. +func (s *ModelConfig) GetTokenLimits() OptTokenLimits { + return s.TokenLimits +} + +// GetInputModalities returns the value of InputModalities. +func (s *ModelConfig) GetInputModalities() []string { + return s.InputModalities +} + +// GetProvider returns the value of Provider. +func (s *ModelConfig) GetProvider() OptString { + return s.Provider +} + +// GetThinking returns the value of Thinking. +func (s *ModelConfig) GetThinking() OptString { + return s.Thinking +} + +// GetReasoningEffort returns the value of ReasoningEffort. +func (s *ModelConfig) GetReasoningEffort() OptString { + return s.ReasoningEffort +} + // SetID sets the value of ID. func (s *ModelConfig) SetID(val string) { s.ID = val @@ -668,6 +889,31 @@ func (s *ModelConfig) SetSpeed(val OptModelSpeed) { s.Speed = val } +// SetTokenLimits sets the value of TokenLimits. +func (s *ModelConfig) SetTokenLimits(val OptTokenLimits) { + s.TokenLimits = val +} + +// SetInputModalities sets the value of InputModalities. +func (s *ModelConfig) SetInputModalities(val []string) { + s.InputModalities = val +} + +// SetProvider sets the value of Provider. +func (s *ModelConfig) SetProvider(val OptString) { + s.Provider = val +} + +// SetThinking sets the value of Thinking. +func (s *ModelConfig) SetThinking(val OptString) { + s.Thinking = val +} + +// SetReasoningEffort sets the value of ReasoningEffort. +func (s *ModelConfig) SetReasoningEffort(val OptString) { + s.ReasoningEffort = val +} + // 模型速度档位。空字符串走默认。. // Ref: #/components/schemas/ModelSpeed type ModelSpeed string @@ -915,6 +1161,98 @@ func (o OptCreateAgentRequestMetadata) Or(d CreateAgentRequestMetadata) CreateAg return d } +// NewOptCustomToolInputSchema returns new OptCustomToolInputSchema with value set to v. +func NewOptCustomToolInputSchema(v CustomToolInputSchema) OptCustomToolInputSchema { + return OptCustomToolInputSchema{ + Value: v, + Set: true, + } +} + +// OptCustomToolInputSchema is optional CustomToolInputSchema. +type OptCustomToolInputSchema struct { + Value CustomToolInputSchema + Set bool +} + +// IsSet returns true if OptCustomToolInputSchema was set. +func (o OptCustomToolInputSchema) IsSet() bool { return o.Set } + +// Reset unsets value. +func (o *OptCustomToolInputSchema) Reset() { + var v CustomToolInputSchema + o.Value = v + o.Set = false +} + +// SetTo sets value to v. +func (o *OptCustomToolInputSchema) SetTo(v CustomToolInputSchema) { + o.Set = true + o.Value = v +} + +// Get returns value and boolean that denotes whether value was set. +func (o OptCustomToolInputSchema) Get() (v CustomToolInputSchema, ok bool) { + if !o.Set { + return v, false + } + return o.Value, true +} + +// Or returns value if set, or given parameter if does not. +func (o OptCustomToolInputSchema) Or(d CustomToolInputSchema) CustomToolInputSchema { + if v, ok := o.Get(); ok { + return v + } + return d +} + +// NewOptCustomToolInputSchemaProperties returns new OptCustomToolInputSchemaProperties with value set to v. +func NewOptCustomToolInputSchemaProperties(v CustomToolInputSchemaProperties) OptCustomToolInputSchemaProperties { + return OptCustomToolInputSchemaProperties{ + Value: v, + Set: true, + } +} + +// OptCustomToolInputSchemaProperties is optional CustomToolInputSchemaProperties. +type OptCustomToolInputSchemaProperties struct { + Value CustomToolInputSchemaProperties + Set bool +} + +// IsSet returns true if OptCustomToolInputSchemaProperties was set. +func (o OptCustomToolInputSchemaProperties) IsSet() bool { return o.Set } + +// Reset unsets value. +func (o *OptCustomToolInputSchemaProperties) Reset() { + var v CustomToolInputSchemaProperties + o.Value = v + o.Set = false +} + +// SetTo sets value to v. +func (o *OptCustomToolInputSchemaProperties) SetTo(v CustomToolInputSchemaProperties) { + o.Set = true + o.Value = v +} + +// Get returns value and boolean that denotes whether value was set. +func (o OptCustomToolInputSchemaProperties) Get() (v CustomToolInputSchemaProperties, ok bool) { + if !o.Set { + return v, false + } + return o.Value, true +} + +// Or returns value if set, or given parameter if does not. +func (o OptCustomToolInputSchemaProperties) Or(d CustomToolInputSchemaProperties) CustomToolInputSchemaProperties { + if v, ok := o.Get(); ok { + return v + } + return d +} + // NewOptInt32 returns new OptInt32 with value set to v. func NewOptInt32(v int32) OptInt32 { return OptInt32{ @@ -1237,6 +1575,52 @@ func (o OptString) Or(d string) string { return d } +// NewOptTokenLimits returns new OptTokenLimits with value set to v. +func NewOptTokenLimits(v TokenLimits) OptTokenLimits { + return OptTokenLimits{ + Value: v, + Set: true, + } +} + +// OptTokenLimits is optional TokenLimits. +type OptTokenLimits struct { + Value TokenLimits + Set bool +} + +// IsSet returns true if OptTokenLimits was set. +func (o OptTokenLimits) IsSet() bool { return o.Set } + +// Reset unsets value. +func (o *OptTokenLimits) Reset() { + var v TokenLimits + o.Value = v + o.Set = false +} + +// SetTo sets value to v. +func (o *OptTokenLimits) SetTo(v TokenLimits) { + o.Set = true + o.Value = v +} + +// Get returns value and boolean that denotes whether value was set. +func (o OptTokenLimits) Get() (v TokenLimits, ok bool) { + if !o.Set { + return v, false + } + return o.Value, true +} + +// Or returns value if set, or given parameter if does not. +func (o OptTokenLimits) Or(d TokenLimits) TokenLimits { + if v, ok := o.Get(); ok { + return v + } + return d +} + // NewOptToolDefaultConfig returns new OptToolDefaultConfig with value set to v. func NewOptToolDefaultConfig(v ToolDefaultConfig) OptToolDefaultConfig { return OptToolDefaultConfig{ @@ -1468,6 +1852,47 @@ func (s *Tag) SetValue(val OptString) { s.Value = val } +// 模型上下文与输入输出 token 限制。. +// Ref: #/components/schemas/TokenLimits +type TokenLimits struct { + // 模型上下文窗口。. + ContextWindow OptInt64 `json:"context_window"` + // 最大输入 token 长度。. + MaxInputTokenLength OptInt64 `json:"max_input_token_length"` + // 最大输出 token 长度。. + MaxOutputTokenLength OptInt64 `json:"max_output_token_length"` +} + +// GetContextWindow returns the value of ContextWindow. +func (s *TokenLimits) GetContextWindow() OptInt64 { + return s.ContextWindow +} + +// GetMaxInputTokenLength returns the value of MaxInputTokenLength. +func (s *TokenLimits) GetMaxInputTokenLength() OptInt64 { + return s.MaxInputTokenLength +} + +// GetMaxOutputTokenLength returns the value of MaxOutputTokenLength. +func (s *TokenLimits) GetMaxOutputTokenLength() OptInt64 { + return s.MaxOutputTokenLength +} + +// SetContextWindow sets the value of ContextWindow. +func (s *TokenLimits) SetContextWindow(val OptInt64) { + s.ContextWindow = val +} + +// SetMaxInputTokenLength sets the value of MaxInputTokenLength. +func (s *TokenLimits) SetMaxInputTokenLength(val OptInt64) { + s.MaxInputTokenLength = val +} + +// SetMaxOutputTokenLength sets the value of MaxOutputTokenLength. +func (s *TokenLimits) SetMaxOutputTokenLength(val OptInt64) { + s.MaxOutputTokenLength = val +} + // Toolset 类型工具的单条配置项。. // Ref: #/components/schemas/ToolConfig type ToolConfig struct { @@ -1542,6 +1967,7 @@ func (s *ToolDefaultConfig) SetPermissionPolicy(val OptPermissionPolicy) { // - `agent_toolset_`:内置工具集(当前默认 `agent_toolset_20260701`; // 存量 `agent_toolset_20260401` 仍兼容) // - `mcp_toolset`:来自 `mcp_servers[]` 的工具集 +// - `evolution`:自演进类工具 // - `custom`:客户端执行的自定义工具 // 所有变体字段合并在一个 model 里,未使用的字段留空即可(proto oneof // 风格;wire 上就是同一个 JSON 对象按 `type` 决定语义)。. @@ -1559,9 +1985,8 @@ type ToolItem struct { Name OptString `json:"name"` // `custom` 专用;1–1024 字符。. Description OptString `json:"description"` - // `custom` 专用;承载 JSON Schema 的字符串形态 - // (wire 上是 JSON-encoded string,非 nested object)。. - InputSchema OptString `json:"input_schema"` + // `custom` 专用;承载 JSON Schema 对象。. + InputSchema OptCustomToolInputSchema `json:"input_schema"` } // GetType returns the value of Type. @@ -1595,7 +2020,7 @@ func (s *ToolItem) GetDescription() OptString { } // GetInputSchema returns the value of InputSchema. -func (s *ToolItem) GetInputSchema() OptString { +func (s *ToolItem) GetInputSchema() OptCustomToolInputSchema { return s.InputSchema } @@ -1630,7 +2055,7 @@ func (s *ToolItem) SetDescription(val OptString) { } // SetInputSchema sets the value of InputSchema. -func (s *ToolItem) SetInputSchema(val OptString) { +func (s *ToolItem) SetInputSchema(val OptCustomToolInputSchema) { s.InputSchema = val } @@ -1647,6 +2072,8 @@ type UpdateAgentRequest struct { Version int32 `json:"version"` // 人类可读名称。. Name OptString `json:"name"` + // 展示名。. + DisplayName OptString `json:"display_name"` // 模型配置。. Model OptModelConfig `json:"model"` // 描述信息。. @@ -1663,6 +2090,8 @@ type UpdateAgentRequest struct { Multiagent OptMultiagentConfig `json:"multiagent"` // 用户自定义键值对元数据(patch)。. Metadata OptUpdateAgentRequestMetadata `json:"metadata"` + // 资源标签(整体替换)。. + Tags []Tag `json:"tags"` } // GetVersion returns the value of Version. @@ -1675,6 +2104,11 @@ func (s *UpdateAgentRequest) GetName() OptString { return s.Name } +// GetDisplayName returns the value of DisplayName. +func (s *UpdateAgentRequest) GetDisplayName() OptString { + return s.DisplayName +} + // GetModel returns the value of Model. func (s *UpdateAgentRequest) GetModel() OptModelConfig { return s.Model @@ -1715,6 +2149,11 @@ func (s *UpdateAgentRequest) GetMetadata() OptUpdateAgentRequestMetadata { return s.Metadata } +// GetTags returns the value of Tags. +func (s *UpdateAgentRequest) GetTags() []Tag { + return s.Tags +} + // SetVersion sets the value of Version. func (s *UpdateAgentRequest) SetVersion(val int32) { s.Version = val @@ -1725,6 +2164,11 @@ func (s *UpdateAgentRequest) SetName(val OptString) { s.Name = val } +// SetDisplayName sets the value of DisplayName. +func (s *UpdateAgentRequest) SetDisplayName(val OptString) { + s.DisplayName = val +} + // SetModel sets the value of Model. func (s *UpdateAgentRequest) SetModel(val OptModelConfig) { s.Model = val @@ -1765,6 +2209,11 @@ func (s *UpdateAgentRequest) SetMetadata(val OptUpdateAgentRequestMetadata) { s.Metadata = val } +// SetTags sets the value of Tags. +func (s *UpdateAgentRequest) SetTags(val []Tag) { + s.Tags = val +} + // 用户自定义键值对元数据(patch)。. type UpdateAgentRequestMetadata map[string]string diff --git a/arkruntime/model/agent/oas_validators_gen.go b/arkruntime/model/agent/oas_validators_gen.go index 7276781..ae021d6 100644 --- a/arkruntime/model/agent/oas_validators_gen.go +++ b/arkruntime/model/agent/oas_validators_gen.go @@ -163,6 +163,99 @@ func (s *AgentRef) Validate() error { Error: err, }) } + if err := func() error { + if value, ok := s.Model.Get(); ok { + if err := func() error { + if err := value.Validate(); err != nil { + return err + } + return nil + }(); err != nil { + return err + } + } + return nil + }(); err != nil { + failures = append(failures, validate.FieldError{ + Name: "model", + Error: err, + }) + } + if err := func() error { + var failures []validate.FieldError + for i, elem := range s.Tools { + if err := func() error { + if err := elem.Validate(); err != nil { + return err + } + return nil + }(); err != nil { + failures = append(failures, validate.FieldError{ + Name: fmt.Sprintf("[%d]", i), + Error: err, + }) + } + } + if len(failures) > 0 { + return &validate.Error{Fields: failures} + } + return nil + }(); err != nil { + failures = append(failures, validate.FieldError{ + Name: "tools", + Error: err, + }) + } + if err := func() error { + var failures []validate.FieldError + for i, elem := range s.McpServers { + if err := func() error { + if err := elem.Validate(); err != nil { + return err + } + return nil + }(); err != nil { + failures = append(failures, validate.FieldError{ + Name: fmt.Sprintf("[%d]", i), + Error: err, + }) + } + } + if len(failures) > 0 { + return &validate.Error{Fields: failures} + } + return nil + }(); err != nil { + failures = append(failures, validate.FieldError{ + Name: "mcp_servers", + Error: err, + }) + } + if err := func() error { + var failures []validate.FieldError + for i, elem := range s.Skills { + if err := func() error { + if err := elem.Validate(); err != nil { + return err + } + return nil + }(); err != nil { + failures = append(failures, validate.FieldError{ + Name: fmt.Sprintf("[%d]", i), + Error: err, + }) + } + } + if len(failures) > 0 { + return &validate.Error{Fields: failures} + } + return nil + }(); err != nil { + failures = append(failures, validate.FieldError{ + Name: "skills", + Error: err, + }) + } if len(failures) > 0 { return &validate.Error{Fields: failures} } diff --git a/arkruntime/model/environment/oas_json_gen.go b/arkruntime/model/environment/oas_json_gen.go index 2875d98..93d35bc 100644 --- a/arkruntime/model/environment/oas_json_gen.go +++ b/arkruntime/model/environment/oas_json_gen.go @@ -378,13 +378,27 @@ func (s *EnvConfig) encodeFields(e *jx.Encoder) { s.Env.Encode(e) } } + { + if s.SetupScript.Set { + e.FieldStart("setup_script") + s.SetupScript.Encode(e) + } + } + { + if s.Tos.Set { + e.FieldStart("tos") + s.Tos.Encode(e) + } + } } -var jsonFieldsNameOfEnvConfig = [4]string{ +var jsonFieldsNameOfEnvConfig = [6]string{ 0: "type", 1: "networking", 2: "packages", 3: "env", + 4: "setup_script", + 5: "tos", } // Decode decodes EnvConfig from json. @@ -436,6 +450,26 @@ func (s *EnvConfig) Decode(d *jx.Decoder) error { }(); err != nil { return errors.Wrap(err, "decode field \"env\"") } + case "setup_script": + if err := func() error { + s.SetupScript.Reset() + if err := s.SetupScript.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"setup_script\"") + } + case "tos": + if err := func() error { + s.Tos.Reset() + if err := s.Tos.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"tos\"") + } default: return d.Skip() } @@ -641,9 +675,19 @@ func (s *Environment) encodeFields(e *jx.Encoder) { e.FieldStart("updated_at") e.Str(s.UpdatedAt) } + { + if s.OverriddenFields != nil { + e.FieldStart("overridden_fields") + e.ArrStart() + for _, elem := range s.OverriddenFields { + e.Str(elem) + } + e.ArrEnd() + } + } } -var jsonFieldsNameOfEnvironment = [9]string{ +var jsonFieldsNameOfEnvironment = [10]string{ 0: "id", 1: "type", 2: "name", @@ -653,6 +697,7 @@ var jsonFieldsNameOfEnvironment = [9]string{ 6: "scope", 7: "created_at", 8: "updated_at", + 9: "overridden_fields", } // Decode decodes Environment from json. @@ -762,6 +807,25 @@ func (s *Environment) Decode(d *jx.Decoder) error { }(); err != nil { return errors.Wrap(err, "decode field \"updated_at\"") } + case "overridden_fields": + if err := func() error { + s.OverriddenFields = make([]string, 0) + if err := d.Arr(func(d *jx.Decoder) error { + var elem string + v, err := d.Str() + elem = string(v) + if err != nil { + return err + } + s.OverriddenFields = append(s.OverriddenFields, elem) + return nil + }); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"overridden_fields\"") + } default: return d.Skip() } @@ -953,6 +1017,204 @@ func (s *EnvironmentType) UnmarshalJSON(data []byte) error { return s.Decode(d) } +// Encode implements json.Marshaler. +func (s *HeartbeatWorkResponse) Encode(e *jx.Encoder) { + e.ObjStart() + s.encodeFields(e) + e.ObjEnd() +} + +// encodeFields encodes fields. +func (s *HeartbeatWorkResponse) encodeFields(e *jx.Encoder) { + { + e.FieldStart("last_heartbeat") + e.Str(s.LastHeartbeat) + } + { + e.FieldStart("lease_extended") + e.Bool(s.LeaseExtended) + } + { + e.FieldStart("state") + s.State.Encode(e) + } + { + e.FieldStart("ttl_seconds") + e.Int64(s.TTLSeconds) + } + { + e.FieldStart("type") + s.Type.Encode(e) + } +} + +var jsonFieldsNameOfHeartbeatWorkResponse = [5]string{ + 0: "last_heartbeat", + 1: "lease_extended", + 2: "state", + 3: "ttl_seconds", + 4: "type", +} + +// Decode decodes HeartbeatWorkResponse from json. +func (s *HeartbeatWorkResponse) Decode(d *jx.Decoder) error { + if s == nil { + return errors.New("invalid: unable to decode HeartbeatWorkResponse to nil") + } + var requiredBitSet [1]uint8 + + if err := d.ObjBytes(func(d *jx.Decoder, k []byte) error { + switch string(k) { + case "last_heartbeat": + requiredBitSet[0] |= 1 << 0 + if err := func() error { + v, err := d.Str() + s.LastHeartbeat = string(v) + if err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"last_heartbeat\"") + } + case "lease_extended": + requiredBitSet[0] |= 1 << 1 + if err := func() error { + v, err := d.Bool() + s.LeaseExtended = bool(v) + if err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"lease_extended\"") + } + case "state": + requiredBitSet[0] |= 1 << 2 + if err := func() error { + if err := s.State.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"state\"") + } + case "ttl_seconds": + requiredBitSet[0] |= 1 << 3 + if err := func() error { + v, err := d.Int64() + s.TTLSeconds = int64(v) + if err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"ttl_seconds\"") + } + case "type": + requiredBitSet[0] |= 1 << 4 + if err := func() error { + if err := s.Type.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"type\"") + } + default: + return d.Skip() + } + return nil + }); err != nil { + return errors.Wrap(err, "decode HeartbeatWorkResponse") + } + // Validate required fields. + var failures []validate.FieldError + for i, mask := range [1]uint8{ + 0b00011111, + } { + if result := (requiredBitSet[i] & mask) ^ mask; result != 0 { + // Mask only required fields and check equality to mask using XOR. + // + // If XOR result is not zero, result is not equal to expected, so some fields are missed. + // Bits of fields which would be set are actually bits of missed fields. + missed := bits.OnesCount8(result) + for bitN := 0; bitN < missed; bitN++ { + bitIdx := bits.TrailingZeros8(result) + fieldIdx := i*8 + bitIdx + var name string + if fieldIdx < len(jsonFieldsNameOfHeartbeatWorkResponse) { + name = jsonFieldsNameOfHeartbeatWorkResponse[fieldIdx] + } else { + name = strconv.Itoa(fieldIdx) + } + failures = append(failures, validate.FieldError{ + Name: name, + Error: validate.ErrFieldRequired, + }) + // Reset bit. + result &^= 1 << bitIdx + } + } + } + if len(failures) > 0 { + return &validate.Error{Fields: failures} + } + + return nil +} + +// MarshalJSON implements stdjson.Marshaler. +func (s *HeartbeatWorkResponse) MarshalJSON() ([]byte, error) { + e := jx.Encoder{} + s.Encode(&e) + return e.Bytes(), nil +} + +// UnmarshalJSON implements stdjson.Unmarshaler. +func (s *HeartbeatWorkResponse) UnmarshalJSON(data []byte) error { + d := jx.DecodeBytes(data) + return s.Decode(d) +} + +// Encode encodes HeartbeatWorkResponseType as json. +func (s HeartbeatWorkResponseType) Encode(e *jx.Encoder) { + e.Str(string(s)) +} + +// Decode decodes HeartbeatWorkResponseType from json. +func (s *HeartbeatWorkResponseType) Decode(d *jx.Decoder) error { + if s == nil { + return errors.New("invalid: unable to decode HeartbeatWorkResponseType to nil") + } + v, err := d.StrBytes() + if err != nil { + return err + } + // Try to use constant string. + switch HeartbeatWorkResponseType(v) { + case HeartbeatWorkResponseTypeWorkHeartbeat: + *s = HeartbeatWorkResponseTypeWorkHeartbeat + default: + *s = HeartbeatWorkResponseType(v) + } + + return nil +} + +// MarshalJSON implements stdjson.Marshaler. +func (s HeartbeatWorkResponseType) MarshalJSON() ([]byte, error) { + e := jx.Encoder{} + s.Encode(&e) + return e.Bytes(), nil +} + +// UnmarshalJSON implements stdjson.Unmarshaler. +func (s *HeartbeatWorkResponseType) UnmarshalJSON(data []byte) error { + d := jx.DecodeBytes(data) + return s.Decode(d) +} + // Encode implements json.Marshaler. func (s *ListEnvironmentsResponse) Encode(e *jx.Encoder) { e.ObjStart() @@ -1576,6 +1838,39 @@ func (s *OptPackagesConfigType) UnmarshalJSON(data []byte) error { return s.Decode(d) } +// Encode encodes StopWorkBody as json. +func (o OptStopWorkBody) Encode(e *jx.Encoder) { + if !o.Set { + return + } + o.Value.Encode(e) +} + +// Decode decodes StopWorkBody from json. +func (o *OptStopWorkBody) Decode(d *jx.Decoder) error { + if o == nil { + return errors.New("invalid: unable to decode OptStopWorkBody to nil") + } + o.Set = true + if err := o.Value.Decode(d); err != nil { + return err + } + return nil +} + +// MarshalJSON implements stdjson.Marshaler. +func (s OptStopWorkBody) MarshalJSON() ([]byte, error) { + e := jx.Encoder{} + s.Encode(&e) + return e.Bytes(), nil +} + +// UnmarshalJSON implements stdjson.Unmarshaler. +func (s *OptStopWorkBody) UnmarshalJSON(data []byte) error { + d := jx.DecodeBytes(data) + return s.Decode(d) +} + // Encode encodes string as json. func (o OptString) Encode(e *jx.Encoder) { if !o.Set { @@ -1611,6 +1906,39 @@ func (s *OptString) UnmarshalJSON(data []byte) error { return s.Decode(d) } +// Encode encodes TosConfig as json. +func (o OptTosConfig) Encode(e *jx.Encoder) { + if !o.Set { + return + } + o.Value.Encode(e) +} + +// Decode decodes TosConfig from json. +func (o *OptTosConfig) Decode(d *jx.Decoder) error { + if o == nil { + return errors.New("invalid: unable to decode OptTosConfig to nil") + } + o.Set = true + if err := o.Value.Decode(d); err != nil { + return err + } + return nil +} + +// MarshalJSON implements stdjson.Marshaler. +func (s OptTosConfig) MarshalJSON() ([]byte, error) { + e := jx.Encoder{} + s.Encode(&e) + return e.Bytes(), nil +} + +// UnmarshalJSON implements stdjson.Unmarshaler. +func (s *OptTosConfig) UnmarshalJSON(data []byte) error { + d := jx.DecodeBytes(data) + return s.Decode(d) +} + // Encode encodes UpdateEnvironmentRequestMetadata as json. func (o OptUpdateEnvironmentRequestMetadata) Encode(e *jx.Encoder) { if !o.Set { @@ -1927,29 +2255,172 @@ func (s *PackagesConfigType) UnmarshalJSON(data []byte) error { } // Encode implements json.Marshaler. -func (s *UpdateEnvironmentRequest) Encode(e *jx.Encoder) { +func (s *StopWorkBody) Encode(e *jx.Encoder) { e.ObjStart() s.encodeFields(e) e.ObjEnd() } // encodeFields encodes fields. -func (s *UpdateEnvironmentRequest) encodeFields(e *jx.Encoder) { +func (s *StopWorkBody) encodeFields(e *jx.Encoder) { { - if s.Name.Set { - e.FieldStart("name") - s.Name.Encode(e) + if s.Force.Set { + e.FieldStart("force") + s.Force.Encode(e) } } - { - if s.Config.Set { - e.FieldStart("config") - s.Config.Encode(e) - } +} + +var jsonFieldsNameOfStopWorkBody = [1]string{ + 0: "force", +} + +// Decode decodes StopWorkBody from json. +func (s *StopWorkBody) Decode(d *jx.Decoder) error { + if s == nil { + return errors.New("invalid: unable to decode StopWorkBody to nil") } - { - if s.Description.Set { - e.FieldStart("description") + + if err := d.ObjBytes(func(d *jx.Decoder, k []byte) error { + switch string(k) { + case "force": + if err := func() error { + s.Force.Reset() + if err := s.Force.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"force\"") + } + default: + return d.Skip() + } + return nil + }); err != nil { + return errors.Wrap(err, "decode StopWorkBody") + } + + return nil +} + +// MarshalJSON implements stdjson.Marshaler. +func (s *StopWorkBody) MarshalJSON() ([]byte, error) { + e := jx.Encoder{} + s.Encode(&e) + return e.Bytes(), nil +} + +// UnmarshalJSON implements stdjson.Unmarshaler. +func (s *StopWorkBody) UnmarshalJSON(data []byte) error { + d := jx.DecodeBytes(data) + return s.Decode(d) +} + +// Encode implements json.Marshaler. +func (s *TosConfig) Encode(e *jx.Encoder) { + e.ObjStart() + s.encodeFields(e) + e.ObjEnd() +} + +// encodeFields encodes fields. +func (s *TosConfig) encodeFields(e *jx.Encoder) { + { + if s.Bucket.Set { + e.FieldStart("bucket") + s.Bucket.Encode(e) + } + } + { + if s.Prefix.Set { + e.FieldStart("prefix") + s.Prefix.Encode(e) + } + } +} + +var jsonFieldsNameOfTosConfig = [2]string{ + 0: "bucket", + 1: "prefix", +} + +// Decode decodes TosConfig from json. +func (s *TosConfig) Decode(d *jx.Decoder) error { + if s == nil { + return errors.New("invalid: unable to decode TosConfig to nil") + } + + if err := d.ObjBytes(func(d *jx.Decoder, k []byte) error { + switch string(k) { + case "bucket": + if err := func() error { + s.Bucket.Reset() + if err := s.Bucket.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"bucket\"") + } + case "prefix": + if err := func() error { + s.Prefix.Reset() + if err := s.Prefix.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"prefix\"") + } + default: + return d.Skip() + } + return nil + }); err != nil { + return errors.Wrap(err, "decode TosConfig") + } + + return nil +} + +// MarshalJSON implements stdjson.Marshaler. +func (s *TosConfig) MarshalJSON() ([]byte, error) { + e := jx.Encoder{} + s.Encode(&e) + return e.Bytes(), nil +} + +// UnmarshalJSON implements stdjson.Unmarshaler. +func (s *TosConfig) UnmarshalJSON(data []byte) error { + d := jx.DecodeBytes(data) + return s.Decode(d) +} + +// Encode implements json.Marshaler. +func (s *UpdateEnvironmentRequest) Encode(e *jx.Encoder) { + e.ObjStart() + s.encodeFields(e) + e.ObjEnd() +} + +// encodeFields encodes fields. +func (s *UpdateEnvironmentRequest) encodeFields(e *jx.Encoder) { + { + if s.Name.Set { + e.FieldStart("name") + s.Name.Encode(e) + } + } + { + if s.Config.Set { + e.FieldStart("config") + s.Config.Encode(e) + } + } + { + if s.Description.Set { + e.FieldStart("description") s.Description.Encode(e) } } @@ -2112,3 +2583,619 @@ func (s *UpdateEnvironmentRequestMetadata) UnmarshalJSON(data []byte) error { d := jx.DecodeBytes(data) return s.Decode(d) } + +// Encode implements json.Marshaler. +func (s *VolcTag) Encode(e *jx.Encoder) { + e.ObjStart() + s.encodeFields(e) + e.ObjEnd() +} + +// encodeFields encodes fields. +func (s *VolcTag) encodeFields(e *jx.Encoder) { + { + e.FieldStart("key") + e.Str(s.Key) + } + { + if s.Value.Set { + e.FieldStart("value") + s.Value.Encode(e) + } + } +} + +var jsonFieldsNameOfVolcTag = [2]string{ + 0: "key", + 1: "value", +} + +// Decode decodes VolcTag from json. +func (s *VolcTag) Decode(d *jx.Decoder) error { + if s == nil { + return errors.New("invalid: unable to decode VolcTag to nil") + } + var requiredBitSet [1]uint8 + + if err := d.ObjBytes(func(d *jx.Decoder, k []byte) error { + switch string(k) { + case "key": + requiredBitSet[0] |= 1 << 0 + if err := func() error { + v, err := d.Str() + s.Key = string(v) + if err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"key\"") + } + case "value": + if err := func() error { + s.Value.Reset() + if err := s.Value.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"value\"") + } + default: + return d.Skip() + } + return nil + }); err != nil { + return errors.Wrap(err, "decode VolcTag") + } + // Validate required fields. + var failures []validate.FieldError + for i, mask := range [1]uint8{ + 0b00000001, + } { + if result := (requiredBitSet[i] & mask) ^ mask; result != 0 { + // Mask only required fields and check equality to mask using XOR. + // + // If XOR result is not zero, result is not equal to expected, so some fields are missed. + // Bits of fields which would be set are actually bits of missed fields. + missed := bits.OnesCount8(result) + for bitN := 0; bitN < missed; bitN++ { + bitIdx := bits.TrailingZeros8(result) + fieldIdx := i*8 + bitIdx + var name string + if fieldIdx < len(jsonFieldsNameOfVolcTag) { + name = jsonFieldsNameOfVolcTag[fieldIdx] + } else { + name = strconv.Itoa(fieldIdx) + } + failures = append(failures, validate.FieldError{ + Name: name, + Error: validate.ErrFieldRequired, + }) + // Reset bit. + result &^= 1 << bitIdx + } + } + } + if len(failures) > 0 { + return &validate.Error{Fields: failures} + } + + return nil +} + +// MarshalJSON implements stdjson.Marshaler. +func (s *VolcTag) MarshalJSON() ([]byte, error) { + e := jx.Encoder{} + s.Encode(&e) + return e.Bytes(), nil +} + +// UnmarshalJSON implements stdjson.Unmarshaler. +func (s *VolcTag) UnmarshalJSON(data []byte) error { + d := jx.DecodeBytes(data) + return s.Decode(d) +} + +// Encode implements json.Marshaler. +func (s *WorkData) Encode(e *jx.Encoder) { + e.ObjStart() + s.encodeFields(e) + e.ObjEnd() +} + +// encodeFields encodes fields. +func (s *WorkData) encodeFields(e *jx.Encoder) { + { + e.FieldStart("id") + e.Str(s.ID) + } + { + e.FieldStart("type") + e.Str(s.Type) + } +} + +var jsonFieldsNameOfWorkData = [2]string{ + 0: "id", + 1: "type", +} + +// Decode decodes WorkData from json. +func (s *WorkData) Decode(d *jx.Decoder) error { + if s == nil { + return errors.New("invalid: unable to decode WorkData to nil") + } + var requiredBitSet [1]uint8 + + if err := d.ObjBytes(func(d *jx.Decoder, k []byte) error { + switch string(k) { + case "id": + requiredBitSet[0] |= 1 << 0 + if err := func() error { + v, err := d.Str() + s.ID = string(v) + if err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"id\"") + } + case "type": + requiredBitSet[0] |= 1 << 1 + if err := func() error { + v, err := d.Str() + s.Type = string(v) + if err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"type\"") + } + default: + return d.Skip() + } + return nil + }); err != nil { + return errors.Wrap(err, "decode WorkData") + } + // Validate required fields. + var failures []validate.FieldError + for i, mask := range [1]uint8{ + 0b00000011, + } { + if result := (requiredBitSet[i] & mask) ^ mask; result != 0 { + // Mask only required fields and check equality to mask using XOR. + // + // If XOR result is not zero, result is not equal to expected, so some fields are missed. + // Bits of fields which would be set are actually bits of missed fields. + missed := bits.OnesCount8(result) + for bitN := 0; bitN < missed; bitN++ { + bitIdx := bits.TrailingZeros8(result) + fieldIdx := i*8 + bitIdx + var name string + if fieldIdx < len(jsonFieldsNameOfWorkData) { + name = jsonFieldsNameOfWorkData[fieldIdx] + } else { + name = strconv.Itoa(fieldIdx) + } + failures = append(failures, validate.FieldError{ + Name: name, + Error: validate.ErrFieldRequired, + }) + // Reset bit. + result &^= 1 << bitIdx + } + } + } + if len(failures) > 0 { + return &validate.Error{Fields: failures} + } + + return nil +} + +// MarshalJSON implements stdjson.Marshaler. +func (s *WorkData) MarshalJSON() ([]byte, error) { + e := jx.Encoder{} + s.Encode(&e) + return e.Bytes(), nil +} + +// UnmarshalJSON implements stdjson.Unmarshaler. +func (s *WorkData) UnmarshalJSON(data []byte) error { + d := jx.DecodeBytes(data) + return s.Decode(d) +} + +// Encode implements json.Marshaler. +func (s *WorkItem) Encode(e *jx.Encoder) { + e.ObjStart() + s.encodeFields(e) + e.ObjEnd() +} + +// encodeFields encodes fields. +func (s *WorkItem) encodeFields(e *jx.Encoder) { + { + e.FieldStart("id") + e.Str(s.ID) + } + { + if s.AcknowledgedAt.Set { + e.FieldStart("acknowledged_at") + s.AcknowledgedAt.Encode(e) + } + } + { + e.FieldStart("created_at") + e.Str(s.CreatedAt) + } + { + e.FieldStart("data") + s.Data.Encode(e) + } + { + e.FieldStart("environment_id") + e.Str(s.EnvironmentID) + } + { + if s.LatestHeartbeatAt.Set { + e.FieldStart("latest_heartbeat_at") + s.LatestHeartbeatAt.Encode(e) + } + } + { + if s.Tags != nil { + e.FieldStart("tags") + e.ArrStart() + for _, elem := range s.Tags { + elem.Encode(e) + } + e.ArrEnd() + } + } + { + if s.Secret.Set { + e.FieldStart("secret") + s.Secret.Encode(e) + } + } + { + if s.StartedAt.Set { + e.FieldStart("started_at") + s.StartedAt.Encode(e) + } + } + { + e.FieldStart("state") + s.State.Encode(e) + } + { + if s.StopRequestedAt.Set { + e.FieldStart("stop_requested_at") + s.StopRequestedAt.Encode(e) + } + } + { + if s.StoppedAt.Set { + e.FieldStart("stopped_at") + s.StoppedAt.Encode(e) + } + } + { + e.FieldStart("type") + s.Type.Encode(e) + } +} + +var jsonFieldsNameOfWorkItem = [13]string{ + 0: "id", + 1: "acknowledged_at", + 2: "created_at", + 3: "data", + 4: "environment_id", + 5: "latest_heartbeat_at", + 6: "tags", + 7: "secret", + 8: "started_at", + 9: "state", + 10: "stop_requested_at", + 11: "stopped_at", + 12: "type", +} + +// Decode decodes WorkItem from json. +func (s *WorkItem) Decode(d *jx.Decoder) error { + if s == nil { + return errors.New("invalid: unable to decode WorkItem to nil") + } + var requiredBitSet [2]uint8 + + if err := d.ObjBytes(func(d *jx.Decoder, k []byte) error { + switch string(k) { + case "id": + requiredBitSet[0] |= 1 << 0 + if err := func() error { + v, err := d.Str() + s.ID = string(v) + if err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"id\"") + } + case "acknowledged_at": + if err := func() error { + s.AcknowledgedAt.Reset() + if err := s.AcknowledgedAt.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"acknowledged_at\"") + } + case "created_at": + requiredBitSet[0] |= 1 << 2 + if err := func() error { + v, err := d.Str() + s.CreatedAt = string(v) + if err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"created_at\"") + } + case "data": + requiredBitSet[0] |= 1 << 3 + if err := func() error { + if err := s.Data.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"data\"") + } + case "environment_id": + requiredBitSet[0] |= 1 << 4 + if err := func() error { + v, err := d.Str() + s.EnvironmentID = string(v) + if err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"environment_id\"") + } + case "latest_heartbeat_at": + if err := func() error { + s.LatestHeartbeatAt.Reset() + if err := s.LatestHeartbeatAt.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"latest_heartbeat_at\"") + } + case "tags": + if err := func() error { + s.Tags = make([]VolcTag, 0) + if err := d.Arr(func(d *jx.Decoder) error { + var elem VolcTag + if err := elem.Decode(d); err != nil { + return err + } + s.Tags = append(s.Tags, elem) + return nil + }); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"tags\"") + } + case "secret": + if err := func() error { + s.Secret.Reset() + if err := s.Secret.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"secret\"") + } + case "started_at": + if err := func() error { + s.StartedAt.Reset() + if err := s.StartedAt.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"started_at\"") + } + case "state": + requiredBitSet[1] |= 1 << 1 + if err := func() error { + if err := s.State.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"state\"") + } + case "stop_requested_at": + if err := func() error { + s.StopRequestedAt.Reset() + if err := s.StopRequestedAt.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"stop_requested_at\"") + } + case "stopped_at": + if err := func() error { + s.StoppedAt.Reset() + if err := s.StoppedAt.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"stopped_at\"") + } + case "type": + requiredBitSet[1] |= 1 << 4 + if err := func() error { + if err := s.Type.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"type\"") + } + default: + return d.Skip() + } + return nil + }); err != nil { + return errors.Wrap(err, "decode WorkItem") + } + // Validate required fields. + var failures []validate.FieldError + for i, mask := range [2]uint8{ + 0b00011101, + 0b00010010, + } { + if result := (requiredBitSet[i] & mask) ^ mask; result != 0 { + // Mask only required fields and check equality to mask using XOR. + // + // If XOR result is not zero, result is not equal to expected, so some fields are missed. + // Bits of fields which would be set are actually bits of missed fields. + missed := bits.OnesCount8(result) + for bitN := 0; bitN < missed; bitN++ { + bitIdx := bits.TrailingZeros8(result) + fieldIdx := i*8 + bitIdx + var name string + if fieldIdx < len(jsonFieldsNameOfWorkItem) { + name = jsonFieldsNameOfWorkItem[fieldIdx] + } else { + name = strconv.Itoa(fieldIdx) + } + failures = append(failures, validate.FieldError{ + Name: name, + Error: validate.ErrFieldRequired, + }) + // Reset bit. + result &^= 1 << bitIdx + } + } + } + if len(failures) > 0 { + return &validate.Error{Fields: failures} + } + + return nil +} + +// MarshalJSON implements stdjson.Marshaler. +func (s *WorkItem) MarshalJSON() ([]byte, error) { + e := jx.Encoder{} + s.Encode(&e) + return e.Bytes(), nil +} + +// UnmarshalJSON implements stdjson.Unmarshaler. +func (s *WorkItem) UnmarshalJSON(data []byte) error { + d := jx.DecodeBytes(data) + return s.Decode(d) +} + +// Encode encodes WorkItemType as json. +func (s WorkItemType) Encode(e *jx.Encoder) { + e.Str(string(s)) +} + +// Decode decodes WorkItemType from json. +func (s *WorkItemType) Decode(d *jx.Decoder) error { + if s == nil { + return errors.New("invalid: unable to decode WorkItemType to nil") + } + v, err := d.StrBytes() + if err != nil { + return err + } + // Try to use constant string. + switch WorkItemType(v) { + case WorkItemTypeWork: + *s = WorkItemTypeWork + default: + *s = WorkItemType(v) + } + + return nil +} + +// MarshalJSON implements stdjson.Marshaler. +func (s WorkItemType) MarshalJSON() ([]byte, error) { + e := jx.Encoder{} + s.Encode(&e) + return e.Bytes(), nil +} + +// UnmarshalJSON implements stdjson.Unmarshaler. +func (s *WorkItemType) UnmarshalJSON(data []byte) error { + d := jx.DecodeBytes(data) + return s.Decode(d) +} + +// Encode encodes WorkState as json. +func (s WorkState) Encode(e *jx.Encoder) { + e.Str(string(s)) +} + +// Decode decodes WorkState from json. +func (s *WorkState) Decode(d *jx.Decoder) error { + if s == nil { + return errors.New("invalid: unable to decode WorkState to nil") + } + v, err := d.StrBytes() + if err != nil { + return err + } + // Try to use constant string. + switch WorkState(v) { + case WorkStateQueued: + *s = WorkStateQueued + case WorkStateStarting: + *s = WorkStateStarting + case WorkStateActive: + *s = WorkStateActive + case WorkStateStopping: + *s = WorkStateStopping + case WorkStateStopped: + *s = WorkStateStopped + default: + *s = WorkState(v) + } + + return nil +} + +// MarshalJSON implements stdjson.Marshaler. +func (s WorkState) MarshalJSON() ([]byte, error) { + e := jx.Encoder{} + s.Encode(&e) + return e.Bytes(), nil +} + +// UnmarshalJSON implements stdjson.Unmarshaler. +func (s *WorkState) UnmarshalJSON(data []byte) error { + d := jx.DecodeBytes(data) + return s.Decode(d) +} diff --git a/arkruntime/model/environment/oas_parameters_gen.go b/arkruntime/model/environment/oas_parameters_gen.go index e021b45..98c60bf 100644 --- a/arkruntime/model/environment/oas_parameters_gen.go +++ b/arkruntime/model/environment/oas_parameters_gen.go @@ -5,6 +5,41 @@ package environment +// EnvironmentWorkAckParams is parameters of EnvironmentWork_ack operation. +type EnvironmentWorkAckParams struct { + EnvironmentId string + WorkId string + // Worker 实例 ID,用于控制面记录 work 归属和排障。. + ArkWorkerID OptString `json:",omitempty,omitzero"` +} + +// EnvironmentWorkHeartbeatParams is parameters of EnvironmentWork_heartbeat operation. +type EnvironmentWorkHeartbeatParams struct { + EnvironmentId string + WorkId string + // 期望刷新后的 TTL 秒数。. + DesiredTTLSeconds OptInt64 `json:",omitempty,omitzero" query:"desired_ttl_seconds"` + // Worker 上一次看到的 last_heartbeat,用于并发校验。. + ExpectedLastHeartbeat OptString `json:",omitempty,omitzero" query:"expected_last_heartbeat"` +} + +// EnvironmentWorkPollParams is parameters of EnvironmentWork_poll operation. +type EnvironmentWorkPollParams struct { + EnvironmentId string + // 长轮询阻塞时长,单位毫秒。. + BlockMs OptInt64 `json:",omitempty,omitzero" query:"block_ms"` + // 允许控制面回收超过指定时长未 heartbeat 的 work,单位毫秒。. + ReclaimOlderThanMs OptInt64 `json:",omitempty,omitzero" query:"reclaim_older_than_ms"` + // Worker 实例 ID,用于控制面记录 work 归属和排障。. + ArkWorkerID OptString `json:",omitempty,omitzero"` +} + +// EnvironmentWorkStopParams is parameters of EnvironmentWork_stop operation. +type EnvironmentWorkStopParams struct { + EnvironmentId string + WorkId string +} + // EnvironmentsListParams is parameters of Environments_list operation. type EnvironmentsListParams struct { // 下一页 cursor(opaque)。从上一次响应的 `next_page` 中取。 diff --git a/arkruntime/model/environment/oas_schemas_gen.go b/arkruntime/model/environment/oas_schemas_gen.go index 45b98fb..9aa7f1f 100644 --- a/arkruntime/model/environment/oas_schemas_gen.go +++ b/arkruntime/model/environment/oas_schemas_gen.go @@ -126,6 +126,10 @@ type EnvConfig struct { Packages OptPackagesConfig `json:"packages"` // 容器启动时注入的环境变量。. Env OptEnvConfigEnv `json:"env"` + // 沙箱启动阶段执行的初始化脚本。. + SetupScript OptString `json:"setup_script"` + // Environment outputs 的 TOS 存储配置。. + Tos OptTosConfig `json:"tos"` } // GetType returns the value of Type. @@ -148,6 +152,16 @@ func (s *EnvConfig) GetEnv() OptEnvConfigEnv { return s.Env } +// GetSetupScript returns the value of SetupScript. +func (s *EnvConfig) GetSetupScript() OptString { + return s.SetupScript +} + +// GetTos returns the value of Tos. +func (s *EnvConfig) GetTos() OptTosConfig { + return s.Tos +} + // SetType sets the value of Type. func (s *EnvConfig) SetType(val EnvConfigType) { s.Type = val @@ -168,6 +182,16 @@ func (s *EnvConfig) SetEnv(val OptEnvConfigEnv) { s.Env = val } +// SetSetupScript sets the value of SetupScript. +func (s *EnvConfig) SetSetupScript(val OptString) { + s.SetupScript = val +} + +// SetTos sets the value of Tos. +func (s *EnvConfig) SetTos(val OptTosConfig) { + s.Tos = val +} + // 容器启动时注入的环境变量。. type EnvConfigEnv map[string]string @@ -244,6 +268,8 @@ type Environment struct { CreatedAt string `json:"created_at"` // RFC 3339 时间。. UpdatedAt string `json:"updated_at"` + // Session 使用 EnvironmentWithOverrides 时,本次被覆写的 config 子字段。. + OverriddenFields []string `json:"overridden_fields"` } // GetID returns the value of ID. @@ -291,6 +317,11 @@ func (s *Environment) GetUpdatedAt() string { return s.UpdatedAt } +// GetOverriddenFields returns the value of OverriddenFields. +func (s *Environment) GetOverriddenFields() []string { + return s.OverriddenFields +} + // SetID sets the value of ID. func (s *Environment) SetID(val string) { s.ID = val @@ -336,6 +367,11 @@ func (s *Environment) SetUpdatedAt(val string) { s.UpdatedAt = val } +// SetOverriddenFields sets the value of OverriddenFields. +func (s *Environment) SetOverriddenFields(val []string) { + s.OverriddenFields = val +} + // 用户自定义键值对元数据。. type EnvironmentMetadata map[string]string @@ -426,6 +462,106 @@ func (s *EnvironmentType) UnmarshalText(data []byte) error { } } +// Heartbeat work 的响应体。. +// Ref: #/components/schemas/HeartbeatWorkResponse +type HeartbeatWorkResponse struct { + // 控制面接受的 heartbeat 时间,RFC 3339。. + LastHeartbeat string `json:"last_heartbeat"` + // Lease 是否被刷新。. + LeaseExtended bool `json:"lease_extended"` + // Work 生命周期状态。. + State WorkState `json:"state"` + // Lease TTL 秒数。. + TTLSeconds int64 `json:"ttl_seconds"` + // 对象类型,固定为 `work_heartbeat`。. + Type HeartbeatWorkResponseType `json:"type"` +} + +// GetLastHeartbeat returns the value of LastHeartbeat. +func (s *HeartbeatWorkResponse) GetLastHeartbeat() string { + return s.LastHeartbeat +} + +// GetLeaseExtended returns the value of LeaseExtended. +func (s *HeartbeatWorkResponse) GetLeaseExtended() bool { + return s.LeaseExtended +} + +// GetState returns the value of State. +func (s *HeartbeatWorkResponse) GetState() WorkState { + return s.State +} + +// GetTTLSeconds returns the value of TTLSeconds. +func (s *HeartbeatWorkResponse) GetTTLSeconds() int64 { + return s.TTLSeconds +} + +// GetType returns the value of Type. +func (s *HeartbeatWorkResponse) GetType() HeartbeatWorkResponseType { + return s.Type +} + +// SetLastHeartbeat sets the value of LastHeartbeat. +func (s *HeartbeatWorkResponse) SetLastHeartbeat(val string) { + s.LastHeartbeat = val +} + +// SetLeaseExtended sets the value of LeaseExtended. +func (s *HeartbeatWorkResponse) SetLeaseExtended(val bool) { + s.LeaseExtended = val +} + +// SetState sets the value of State. +func (s *HeartbeatWorkResponse) SetState(val WorkState) { + s.State = val +} + +// SetTTLSeconds sets the value of TTLSeconds. +func (s *HeartbeatWorkResponse) SetTTLSeconds(val int64) { + s.TTLSeconds = val +} + +// SetType sets the value of Type. +func (s *HeartbeatWorkResponse) SetType(val HeartbeatWorkResponseType) { + s.Type = val +} + +// 对象类型,固定为 `work_heartbeat`。. +type HeartbeatWorkResponseType string + +const ( + HeartbeatWorkResponseTypeWorkHeartbeat HeartbeatWorkResponseType = "work_heartbeat" +) + +// AllValues returns all HeartbeatWorkResponseType values. +func (HeartbeatWorkResponseType) AllValues() []HeartbeatWorkResponseType { + return []HeartbeatWorkResponseType{ + HeartbeatWorkResponseTypeWorkHeartbeat, + } +} + +// MarshalText implements encoding.TextMarshaler. +func (s HeartbeatWorkResponseType) MarshalText() ([]byte, error) { + switch s { + case HeartbeatWorkResponseTypeWorkHeartbeat: + return []byte(s), nil + default: + return nil, errors.Errorf("invalid value: %q", s) + } +} + +// UnmarshalText implements encoding.TextUnmarshaler. +func (s *HeartbeatWorkResponseType) UnmarshalText(data []byte) error { + switch HeartbeatWorkResponseType(data) { + case HeartbeatWorkResponseTypeWorkHeartbeat: + *s = HeartbeatWorkResponseTypeWorkHeartbeat + return nil + default: + return errors.Errorf("invalid value: %q", data) + } +} + // List Environments 响应体。. // Ref: #/components/schemas/ListEnvironmentsResponse type ListEnvironmentsResponse struct { @@ -876,6 +1012,52 @@ func (o OptInt32) Or(d int32) int32 { return d } +// NewOptInt64 returns new OptInt64 with value set to v. +func NewOptInt64(v int64) OptInt64 { + return OptInt64{ + Value: v, + Set: true, + } +} + +// OptInt64 is optional int64. +type OptInt64 struct { + Value int64 + Set bool +} + +// IsSet returns true if OptInt64 was set. +func (o OptInt64) IsSet() bool { return o.Set } + +// Reset unsets value. +func (o *OptInt64) Reset() { + var v int64 + o.Value = v + o.Set = false +} + +// SetTo sets value to v. +func (o *OptInt64) SetTo(v int64) { + o.Set = true + o.Value = v +} + +// Get returns value and boolean that denotes whether value was set. +func (o OptInt64) Get() (v int64, ok bool) { + if !o.Set { + return v, false + } + return o.Value, true +} + +// Or returns value if set, or given parameter if does not. +func (o OptInt64) Or(d int64) int64 { + if v, ok := o.Get(); ok { + return v + } + return d +} + // NewOptNetworkingConfig returns new OptNetworkingConfig with value set to v. func NewOptNetworkingConfig(v NetworkingConfig) OptNetworkingConfig { return OptNetworkingConfig{ @@ -1014,6 +1196,52 @@ func (o OptPackagesConfigType) Or(d PackagesConfigType) PackagesConfigType { return d } +// NewOptStopWorkBody returns new OptStopWorkBody with value set to v. +func NewOptStopWorkBody(v StopWorkBody) OptStopWorkBody { + return OptStopWorkBody{ + Value: v, + Set: true, + } +} + +// OptStopWorkBody is optional StopWorkBody. +type OptStopWorkBody struct { + Value StopWorkBody + Set bool +} + +// IsSet returns true if OptStopWorkBody was set. +func (o OptStopWorkBody) IsSet() bool { return o.Set } + +// Reset unsets value. +func (o *OptStopWorkBody) Reset() { + var v StopWorkBody + o.Value = v + o.Set = false +} + +// SetTo sets value to v. +func (o *OptStopWorkBody) SetTo(v StopWorkBody) { + o.Set = true + o.Value = v +} + +// Get returns value and boolean that denotes whether value was set. +func (o OptStopWorkBody) Get() (v StopWorkBody, ok bool) { + if !o.Set { + return v, false + } + return o.Value, true +} + +// Or returns value if set, or given parameter if does not. +func (o OptStopWorkBody) Or(d StopWorkBody) StopWorkBody { + if v, ok := o.Get(); ok { + return v + } + return d +} + // NewOptString returns new OptString with value set to v. func NewOptString(v string) OptString { return OptString{ @@ -1060,6 +1288,52 @@ func (o OptString) Or(d string) string { return d } +// NewOptTosConfig returns new OptTosConfig with value set to v. +func NewOptTosConfig(v TosConfig) OptTosConfig { + return OptTosConfig{ + Value: v, + Set: true, + } +} + +// OptTosConfig is optional TosConfig. +type OptTosConfig struct { + Value TosConfig + Set bool +} + +// IsSet returns true if OptTosConfig was set. +func (o OptTosConfig) IsSet() bool { return o.Set } + +// Reset unsets value. +func (o *OptTosConfig) Reset() { + var v TosConfig + o.Value = v + o.Set = false +} + +// SetTo sets value to v. +func (o *OptTosConfig) SetTo(v TosConfig) { + o.Set = true + o.Value = v +} + +// Get returns value and boolean that denotes whether value was set. +func (o OptTosConfig) Get() (v TosConfig, ok bool) { + if !o.Set { + return v, false + } + return o.Value, true +} + +// Or returns value if set, or given parameter if does not. +func (o OptTosConfig) Or(d TosConfig) TosConfig { + if v, ok := o.Get(); ok { + return v + } + return d +} + // NewOptUpdateEnvironmentRequestMetadata returns new OptUpdateEnvironmentRequestMetadata with value set to v. func NewOptUpdateEnvironmentRequestMetadata(v UpdateEnvironmentRequestMetadata) OptUpdateEnvironmentRequestMetadata { return OptUpdateEnvironmentRequestMetadata{ @@ -1231,6 +1505,53 @@ func (s *PackagesConfigType) UnmarshalText(data []byte) error { } } +// Stop work 的请求体。. +// Ref: #/components/schemas/StopWorkBody +type StopWorkBody struct { + // 是否强制停止。. + Force OptBool `json:"force"` +} + +// GetForce returns the value of Force. +func (s *StopWorkBody) GetForce() OptBool { + return s.Force +} + +// SetForce sets the value of Force. +func (s *StopWorkBody) SetForce(val OptBool) { + s.Force = val +} + +// Environment 产物存储位置。设置后 outputs 文件会注册到用户指定的 TOS +// bucket/prefix;不设置则走方舟默认存储。. +// Ref: #/components/schemas/TosConfig +type TosConfig struct { + // TOS bucket 名称。. + Bucket OptString `json:"bucket"` + // TOS 前缀。. + Prefix OptString `json:"prefix"` +} + +// GetBucket returns the value of Bucket. +func (s *TosConfig) GetBucket() OptString { + return s.Bucket +} + +// GetPrefix returns the value of Prefix. +func (s *TosConfig) GetPrefix() OptString { + return s.Prefix +} + +// SetBucket sets the value of Bucket. +func (s *TosConfig) SetBucket(val OptString) { + s.Bucket = val +} + +// SetPrefix sets the value of Prefix. +func (s *TosConfig) SetPrefix(val OptString) { + s.Prefix = val +} + // 更新 Environment 的请求体。语义: // - 省略字段 = 保留原值 // - `description` 显式空字符串或 null = 清空 @@ -1311,3 +1632,321 @@ func (s *UpdateEnvironmentRequestMetadata) init() UpdateEnvironmentRequestMetada } return m } + +// Work 关联的标签。. +// Ref: #/components/schemas/VolcTag +type VolcTag struct { + // 标签 key。. + Key string `json:"key"` + // 标签 value。. + Value OptString `json:"value"` +} + +// GetKey returns the value of Key. +func (s *VolcTag) GetKey() string { + return s.Key +} + +// GetValue returns the value of Value. +func (s *VolcTag) GetValue() OptString { + return s.Value +} + +// SetKey sets the value of Key. +func (s *VolcTag) SetKey(val string) { + s.Key = val +} + +// SetValue sets the value of Value. +func (s *VolcTag) SetValue(val OptString) { + s.Value = val +} + +// Work 的业务载荷。. +// Ref: #/components/schemas/WorkData +type WorkData struct { + // 业务对象 ID,例如 session ID。. + ID string `json:"id"` + // 业务载荷类型,例如 `session`。. + Type string `json:"type"` +} + +// GetID returns the value of ID. +func (s *WorkData) GetID() string { + return s.ID +} + +// GetType returns the value of Type. +func (s *WorkData) GetType() string { + return s.Type +} + +// SetID sets the value of ID. +func (s *WorkData) SetID(val string) { + s.ID = val +} + +// SetType sets the value of Type. +func (s *WorkData) SetType(val string) { + s.Type = val +} + +// Worker queue 中的一条 work。. +// Ref: #/components/schemas/WorkItem +type WorkItem struct { + // Work ID。. + ID string `json:"id"` + // Work ack 时间,RFC 3339。. + AcknowledgedAt OptString `json:"acknowledged_at"` + // Work 创建时间,RFC 3339。. + CreatedAt string `json:"created_at"` + // 业务载荷。. + Data WorkData `json:"data"` + // Environment ID。. + EnvironmentID string `json:"environment_id"` + // 最近 heartbeat 时间,RFC 3339。. + LatestHeartbeatAt OptString `json:"latest_heartbeat_at"` + // Work 标签。. + Tags []VolcTag `json:"tags"` + // Work secret;仅 poll 时返回,ack / stop 响应会抹掉。. + Secret OptString `json:"secret"` + // Work 开始时间,RFC 3339。. + StartedAt OptString `json:"started_at"` + // Work 生命周期状态。. + State WorkState `json:"state"` + // 控制面请求停止的时间,RFC 3339。. + StopRequestedAt OptString `json:"stop_requested_at"` + // Work 停止时间,RFC 3339。. + StoppedAt OptString `json:"stopped_at"` + // 对象类型,固定为 `work`。. + Type WorkItemType `json:"type"` +} + +// GetID returns the value of ID. +func (s *WorkItem) GetID() string { + return s.ID +} + +// GetAcknowledgedAt returns the value of AcknowledgedAt. +func (s *WorkItem) GetAcknowledgedAt() OptString { + return s.AcknowledgedAt +} + +// GetCreatedAt returns the value of CreatedAt. +func (s *WorkItem) GetCreatedAt() string { + return s.CreatedAt +} + +// GetData returns the value of Data. +func (s *WorkItem) GetData() WorkData { + return s.Data +} + +// GetEnvironmentID returns the value of EnvironmentID. +func (s *WorkItem) GetEnvironmentID() string { + return s.EnvironmentID +} + +// GetLatestHeartbeatAt returns the value of LatestHeartbeatAt. +func (s *WorkItem) GetLatestHeartbeatAt() OptString { + return s.LatestHeartbeatAt +} + +// GetTags returns the value of Tags. +func (s *WorkItem) GetTags() []VolcTag { + return s.Tags +} + +// GetSecret returns the value of Secret. +func (s *WorkItem) GetSecret() OptString { + return s.Secret +} + +// GetStartedAt returns the value of StartedAt. +func (s *WorkItem) GetStartedAt() OptString { + return s.StartedAt +} + +// GetState returns the value of State. +func (s *WorkItem) GetState() WorkState { + return s.State +} + +// GetStopRequestedAt returns the value of StopRequestedAt. +func (s *WorkItem) GetStopRequestedAt() OptString { + return s.StopRequestedAt +} + +// GetStoppedAt returns the value of StoppedAt. +func (s *WorkItem) GetStoppedAt() OptString { + return s.StoppedAt +} + +// GetType returns the value of Type. +func (s *WorkItem) GetType() WorkItemType { + return s.Type +} + +// SetID sets the value of ID. +func (s *WorkItem) SetID(val string) { + s.ID = val +} + +// SetAcknowledgedAt sets the value of AcknowledgedAt. +func (s *WorkItem) SetAcknowledgedAt(val OptString) { + s.AcknowledgedAt = val +} + +// SetCreatedAt sets the value of CreatedAt. +func (s *WorkItem) SetCreatedAt(val string) { + s.CreatedAt = val +} + +// SetData sets the value of Data. +func (s *WorkItem) SetData(val WorkData) { + s.Data = val +} + +// SetEnvironmentID sets the value of EnvironmentID. +func (s *WorkItem) SetEnvironmentID(val string) { + s.EnvironmentID = val +} + +// SetLatestHeartbeatAt sets the value of LatestHeartbeatAt. +func (s *WorkItem) SetLatestHeartbeatAt(val OptString) { + s.LatestHeartbeatAt = val +} + +// SetTags sets the value of Tags. +func (s *WorkItem) SetTags(val []VolcTag) { + s.Tags = val +} + +// SetSecret sets the value of Secret. +func (s *WorkItem) SetSecret(val OptString) { + s.Secret = val +} + +// SetStartedAt sets the value of StartedAt. +func (s *WorkItem) SetStartedAt(val OptString) { + s.StartedAt = val +} + +// SetState sets the value of State. +func (s *WorkItem) SetState(val WorkState) { + s.State = val +} + +// SetStopRequestedAt sets the value of StopRequestedAt. +func (s *WorkItem) SetStopRequestedAt(val OptString) { + s.StopRequestedAt = val +} + +// SetStoppedAt sets the value of StoppedAt. +func (s *WorkItem) SetStoppedAt(val OptString) { + s.StoppedAt = val +} + +// SetType sets the value of Type. +func (s *WorkItem) SetType(val WorkItemType) { + s.Type = val +} + +// 对象类型,固定为 `work`。. +type WorkItemType string + +const ( + WorkItemTypeWork WorkItemType = "work" +) + +// AllValues returns all WorkItemType values. +func (WorkItemType) AllValues() []WorkItemType { + return []WorkItemType{ + WorkItemTypeWork, + } +} + +// MarshalText implements encoding.TextMarshaler. +func (s WorkItemType) MarshalText() ([]byte, error) { + switch s { + case WorkItemTypeWork: + return []byte(s), nil + default: + return nil, errors.Errorf("invalid value: %q", s) + } +} + +// UnmarshalText implements encoding.TextUnmarshaler. +func (s *WorkItemType) UnmarshalText(data []byte) error { + switch WorkItemType(data) { + case WorkItemTypeWork: + *s = WorkItemTypeWork + return nil + default: + return errors.Errorf("invalid value: %q", data) + } +} + +// Work 生命周期状态。. +// Ref: #/components/schemas/WorkState +type WorkState string + +const ( + WorkStateQueued WorkState = "queued" + WorkStateStarting WorkState = "starting" + WorkStateActive WorkState = "active" + WorkStateStopping WorkState = "stopping" + WorkStateStopped WorkState = "stopped" +) + +// AllValues returns all WorkState values. +func (WorkState) AllValues() []WorkState { + return []WorkState{ + WorkStateQueued, + WorkStateStarting, + WorkStateActive, + WorkStateStopping, + WorkStateStopped, + } +} + +// MarshalText implements encoding.TextMarshaler. +func (s WorkState) MarshalText() ([]byte, error) { + switch s { + case WorkStateQueued: + return []byte(s), nil + case WorkStateStarting: + return []byte(s), nil + case WorkStateActive: + return []byte(s), nil + case WorkStateStopping: + return []byte(s), nil + case WorkStateStopped: + return []byte(s), nil + default: + return nil, errors.Errorf("invalid value: %q", s) + } +} + +// UnmarshalText implements encoding.TextUnmarshaler. +func (s *WorkState) UnmarshalText(data []byte) error { + switch WorkState(data) { + case WorkStateQueued: + *s = WorkStateQueued + return nil + case WorkStateStarting: + *s = WorkStateStarting + return nil + case WorkStateActive: + *s = WorkStateActive + return nil + case WorkStateStopping: + *s = WorkStateStopping + return nil + case WorkStateStopped: + *s = WorkStateStopped + return nil + default: + return errors.Errorf("invalid value: %q", data) + } +} diff --git a/arkruntime/model/environment/oas_validators_gen.go b/arkruntime/model/environment/oas_validators_gen.go index c3fef22..1fa0e44 100644 --- a/arkruntime/model/environment/oas_validators_gen.go +++ b/arkruntime/model/environment/oas_validators_gen.go @@ -209,6 +209,49 @@ func (s EnvironmentType) Validate() error { } } +func (s *HeartbeatWorkResponse) Validate() error { + if s == nil { + return validate.ErrNilPointer + } + + var failures []validate.FieldError + if err := func() error { + if err := s.State.Validate(); err != nil { + return err + } + return nil + }(); err != nil { + failures = append(failures, validate.FieldError{ + Name: "state", + Error: err, + }) + } + if err := func() error { + if err := s.Type.Validate(); err != nil { + return err + } + return nil + }(); err != nil { + failures = append(failures, validate.FieldError{ + Name: "type", + Error: err, + }) + } + if len(failures) > 0 { + return &validate.Error{Fields: failures} + } + return nil +} + +func (s HeartbeatWorkResponseType) Validate() error { + switch s { + case "work_heartbeat": + return nil + default: + return errors.Errorf("invalid value: %v", s) + } +} + func (s *ListEnvironmentsResponse) Validate() error { if s == nil { return validate.ErrNilPointer @@ -369,3 +412,63 @@ func (s *UpdateEnvironmentRequest) Validate() error { } return nil } + +func (s *WorkItem) Validate() error { + if s == nil { + return validate.ErrNilPointer + } + + var failures []validate.FieldError + if err := func() error { + if err := s.State.Validate(); err != nil { + return err + } + return nil + }(); err != nil { + failures = append(failures, validate.FieldError{ + Name: "state", + Error: err, + }) + } + if err := func() error { + if err := s.Type.Validate(); err != nil { + return err + } + return nil + }(); err != nil { + failures = append(failures, validate.FieldError{ + Name: "type", + Error: err, + }) + } + if len(failures) > 0 { + return &validate.Error{Fields: failures} + } + return nil +} + +func (s WorkItemType) Validate() error { + switch s { + case "work": + return nil + default: + return errors.Errorf("invalid value: %v", s) + } +} + +func (s WorkState) Validate() error { + switch s { + case "queued": + return nil + case "starting": + return nil + case "active": + return nil + case "stopping": + return nil + case "stopped": + return nil + default: + return errors.Errorf("invalid value: %v", s) + } +} diff --git a/arkruntime/model/environment/work_shim.go b/arkruntime/model/environment/work_shim.go new file mode 100644 index 0000000..77c2108 --- /dev/null +++ b/arkruntime/model/environment/work_shim.go @@ -0,0 +1,108 @@ +// Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +// SPDX-License-Identifier: Apache-2.0 + +package environment + +import ( + "bytes" + + "github.com/volcengine/ark-runtime-go/arkruntime/model" +) + +const ( + // DefaultWorkerClientType is kept for source compatibility. + // + // Deprecated: the MA work API no longer accepts worker client headers. + DefaultWorkerClientType = "ark-self-hosted-worker-go" + // DefaultWorkerClientVersion is kept for source compatibility. + // + // Deprecated: the MA work API no longer accepts worker client headers. + DefaultWorkerClientVersion = "0.1.0" +) + +// PollWorkRequest is the ergonomic SDK request for polling one work item. +type PollWorkRequest struct { + EnvironmentID string `json:"environment_id"` + WorkerID string `json:"worker_id,omitempty"` + // Deprecated: MA poll returns at most one work item. + MaxItems int `json:"max_items,omitempty"` + BlockMS int `json:"block_ms,omitempty"` + ReclaimOlderThanMS int `json:"reclaim_older_than_ms,omitempty"` + // Deprecated: MA work API no longer accepts worker client headers. + WorkerClientType string `json:"-"` + // Deprecated: MA work API no longer accepts worker client headers. + WorkerClientVersion string `json:"-"` +} + +// AckWorkRequest is the ergonomic SDK request for acknowledging one work item. +type AckWorkRequest struct { + EnvironmentID string `json:"environment_id"` + WorkID string `json:"work_id"` + WorkerID OptString `json:"worker_id,omitempty"` +} + +// HeartbeatWorkRequest is the ergonomic SDK request for refreshing one work lease. +type HeartbeatWorkRequest struct { + EnvironmentID string `json:"environment_id"` + WorkID string `json:"work_id"` + ExpectedLastHeartbeat OptString `json:"expected_last_heartbeat,omitempty"` + DesiredTTLSeconds OptInt64 `json:"desired_ttl_seconds,omitempty"` +} + +// StopWorkRequest is the ergonomic SDK request for stopping one work item. +type StopWorkRequest struct { + EnvironmentID string `json:"environment_id"` + WorkID string `json:"work_id"` + Force OptBool `json:"force,omitempty"` +} + +// SessionIDValue returns the session id carried by the work item. +func (w WorkItem) SessionIDValue() string { + if w.Data.ID != "" && (w.Data.Type == "" || w.Data.Type == "session") { + return w.Data.ID + } + return "" +} + +// LeaseTTLSeconds returns the lease TTL carried by legacy MA work payloads. +// +// Deprecated: MA work items no longer carry TTL fields. Use HeartbeatWorkResponse.TTLSeconds. +func (w WorkItem) LeaseTTLSeconds() int { + return 0 +} + +// LatestHeartbeatValue returns the heartbeat CAS value from MA work payloads. +func (w WorkItem) LatestHeartbeatValue() string { + if latestHeartbeatAt, ok := w.LatestHeartbeatAt.Get(); ok && latestHeartbeatAt != "" { + return latestHeartbeatAt + } + return "" +} + +// WorkItemResponse wraps WorkItem so it satisfies model.Response. +type WorkItemResponse struct { + WorkItem + model.HttpHeader +} + +// UnmarshalJSON accepts MA's empty-poll response (`200 {}`) while preserving +// generated validation for real work items. +func (r *WorkItemResponse) UnmarshalJSON(data []byte) error { + trimmed := bytes.TrimSpace(data) + if bytes.Equal(trimmed, []byte("{}")) || bytes.Equal(trimmed, []byte("null")) { + r.WorkItem = WorkItem{} + return nil + } + var item WorkItem + if err := item.UnmarshalJSON(data); err != nil { + return err + } + r.WorkItem = item + return nil +} + +// HeartbeatWorkResponseWrapper wraps HeartbeatWorkResponse. +type HeartbeatWorkResponseWrapper struct { + HeartbeatWorkResponse + model.HttpHeader +} diff --git a/arkruntime/model/session/oas_json_gen.go b/arkruntime/model/session/oas_json_gen.go index 27ce89b..bf4b1ac 100644 --- a/arkruntime/model/session/oas_json_gen.go +++ b/arkruntime/model/session/oas_json_gen.go @@ -73,11 +73,13 @@ func (s *AgentRef) Encode(e *jx.Encoder) { func (s *AgentRef) encodeFields(e *jx.Encoder) { { e.FieldStart("type") - s.Type.Encode(e) + e.Str(s.Type) } { - e.FieldStart("id") - e.Str(s.ID) + if s.ID.Set { + e.FieldStart("id") + s.ID.Encode(e) + } } { if s.Version.Set { @@ -85,12 +87,73 @@ func (s *AgentRef) encodeFields(e *jx.Encoder) { s.Version.Encode(e) } } + { + if s.System.Set { + e.FieldStart("system") + s.System.Encode(e) + } + } + { + if s.Tools != nil { + e.FieldStart("tools") + e.ArrStart() + for _, elem := range s.Tools { + elem.Encode(e) + } + e.ArrEnd() + } + } + { + if s.McpServers != nil { + e.FieldStart("mcp_servers") + e.ArrStart() + for _, elem := range s.McpServers { + elem.Encode(e) + } + e.ArrEnd() + } + } + { + if s.Skills != nil { + e.FieldStart("skills") + e.ArrStart() + for _, elem := range s.Skills { + elem.Encode(e) + } + e.ArrEnd() + } + } + { + if s.Multiagent.Set { + e.FieldStart("multiagent") + s.Multiagent.Encode(e) + } + } + { + if s.DisplayName.Set { + e.FieldStart("display_name") + s.DisplayName.Encode(e) + } + } + { + if s.Model.Set { + e.FieldStart("model") + s.Model.Encode(e) + } + } } -var jsonFieldsNameOfAgentRef = [3]string{ +var jsonFieldsNameOfAgentRef = [10]string{ 0: "type", 1: "id", 2: "version", + 3: "system", + 4: "tools", + 5: "mcp_servers", + 6: "skills", + 7: "multiagent", + 8: "display_name", + 9: "model", } // Decode decodes AgentRef from json. @@ -98,14 +161,16 @@ func (s *AgentRef) Decode(d *jx.Decoder) error { if s == nil { return errors.New("invalid: unable to decode AgentRef to nil") } - var requiredBitSet [1]uint8 + var requiredBitSet [2]uint8 if err := d.ObjBytes(func(d *jx.Decoder, k []byte) error { switch string(k) { case "type": requiredBitSet[0] |= 1 << 0 if err := func() error { - if err := s.Type.Decode(d); err != nil { + v, err := d.Str() + s.Type = string(v) + if err != nil { return err } return nil @@ -113,11 +178,9 @@ func (s *AgentRef) Decode(d *jx.Decoder) error { return errors.Wrap(err, "decode field \"type\"") } case "id": - requiredBitSet[0] |= 1 << 1 if err := func() error { - v, err := d.Str() - s.ID = string(v) - if err != nil { + s.ID.Reset() + if err := s.ID.Decode(d); err != nil { return err } return nil @@ -134,6 +197,97 @@ func (s *AgentRef) Decode(d *jx.Decoder) error { }(); err != nil { return errors.Wrap(err, "decode field \"version\"") } + case "system": + if err := func() error { + s.System.Reset() + if err := s.System.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"system\"") + } + case "tools": + if err := func() error { + s.Tools = make([]AgentRefToolsItem, 0) + if err := d.Arr(func(d *jx.Decoder) error { + var elem AgentRefToolsItem + if err := elem.Decode(d); err != nil { + return err + } + s.Tools = append(s.Tools, elem) + return nil + }); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"tools\"") + } + case "mcp_servers": + if err := func() error { + s.McpServers = make([]AgentRefMcpServersItem, 0) + if err := d.Arr(func(d *jx.Decoder) error { + var elem AgentRefMcpServersItem + if err := elem.Decode(d); err != nil { + return err + } + s.McpServers = append(s.McpServers, elem) + return nil + }); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"mcp_servers\"") + } + case "skills": + if err := func() error { + s.Skills = make([]AgentRefSkillsItem, 0) + if err := d.Arr(func(d *jx.Decoder) error { + var elem AgentRefSkillsItem + if err := elem.Decode(d); err != nil { + return err + } + s.Skills = append(s.Skills, elem) + return nil + }); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"skills\"") + } + case "multiagent": + if err := func() error { + s.Multiagent.Reset() + if err := s.Multiagent.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"multiagent\"") + } + case "display_name": + if err := func() error { + s.DisplayName.Reset() + if err := s.DisplayName.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"display_name\"") + } + case "model": + if err := func() error { + s.Model.Reset() + if err := s.Model.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"model\"") + } default: return d.Skip() } @@ -143,8 +297,9 @@ func (s *AgentRef) Decode(d *jx.Decoder) error { } // Validate required fields. var failures []validate.FieldError - for i, mask := range [1]uint8{ - 0b00000011, + for i, mask := range [2]uint8{ + 0b00000001, + 0b00000000, } { if result := (requiredBitSet[i] & mask) ^ mask; result != 0 { // Mask only required fields and check equality to mask using XOR. @@ -190,165 +345,359 @@ func (s *AgentRef) UnmarshalJSON(data []byte) error { return s.Decode(d) } -// Encode encodes AgentRefType as json. -func (s AgentRefType) Encode(e *jx.Encoder) { - e.Str(string(s)) +// Encode implements json.Marshaler. +func (s AgentRefMcpServersItem) Encode(e *jx.Encoder) { + e.ObjStart() + s.encodeFields(e) + e.ObjEnd() } -// Decode decodes AgentRefType from json. -func (s *AgentRefType) Decode(d *jx.Decoder) error { - if s == nil { - return errors.New("invalid: unable to decode AgentRefType to nil") +// encodeFields implements json.Marshaler. +func (s AgentRefMcpServersItem) encodeFields(e *jx.Encoder) { + for k, elem := range s { + e.FieldStart(k) + + if len(elem) != 0 { + e.Raw(elem) + } } - v, err := d.StrBytes() - if err != nil { - return err +} + +// Decode decodes AgentRefMcpServersItem from json. +func (s *AgentRefMcpServersItem) Decode(d *jx.Decoder) error { + if s == nil { + return errors.New("invalid: unable to decode AgentRefMcpServersItem to nil") } - // Try to use constant string. - switch AgentRefType(v) { - case AgentRefTypeAgent: - *s = AgentRefTypeAgent - default: - *s = AgentRefType(v) + m := s.init() + if err := d.ObjBytes(func(d *jx.Decoder, k []byte) error { + var elem jx.Raw + if err := func() error { + v, err := d.RawAppend(nil) + elem = jx.Raw(v) + if err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrapf(err, "decode field %q", k) + } + m[string(k)] = elem + return nil + }); err != nil { + return errors.Wrap(err, "decode AgentRefMcpServersItem") } return nil } // MarshalJSON implements stdjson.Marshaler. -func (s AgentRefType) MarshalJSON() ([]byte, error) { +func (s AgentRefMcpServersItem) MarshalJSON() ([]byte, error) { e := jx.Encoder{} s.Encode(&e) return e.Bytes(), nil } // UnmarshalJSON implements stdjson.Unmarshaler. -func (s *AgentRefType) UnmarshalJSON(data []byte) error { +func (s *AgentRefMcpServersItem) UnmarshalJSON(data []byte) error { d := jx.DecodeBytes(data) return s.Decode(d) } // Encode implements json.Marshaler. -func (s *Base64DocumentSource) Encode(e *jx.Encoder) { +func (s AgentRefMultiagent) Encode(e *jx.Encoder) { e.ObjStart() s.encodeFields(e) e.ObjEnd() } -// encodeFields encodes fields. -func (s *Base64DocumentSource) encodeFields(e *jx.Encoder) { - { - e.FieldStart("data") - e.Str(s.Data) - } - { - e.FieldStart("media_type") - e.Str(s.MediaType) - } -} +// encodeFields implements json.Marshaler. +func (s AgentRefMultiagent) encodeFields(e *jx.Encoder) { + for k, elem := range s { + e.FieldStart(k) -var jsonFieldsNameOfBase64DocumentSource = [2]string{ - 0: "data", - 1: "media_type", + if len(elem) != 0 { + e.Raw(elem) + } + } } -// Decode decodes Base64DocumentSource from json. -func (s *Base64DocumentSource) Decode(d *jx.Decoder) error { +// Decode decodes AgentRefMultiagent from json. +func (s *AgentRefMultiagent) Decode(d *jx.Decoder) error { if s == nil { - return errors.New("invalid: unable to decode Base64DocumentSource to nil") + return errors.New("invalid: unable to decode AgentRefMultiagent to nil") } - var requiredBitSet [1]uint8 - + m := s.init() if err := d.ObjBytes(func(d *jx.Decoder, k []byte) error { - switch string(k) { - case "data": - requiredBitSet[0] |= 1 << 0 - if err := func() error { - v, err := d.Str() - s.Data = string(v) - if err != nil { - return err - } - return nil - }(); err != nil { - return errors.Wrap(err, "decode field \"data\"") - } - case "media_type": - requiredBitSet[0] |= 1 << 1 - if err := func() error { - v, err := d.Str() - s.MediaType = string(v) - if err != nil { - return err - } - return nil - }(); err != nil { - return errors.Wrap(err, "decode field \"media_type\"") + var elem jx.Raw + if err := func() error { + v, err := d.RawAppend(nil) + elem = jx.Raw(v) + if err != nil { + return err } - default: - return d.Skip() + return nil + }(); err != nil { + return errors.Wrapf(err, "decode field %q", k) } + m[string(k)] = elem return nil }); err != nil { - return errors.Wrap(err, "decode Base64DocumentSource") - } - // Validate required fields. - var failures []validate.FieldError - for i, mask := range [1]uint8{ - 0b00000011, - } { - if result := (requiredBitSet[i] & mask) ^ mask; result != 0 { - // Mask only required fields and check equality to mask using XOR. - // - // If XOR result is not zero, result is not equal to expected, so some fields are missed. - // Bits of fields which would be set are actually bits of missed fields. - missed := bits.OnesCount8(result) - for bitN := 0; bitN < missed; bitN++ { - bitIdx := bits.TrailingZeros8(result) - fieldIdx := i*8 + bitIdx - var name string - if fieldIdx < len(jsonFieldsNameOfBase64DocumentSource) { - name = jsonFieldsNameOfBase64DocumentSource[fieldIdx] - } else { - name = strconv.Itoa(fieldIdx) - } - failures = append(failures, validate.FieldError{ - Name: name, - Error: validate.ErrFieldRequired, - }) - // Reset bit. - result &^= 1 << bitIdx - } - } - } - if len(failures) > 0 { - return &validate.Error{Fields: failures} + return errors.Wrap(err, "decode AgentRefMultiagent") } return nil } // MarshalJSON implements stdjson.Marshaler. -func (s *Base64DocumentSource) MarshalJSON() ([]byte, error) { +func (s AgentRefMultiagent) MarshalJSON() ([]byte, error) { e := jx.Encoder{} s.Encode(&e) return e.Bytes(), nil } // UnmarshalJSON implements stdjson.Unmarshaler. -func (s *Base64DocumentSource) UnmarshalJSON(data []byte) error { +func (s *AgentRefMultiagent) UnmarshalJSON(data []byte) error { d := jx.DecodeBytes(data) return s.Decode(d) } // Encode implements json.Marshaler. -func (s *Base64ImageSource) Encode(e *jx.Encoder) { +func (s AgentRefSkillsItem) Encode(e *jx.Encoder) { e.ObjStart() s.encodeFields(e) e.ObjEnd() } -// encodeFields encodes fields. +// encodeFields implements json.Marshaler. +func (s AgentRefSkillsItem) encodeFields(e *jx.Encoder) { + for k, elem := range s { + e.FieldStart(k) + + if len(elem) != 0 { + e.Raw(elem) + } + } +} + +// Decode decodes AgentRefSkillsItem from json. +func (s *AgentRefSkillsItem) Decode(d *jx.Decoder) error { + if s == nil { + return errors.New("invalid: unable to decode AgentRefSkillsItem to nil") + } + m := s.init() + if err := d.ObjBytes(func(d *jx.Decoder, k []byte) error { + var elem jx.Raw + if err := func() error { + v, err := d.RawAppend(nil) + elem = jx.Raw(v) + if err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrapf(err, "decode field %q", k) + } + m[string(k)] = elem + return nil + }); err != nil { + return errors.Wrap(err, "decode AgentRefSkillsItem") + } + + return nil +} + +// MarshalJSON implements stdjson.Marshaler. +func (s AgentRefSkillsItem) MarshalJSON() ([]byte, error) { + e := jx.Encoder{} + s.Encode(&e) + return e.Bytes(), nil +} + +// UnmarshalJSON implements stdjson.Unmarshaler. +func (s *AgentRefSkillsItem) UnmarshalJSON(data []byte) error { + d := jx.DecodeBytes(data) + return s.Decode(d) +} + +// Encode implements json.Marshaler. +func (s AgentRefToolsItem) Encode(e *jx.Encoder) { + e.ObjStart() + s.encodeFields(e) + e.ObjEnd() +} + +// encodeFields implements json.Marshaler. +func (s AgentRefToolsItem) encodeFields(e *jx.Encoder) { + for k, elem := range s { + e.FieldStart(k) + + if len(elem) != 0 { + e.Raw(elem) + } + } +} + +// Decode decodes AgentRefToolsItem from json. +func (s *AgentRefToolsItem) Decode(d *jx.Decoder) error { + if s == nil { + return errors.New("invalid: unable to decode AgentRefToolsItem to nil") + } + m := s.init() + if err := d.ObjBytes(func(d *jx.Decoder, k []byte) error { + var elem jx.Raw + if err := func() error { + v, err := d.RawAppend(nil) + elem = jx.Raw(v) + if err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrapf(err, "decode field %q", k) + } + m[string(k)] = elem + return nil + }); err != nil { + return errors.Wrap(err, "decode AgentRefToolsItem") + } + + return nil +} + +// MarshalJSON implements stdjson.Marshaler. +func (s AgentRefToolsItem) MarshalJSON() ([]byte, error) { + e := jx.Encoder{} + s.Encode(&e) + return e.Bytes(), nil +} + +// UnmarshalJSON implements stdjson.Unmarshaler. +func (s *AgentRefToolsItem) UnmarshalJSON(data []byte) error { + d := jx.DecodeBytes(data) + return s.Decode(d) +} + +// Encode implements json.Marshaler. +func (s *Base64DocumentSource) Encode(e *jx.Encoder) { + e.ObjStart() + s.encodeFields(e) + e.ObjEnd() +} + +// encodeFields encodes fields. +func (s *Base64DocumentSource) encodeFields(e *jx.Encoder) { + { + e.FieldStart("data") + e.Str(s.Data) + } + { + e.FieldStart("media_type") + e.Str(s.MediaType) + } +} + +var jsonFieldsNameOfBase64DocumentSource = [2]string{ + 0: "data", + 1: "media_type", +} + +// Decode decodes Base64DocumentSource from json. +func (s *Base64DocumentSource) Decode(d *jx.Decoder) error { + if s == nil { + return errors.New("invalid: unable to decode Base64DocumentSource to nil") + } + var requiredBitSet [1]uint8 + + if err := d.ObjBytes(func(d *jx.Decoder, k []byte) error { + switch string(k) { + case "data": + requiredBitSet[0] |= 1 << 0 + if err := func() error { + v, err := d.Str() + s.Data = string(v) + if err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"data\"") + } + case "media_type": + requiredBitSet[0] |= 1 << 1 + if err := func() error { + v, err := d.Str() + s.MediaType = string(v) + if err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"media_type\"") + } + default: + return d.Skip() + } + return nil + }); err != nil { + return errors.Wrap(err, "decode Base64DocumentSource") + } + // Validate required fields. + var failures []validate.FieldError + for i, mask := range [1]uint8{ + 0b00000011, + } { + if result := (requiredBitSet[i] & mask) ^ mask; result != 0 { + // Mask only required fields and check equality to mask using XOR. + // + // If XOR result is not zero, result is not equal to expected, so some fields are missed. + // Bits of fields which would be set are actually bits of missed fields. + missed := bits.OnesCount8(result) + for bitN := 0; bitN < missed; bitN++ { + bitIdx := bits.TrailingZeros8(result) + fieldIdx := i*8 + bitIdx + var name string + if fieldIdx < len(jsonFieldsNameOfBase64DocumentSource) { + name = jsonFieldsNameOfBase64DocumentSource[fieldIdx] + } else { + name = strconv.Itoa(fieldIdx) + } + failures = append(failures, validate.FieldError{ + Name: name, + Error: validate.ErrFieldRequired, + }) + // Reset bit. + result &^= 1 << bitIdx + } + } + } + if len(failures) > 0 { + return &validate.Error{Fields: failures} + } + + return nil +} + +// MarshalJSON implements stdjson.Marshaler. +func (s *Base64DocumentSource) MarshalJSON() ([]byte, error) { + e := jx.Encoder{} + s.Encode(&e) + return e.Bytes(), nil +} + +// UnmarshalJSON implements stdjson.Unmarshaler. +func (s *Base64DocumentSource) UnmarshalJSON(data []byte) error { + d := jx.DecodeBytes(data) + return s.Decode(d) +} + +// Encode implements json.Marshaler. +func (s *Base64ImageSource) Encode(e *jx.Encoder) { + e.ObjStart() + s.encodeFields(e) + e.ObjEnd() +} + +// encodeFields encodes fields. func (s *Base64ImageSource) encodeFields(e *jx.Encoder) { { e.FieldStart("data") @@ -548,8 +897,16 @@ func (s *CreateSessionRequest) encodeFields(e *jx.Encoder) { s.Agent.Encode(e) } { - e.FieldStart("environment_id") - e.Str(s.EnvironmentID) + if s.EnvironmentID.Set { + e.FieldStart("environment_id") + s.EnvironmentID.Encode(e) + } + } + { + if s.Environment.Set { + e.FieldStart("environment") + s.Environment.Encode(e) + } } { if s.Tags != nil { @@ -589,13 +946,14 @@ func (s *CreateSessionRequest) encodeFields(e *jx.Encoder) { } } -var jsonFieldsNameOfCreateSessionRequest = [6]string{ +var jsonFieldsNameOfCreateSessionRequest = [7]string{ 0: "agent", 1: "environment_id", - 2: "tags", - 3: "resources", - 4: "title", - 5: "vault_ids", + 2: "environment", + 3: "tags", + 4: "resources", + 5: "title", + 6: "vault_ids", } // Decode decodes CreateSessionRequest from json. @@ -618,17 +976,25 @@ func (s *CreateSessionRequest) Decode(d *jx.Decoder) error { return errors.Wrap(err, "decode field \"agent\"") } case "environment_id": - requiredBitSet[0] |= 1 << 1 if err := func() error { - v, err := d.Str() - s.EnvironmentID = string(v) - if err != nil { + s.EnvironmentID.Reset() + if err := s.EnvironmentID.Decode(d); err != nil { return err } return nil }(); err != nil { return errors.Wrap(err, "decode field \"environment_id\"") } + case "environment": + if err := func() error { + s.Environment.Reset() + if err := s.Environment.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"environment\"") + } case "tags": if err := func() error { s.Tags = make([]Tag, 0) @@ -702,7 +1068,7 @@ func (s *CreateSessionRequest) Decode(d *jx.Decoder) error { // Validate required fields. var failures []validate.FieldError for i, mask := range [1]uint8{ - 0b00000011, + 0b00000001, } { if result := (requiredBitSet[i] & mask) ^ mask; result != 0 { // Mask only required fields and check equality to mask using XOR. @@ -1205,51 +1571,136 @@ func (s *DocumentSourceSum) UnmarshalJSON(data []byte) error { } // Encode implements json.Marshaler. -func (s *FileDocumentSource) Encode(e *jx.Encoder) { +func (s *EnvironmentConfigOverride) Encode(e *jx.Encoder) { e.ObjStart() s.encodeFields(e) e.ObjEnd() } // encodeFields encodes fields. -func (s *FileDocumentSource) encodeFields(e *jx.Encoder) { +func (s *EnvironmentConfigOverride) encodeFields(e *jx.Encoder) { { - e.FieldStart("file_id") - e.Str(s.FileID) + e.FieldStart("type") + e.Str(s.Type) + } + { + if s.Networking.Set { + e.FieldStart("networking") + s.Networking.Encode(e) + } + } + { + if s.Packages.Set { + e.FieldStart("packages") + s.Packages.Encode(e) + } + } + { + if s.Env.Set { + e.FieldStart("env") + s.Env.Encode(e) + } + } + { + if s.SetupScript.Set { + e.FieldStart("setup_script") + s.SetupScript.Encode(e) + } + } + { + if s.Tos.Set { + e.FieldStart("tos") + s.Tos.Encode(e) + } } } -var jsonFieldsNameOfFileDocumentSource = [1]string{ - 0: "file_id", +var jsonFieldsNameOfEnvironmentConfigOverride = [6]string{ + 0: "type", + 1: "networking", + 2: "packages", + 3: "env", + 4: "setup_script", + 5: "tos", } -// Decode decodes FileDocumentSource from json. -func (s *FileDocumentSource) Decode(d *jx.Decoder) error { +// Decode decodes EnvironmentConfigOverride from json. +func (s *EnvironmentConfigOverride) Decode(d *jx.Decoder) error { if s == nil { - return errors.New("invalid: unable to decode FileDocumentSource to nil") + return errors.New("invalid: unable to decode EnvironmentConfigOverride to nil") } var requiredBitSet [1]uint8 if err := d.ObjBytes(func(d *jx.Decoder, k []byte) error { switch string(k) { - case "file_id": + case "type": requiredBitSet[0] |= 1 << 0 if err := func() error { v, err := d.Str() - s.FileID = string(v) + s.Type = string(v) if err != nil { return err } return nil }(); err != nil { - return errors.Wrap(err, "decode field \"file_id\"") + return errors.Wrap(err, "decode field \"type\"") + } + case "networking": + if err := func() error { + s.Networking.Reset() + if err := s.Networking.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"networking\"") + } + case "packages": + if err := func() error { + s.Packages.Reset() + if err := s.Packages.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"packages\"") + } + case "env": + if err := func() error { + s.Env.Reset() + if err := s.Env.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"env\"") + } + case "setup_script": + if err := func() error { + s.SetupScript.Reset() + if err := s.SetupScript.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"setup_script\"") + } + case "tos": + if err := func() error { + s.Tos.Reset() + if err := s.Tos.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"tos\"") } default: return d.Skip() } return nil }); err != nil { - return errors.Wrap(err, "decode FileDocumentSource") + return errors.Wrap(err, "decode EnvironmentConfigOverride") } // Validate required fields. var failures []validate.FieldError @@ -1266,8 +1717,8 @@ func (s *FileDocumentSource) Decode(d *jx.Decoder) error { bitIdx := bits.TrailingZeros8(result) fieldIdx := i*8 + bitIdx var name string - if fieldIdx < len(jsonFieldsNameOfFileDocumentSource) { - name = jsonFieldsNameOfFileDocumentSource[fieldIdx] + if fieldIdx < len(jsonFieldsNameOfEnvironmentConfigOverride) { + name = jsonFieldsNameOfEnvironmentConfigOverride[fieldIdx] } else { name = strconv.Itoa(fieldIdx) } @@ -1288,64 +1739,184 @@ func (s *FileDocumentSource) Decode(d *jx.Decoder) error { } // MarshalJSON implements stdjson.Marshaler. -func (s *FileDocumentSource) MarshalJSON() ([]byte, error) { +func (s *EnvironmentConfigOverride) MarshalJSON() ([]byte, error) { e := jx.Encoder{} s.Encode(&e) return e.Bytes(), nil } // UnmarshalJSON implements stdjson.Unmarshaler. -func (s *FileDocumentSource) UnmarshalJSON(data []byte) error { +func (s *EnvironmentConfigOverride) UnmarshalJSON(data []byte) error { d := jx.DecodeBytes(data) return s.Decode(d) } // Encode implements json.Marshaler. -func (s *FileImageSource) Encode(e *jx.Encoder) { +func (s EnvironmentConfigOverrideEnv) Encode(e *jx.Encoder) { + e.ObjStart() + s.encodeFields(e) + e.ObjEnd() +} + +// encodeFields implements json.Marshaler. +func (s EnvironmentConfigOverrideEnv) encodeFields(e *jx.Encoder) { + for k, elem := range s { + e.FieldStart(k) + + e.Str(elem) + } +} + +// Decode decodes EnvironmentConfigOverrideEnv from json. +func (s *EnvironmentConfigOverrideEnv) Decode(d *jx.Decoder) error { + if s == nil { + return errors.New("invalid: unable to decode EnvironmentConfigOverrideEnv to nil") + } + m := s.init() + if err := d.ObjBytes(func(d *jx.Decoder, k []byte) error { + var elem string + if err := func() error { + v, err := d.Str() + elem = string(v) + if err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrapf(err, "decode field %q", k) + } + m[string(k)] = elem + return nil + }); err != nil { + return errors.Wrap(err, "decode EnvironmentConfigOverrideEnv") + } + + return nil +} + +// MarshalJSON implements stdjson.Marshaler. +func (s EnvironmentConfigOverrideEnv) MarshalJSON() ([]byte, error) { + e := jx.Encoder{} + s.Encode(&e) + return e.Bytes(), nil +} + +// UnmarshalJSON implements stdjson.Unmarshaler. +func (s *EnvironmentConfigOverrideEnv) UnmarshalJSON(data []byte) error { + d := jx.DecodeBytes(data) + return s.Decode(d) +} + +// Encode implements json.Marshaler. +func (s *EnvironmentNetworkingConfig) Encode(e *jx.Encoder) { e.ObjStart() s.encodeFields(e) e.ObjEnd() } // encodeFields encodes fields. -func (s *FileImageSource) encodeFields(e *jx.Encoder) { +func (s *EnvironmentNetworkingConfig) encodeFields(e *jx.Encoder) { { - e.FieldStart("file_id") - e.Str(s.FileID) + e.FieldStart("type") + e.Str(s.Type) + } + { + if s.AllowMcpServers.Set { + e.FieldStart("allow_mcp_servers") + s.AllowMcpServers.Encode(e) + } + } + { + if s.AllowPackageManagers.Set { + e.FieldStart("allow_package_managers") + s.AllowPackageManagers.Encode(e) + } + } + { + if s.AllowedHosts != nil { + e.FieldStart("allowed_hosts") + e.ArrStart() + for _, elem := range s.AllowedHosts { + e.Str(elem) + } + e.ArrEnd() + } } } -var jsonFieldsNameOfFileImageSource = [1]string{ - 0: "file_id", +var jsonFieldsNameOfEnvironmentNetworkingConfig = [4]string{ + 0: "type", + 1: "allow_mcp_servers", + 2: "allow_package_managers", + 3: "allowed_hosts", } -// Decode decodes FileImageSource from json. -func (s *FileImageSource) Decode(d *jx.Decoder) error { +// Decode decodes EnvironmentNetworkingConfig from json. +func (s *EnvironmentNetworkingConfig) Decode(d *jx.Decoder) error { if s == nil { - return errors.New("invalid: unable to decode FileImageSource to nil") + return errors.New("invalid: unable to decode EnvironmentNetworkingConfig to nil") } var requiredBitSet [1]uint8 if err := d.ObjBytes(func(d *jx.Decoder, k []byte) error { switch string(k) { - case "file_id": + case "type": requiredBitSet[0] |= 1 << 0 if err := func() error { v, err := d.Str() - s.FileID = string(v) + s.Type = string(v) if err != nil { return err } return nil }(); err != nil { - return errors.Wrap(err, "decode field \"file_id\"") + return errors.Wrap(err, "decode field \"type\"") + } + case "allow_mcp_servers": + if err := func() error { + s.AllowMcpServers.Reset() + if err := s.AllowMcpServers.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"allow_mcp_servers\"") + } + case "allow_package_managers": + if err := func() error { + s.AllowPackageManagers.Reset() + if err := s.AllowPackageManagers.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"allow_package_managers\"") + } + case "allowed_hosts": + if err := func() error { + s.AllowedHosts = make([]string, 0) + if err := d.Arr(func(d *jx.Decoder) error { + var elem string + v, err := d.Str() + elem = string(v) + if err != nil { + return err + } + s.AllowedHosts = append(s.AllowedHosts, elem) + return nil + }); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"allowed_hosts\"") } default: return d.Skip() } return nil }); err != nil { - return errors.Wrap(err, "decode FileImageSource") + return errors.Wrap(err, "decode EnvironmentNetworkingConfig") } // Validate required fields. var failures []validate.FieldError @@ -1362,8 +1933,8 @@ func (s *FileImageSource) Decode(d *jx.Decoder) error { bitIdx := bits.TrailingZeros8(result) fieldIdx := i*8 + bitIdx var name string - if fieldIdx < len(jsonFieldsNameOfFileImageSource) { - name = jsonFieldsNameOfFileImageSource[fieldIdx] + if fieldIdx < len(jsonFieldsNameOfEnvironmentNetworkingConfig) { + name = jsonFieldsNameOfEnvironmentNetworkingConfig[fieldIdx] } else { name = strconv.Itoa(fieldIdx) } @@ -1384,428 +1955,591 @@ func (s *FileImageSource) Decode(d *jx.Decoder) error { } // MarshalJSON implements stdjson.Marshaler. -func (s *FileImageSource) MarshalJSON() ([]byte, error) { +func (s *EnvironmentNetworkingConfig) MarshalJSON() ([]byte, error) { e := jx.Encoder{} s.Encode(&e) return e.Bytes(), nil } // UnmarshalJSON implements stdjson.Unmarshaler. -func (s *FileImageSource) UnmarshalJSON(data []byte) error { +func (s *EnvironmentNetworkingConfig) UnmarshalJSON(data []byte) error { d := jx.DecodeBytes(data) return s.Decode(d) } // Encode implements json.Marshaler. -func (s *ImageSource) Encode(e *jx.Encoder) { +func (s *EnvironmentPackagesConfig) Encode(e *jx.Encoder) { e.ObjStart() s.encodeFields(e) e.ObjEnd() } // encodeFields encodes fields. -func (s *ImageSource) encodeFields(e *jx.Encoder) { - s.OneOf.encodeFields(e) -} - -var jsonFieldsNameOfImageSource = [0]string{} - -// Decode decodes ImageSource from json. -func (s *ImageSource) Decode(d *jx.Decoder) error { - if s == nil { - return errors.New("invalid: unable to decode ImageSource to nil") - } - if err := d.Capture(func(d *jx.Decoder) error { - return s.OneOf.Decode(d) - }); err != nil { - return errors.Wrap(err, "decode field OneOf") - } - - if err := d.ObjBytes(func(d *jx.Decoder, k []byte) error { - switch string(k) { - default: - return d.Skip() +func (s *EnvironmentPackagesConfig) encodeFields(e *jx.Encoder) { + { + if s.Type.Set { + e.FieldStart("type") + s.Type.Encode(e) } - }); err != nil { - return errors.Wrap(err, "decode ImageSource") } - - return nil -} - -// MarshalJSON implements stdjson.Marshaler. -func (s *ImageSource) MarshalJSON() ([]byte, error) { - e := jx.Encoder{} - s.Encode(&e) - return e.Bytes(), nil -} - -// UnmarshalJSON implements stdjson.Unmarshaler. -func (s *ImageSource) UnmarshalJSON(data []byte) error { - d := jx.DecodeBytes(data) - return s.Decode(d) -} - -// Encode encodes ImageSourceSum as json. -func (s ImageSourceSum) Encode(e *jx.Encoder) { - e.ObjStart() - s.encodeFields(e) - e.ObjEnd() -} - -func (s ImageSourceSum) encodeFields(e *jx.Encoder) { - switch s.Type { - case Base64ImageSourceImageSourceSum: - e.FieldStart("type") - e.Str("base64") - { - s := s.Base64ImageSource - { - e.FieldStart("data") - e.Str(s.Data) + { + if s.Pip != nil { + e.FieldStart("pip") + e.ArrStart() + for _, elem := range s.Pip { + e.Str(elem) } - { - e.FieldStart("media_type") - e.Str(s.MediaType) + e.ArrEnd() + } + } + { + if s.Apt != nil { + e.FieldStart("apt") + e.ArrStart() + for _, elem := range s.Apt { + e.Str(elem) } + e.ArrEnd() } - case UrlImageSourceImageSourceSum: - e.FieldStart("type") - e.Str("url") - { - s := s.UrlImageSource - { - e.FieldStart("url") - e.Str(s.URL) + } + { + if s.Npm != nil { + e.FieldStart("npm") + e.ArrStart() + for _, elem := range s.Npm { + e.Str(elem) } + e.ArrEnd() } - case FileImageSourceImageSourceSum: - e.FieldStart("type") - e.Str("file") - { - s := s.FileImageSource - { - e.FieldStart("file_id") - e.Str(s.FileID) + } + { + if s.Cargo != nil { + e.FieldStart("cargo") + e.ArrStart() + for _, elem := range s.Cargo { + e.Str(elem) } + e.ArrEnd() } } -} - -// Decode decodes ImageSourceSum from json. -func (s *ImageSourceSum) Decode(d *jx.Decoder) error { - if s == nil { - return errors.New("invalid: unable to decode ImageSourceSum to nil") + { + if s.Gem != nil { + e.FieldStart("gem") + e.ArrStart() + for _, elem := range s.Gem { + e.Str(elem) + } + e.ArrEnd() + } } - // Sum type discriminator. - if typ := d.Next(); typ != jx.Object { - return errors.Errorf("unexpected json type %q", typ) + { + if s.Go != nil { + e.FieldStart("go") + e.ArrStart() + for _, elem := range s.Go { + e.Str(elem) + } + e.ArrEnd() + } } +} - var found bool - if err := d.Capture(func(d *jx.Decoder) error { - return d.ObjBytes(func(d *jx.Decoder, key []byte) error { - if found { - return d.Skip() +var jsonFieldsNameOfEnvironmentPackagesConfig = [7]string{ + 0: "type", + 1: "pip", + 2: "apt", + 3: "npm", + 4: "cargo", + 5: "gem", + 6: "go", +} + +// Decode decodes EnvironmentPackagesConfig from json. +func (s *EnvironmentPackagesConfig) Decode(d *jx.Decoder) error { + if s == nil { + return errors.New("invalid: unable to decode EnvironmentPackagesConfig to nil") + } + + if err := d.ObjBytes(func(d *jx.Decoder, k []byte) error { + switch string(k) { + case "type": + if err := func() error { + s.Type.Reset() + if err := s.Type.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"type\"") } - switch string(key) { - case "type": - typ, err := d.Str() - if err != nil { + case "pip": + if err := func() error { + s.Pip = make([]string, 0) + if err := d.Arr(func(d *jx.Decoder) error { + var elem string + v, err := d.Str() + elem = string(v) + if err != nil { + return err + } + s.Pip = append(s.Pip, elem) + return nil + }); err != nil { return err } - switch typ { - case "base64": - s.Type = Base64ImageSourceImageSourceSum - found = true - case "url": - s.Type = UrlImageSourceImageSourceSum - found = true - case "file": - s.Type = FileImageSourceImageSourceSum - found = true - default: - return errors.Errorf("unknown type %s", typ) + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"pip\"") + } + case "apt": + if err := func() error { + s.Apt = make([]string, 0) + if err := d.Arr(func(d *jx.Decoder) error { + var elem string + v, err := d.Str() + elem = string(v) + if err != nil { + return err + } + s.Apt = append(s.Apt, elem) + return nil + }); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"apt\"") + } + case "npm": + if err := func() error { + s.Npm = make([]string, 0) + if err := d.Arr(func(d *jx.Decoder) error { + var elem string + v, err := d.Str() + elem = string(v) + if err != nil { + return err + } + s.Npm = append(s.Npm, elem) + return nil + }); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"npm\"") + } + case "cargo": + if err := func() error { + s.Cargo = make([]string, 0) + if err := d.Arr(func(d *jx.Decoder) error { + var elem string + v, err := d.Str() + elem = string(v) + if err != nil { + return err + } + s.Cargo = append(s.Cargo, elem) + return nil + }); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"cargo\"") + } + case "gem": + if err := func() error { + s.Gem = make([]string, 0) + if err := d.Arr(func(d *jx.Decoder) error { + var elem string + v, err := d.Str() + elem = string(v) + if err != nil { + return err + } + s.Gem = append(s.Gem, elem) + return nil + }); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"gem\"") + } + case "go": + if err := func() error { + s.Go = make([]string, 0) + if err := d.Arr(func(d *jx.Decoder) error { + var elem string + v, err := d.Str() + elem = string(v) + if err != nil { + return err + } + s.Go = append(s.Go, elem) + return nil + }); err != nil { + return err } return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"go\"") } + default: return d.Skip() - }) + } + return nil }); err != nil { - return errors.Wrap(err, "capture") + return errors.Wrap(err, "decode EnvironmentPackagesConfig") } - if !found { - return errors.New("unable to detect sum type variant") + + return nil +} + +// MarshalJSON implements stdjson.Marshaler. +func (s *EnvironmentPackagesConfig) MarshalJSON() ([]byte, error) { + e := jx.Encoder{} + s.Encode(&e) + return e.Bytes(), nil +} + +// UnmarshalJSON implements stdjson.Unmarshaler. +func (s *EnvironmentPackagesConfig) UnmarshalJSON(data []byte) error { + d := jx.DecodeBytes(data) + return s.Decode(d) +} + +// Encode encodes EnvironmentPackagesConfigType as json. +func (s EnvironmentPackagesConfigType) Encode(e *jx.Encoder) { + e.Str(string(s)) +} + +// Decode decodes EnvironmentPackagesConfigType from json. +func (s *EnvironmentPackagesConfigType) Decode(d *jx.Decoder) error { + if s == nil { + return errors.New("invalid: unable to decode EnvironmentPackagesConfigType to nil") } - switch s.Type { - case Base64ImageSourceImageSourceSum: - if err := s.Base64ImageSource.Decode(d); err != nil { - return err - } - case UrlImageSourceImageSourceSum: - if err := s.UrlImageSource.Decode(d); err != nil { - return err - } - case FileImageSourceImageSourceSum: - if err := s.FileImageSource.Decode(d); err != nil { - return err - } + v, err := d.StrBytes() + if err != nil { + return err + } + // Try to use constant string. + switch EnvironmentPackagesConfigType(v) { + case EnvironmentPackagesConfigTypePackages: + *s = EnvironmentPackagesConfigTypePackages default: - return errors.Errorf("inferred invalid type: %s", s.Type) + *s = EnvironmentPackagesConfigType(v) } + return nil } // MarshalJSON implements stdjson.Marshaler. -func (s ImageSourceSum) MarshalJSON() ([]byte, error) { +func (s EnvironmentPackagesConfigType) MarshalJSON() ([]byte, error) { e := jx.Encoder{} s.Encode(&e) return e.Bytes(), nil } // UnmarshalJSON implements stdjson.Unmarshaler. -func (s *ImageSourceSum) UnmarshalJSON(data []byte) error { +func (s *EnvironmentPackagesConfigType) UnmarshalJSON(data []byte) error { d := jx.DecodeBytes(data) return s.Decode(d) } // Encode implements json.Marshaler. -func (s *ListSessionEventsResponseWire) Encode(e *jx.Encoder) { +func (s *EnvironmentTosConfig) Encode(e *jx.Encoder) { e.ObjStart() s.encodeFields(e) e.ObjEnd() } // encodeFields encodes fields. -func (s *ListSessionEventsResponseWire) encodeFields(e *jx.Encoder) { +func (s *EnvironmentTosConfig) encodeFields(e *jx.Encoder) { { - e.FieldStart("data") - e.ArrStart() - for _, elem := range s.Data { - elem.Encode(e) + if s.Bucket.Set { + e.FieldStart("bucket") + s.Bucket.Encode(e) } - e.ArrEnd() } { - if s.NextPage.Set { - e.FieldStart("next_page") - s.NextPage.Encode(e) + if s.Prefix.Set { + e.FieldStart("prefix") + s.Prefix.Encode(e) } } } -var jsonFieldsNameOfListSessionEventsResponseWire = [2]string{ - 0: "data", - 1: "next_page", +var jsonFieldsNameOfEnvironmentTosConfig = [2]string{ + 0: "bucket", + 1: "prefix", } -// Decode decodes ListSessionEventsResponseWire from json. -func (s *ListSessionEventsResponseWire) Decode(d *jx.Decoder) error { +// Decode decodes EnvironmentTosConfig from json. +func (s *EnvironmentTosConfig) Decode(d *jx.Decoder) error { if s == nil { - return errors.New("invalid: unable to decode ListSessionEventsResponseWire to nil") + return errors.New("invalid: unable to decode EnvironmentTosConfig to nil") } - var requiredBitSet [1]uint8 if err := d.ObjBytes(func(d *jx.Decoder, k []byte) error { switch string(k) { - case "data": - requiredBitSet[0] |= 1 << 0 + case "bucket": if err := func() error { - s.Data = make([]ListSessionEventsResponseWireDataItem, 0) - if err := d.Arr(func(d *jx.Decoder) error { - var elem ListSessionEventsResponseWireDataItem - if err := elem.Decode(d); err != nil { - return err - } - s.Data = append(s.Data, elem) - return nil - }); err != nil { + s.Bucket.Reset() + if err := s.Bucket.Decode(d); err != nil { return err } return nil }(); err != nil { - return errors.Wrap(err, "decode field \"data\"") + return errors.Wrap(err, "decode field \"bucket\"") } - case "next_page": + case "prefix": if err := func() error { - s.NextPage.Reset() - if err := s.NextPage.Decode(d); err != nil { + s.Prefix.Reset() + if err := s.Prefix.Decode(d); err != nil { return err } return nil }(); err != nil { - return errors.Wrap(err, "decode field \"next_page\"") + return errors.Wrap(err, "decode field \"prefix\"") } default: return d.Skip() } return nil }); err != nil { - return errors.Wrap(err, "decode ListSessionEventsResponseWire") - } - // Validate required fields. - var failures []validate.FieldError - for i, mask := range [1]uint8{ - 0b00000001, - } { - if result := (requiredBitSet[i] & mask) ^ mask; result != 0 { - // Mask only required fields and check equality to mask using XOR. - // - // If XOR result is not zero, result is not equal to expected, so some fields are missed. - // Bits of fields which would be set are actually bits of missed fields. - missed := bits.OnesCount8(result) - for bitN := 0; bitN < missed; bitN++ { - bitIdx := bits.TrailingZeros8(result) - fieldIdx := i*8 + bitIdx - var name string - if fieldIdx < len(jsonFieldsNameOfListSessionEventsResponseWire) { - name = jsonFieldsNameOfListSessionEventsResponseWire[fieldIdx] - } else { - name = strconv.Itoa(fieldIdx) - } - failures = append(failures, validate.FieldError{ - Name: name, - Error: validate.ErrFieldRequired, - }) - // Reset bit. - result &^= 1 << bitIdx - } - } - } - if len(failures) > 0 { - return &validate.Error{Fields: failures} + return errors.Wrap(err, "decode EnvironmentTosConfig") } return nil } // MarshalJSON implements stdjson.Marshaler. -func (s *ListSessionEventsResponseWire) MarshalJSON() ([]byte, error) { +func (s *EnvironmentTosConfig) MarshalJSON() ([]byte, error) { e := jx.Encoder{} s.Encode(&e) return e.Bytes(), nil } // UnmarshalJSON implements stdjson.Unmarshaler. -func (s *ListSessionEventsResponseWire) UnmarshalJSON(data []byte) error { +func (s *EnvironmentTosConfig) UnmarshalJSON(data []byte) error { d := jx.DecodeBytes(data) return s.Decode(d) } // Encode implements json.Marshaler. -func (s ListSessionEventsResponseWireDataItem) Encode(e *jx.Encoder) { +func (s *EnvironmentWithOverrides) Encode(e *jx.Encoder) { e.ObjStart() s.encodeFields(e) e.ObjEnd() } -// encodeFields implements json.Marshaler. -func (s ListSessionEventsResponseWireDataItem) encodeFields(e *jx.Encoder) { - for k, elem := range s { - e.FieldStart(k) - - if len(elem) != 0 { - e.Raw(elem) +// encodeFields encodes fields. +func (s *EnvironmentWithOverrides) encodeFields(e *jx.Encoder) { + { + e.FieldStart("type") + s.Type.Encode(e) + } + { + e.FieldStart("id") + e.Str(s.ID) + } + { + if s.Config.Set { + e.FieldStart("config") + s.Config.Encode(e) } } } -// Decode decodes ListSessionEventsResponseWireDataItem from json. -func (s *ListSessionEventsResponseWireDataItem) Decode(d *jx.Decoder) error { +var jsonFieldsNameOfEnvironmentWithOverrides = [3]string{ + 0: "type", + 1: "id", + 2: "config", +} + +// Decode decodes EnvironmentWithOverrides from json. +func (s *EnvironmentWithOverrides) Decode(d *jx.Decoder) error { if s == nil { - return errors.New("invalid: unable to decode ListSessionEventsResponseWireDataItem to nil") + return errors.New("invalid: unable to decode EnvironmentWithOverrides to nil") } - m := s.init() + var requiredBitSet [1]uint8 + if err := d.ObjBytes(func(d *jx.Decoder, k []byte) error { - var elem jx.Raw - if err := func() error { - v, err := d.RawAppend(nil) - elem = jx.Raw(v) - if err != nil { - return err + switch string(k) { + case "type": + requiredBitSet[0] |= 1 << 0 + if err := func() error { + if err := s.Type.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"type\"") } - return nil - }(); err != nil { - return errors.Wrapf(err, "decode field %q", k) + case "id": + requiredBitSet[0] |= 1 << 1 + if err := func() error { + v, err := d.Str() + s.ID = string(v) + if err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"id\"") + } + case "config": + if err := func() error { + s.Config.Reset() + if err := s.Config.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"config\"") + } + default: + return d.Skip() } - m[string(k)] = elem return nil }); err != nil { - return errors.Wrap(err, "decode ListSessionEventsResponseWireDataItem") + return errors.Wrap(err, "decode EnvironmentWithOverrides") + } + // Validate required fields. + var failures []validate.FieldError + for i, mask := range [1]uint8{ + 0b00000011, + } { + if result := (requiredBitSet[i] & mask) ^ mask; result != 0 { + // Mask only required fields and check equality to mask using XOR. + // + // If XOR result is not zero, result is not equal to expected, so some fields are missed. + // Bits of fields which would be set are actually bits of missed fields. + missed := bits.OnesCount8(result) + for bitN := 0; bitN < missed; bitN++ { + bitIdx := bits.TrailingZeros8(result) + fieldIdx := i*8 + bitIdx + var name string + if fieldIdx < len(jsonFieldsNameOfEnvironmentWithOverrides) { + name = jsonFieldsNameOfEnvironmentWithOverrides[fieldIdx] + } else { + name = strconv.Itoa(fieldIdx) + } + failures = append(failures, validate.FieldError{ + Name: name, + Error: validate.ErrFieldRequired, + }) + // Reset bit. + result &^= 1 << bitIdx + } + } + } + if len(failures) > 0 { + return &validate.Error{Fields: failures} } return nil } // MarshalJSON implements stdjson.Marshaler. -func (s ListSessionEventsResponseWireDataItem) MarshalJSON() ([]byte, error) { +func (s *EnvironmentWithOverrides) MarshalJSON() ([]byte, error) { e := jx.Encoder{} s.Encode(&e) return e.Bytes(), nil } // UnmarshalJSON implements stdjson.Unmarshaler. -func (s *ListSessionEventsResponseWireDataItem) UnmarshalJSON(data []byte) error { +func (s *EnvironmentWithOverrides) UnmarshalJSON(data []byte) error { + d := jx.DecodeBytes(data) + return s.Decode(d) +} + +// Encode encodes EnvironmentWithOverridesType as json. +func (s EnvironmentWithOverridesType) Encode(e *jx.Encoder) { + e.Str(string(s)) +} + +// Decode decodes EnvironmentWithOverridesType from json. +func (s *EnvironmentWithOverridesType) Decode(d *jx.Decoder) error { + if s == nil { + return errors.New("invalid: unable to decode EnvironmentWithOverridesType to nil") + } + v, err := d.StrBytes() + if err != nil { + return err + } + // Try to use constant string. + switch EnvironmentWithOverridesType(v) { + case EnvironmentWithOverridesTypeEnvironmentWithOverrides: + *s = EnvironmentWithOverridesTypeEnvironmentWithOverrides + default: + *s = EnvironmentWithOverridesType(v) + } + + return nil +} + +// MarshalJSON implements stdjson.Marshaler. +func (s EnvironmentWithOverridesType) MarshalJSON() ([]byte, error) { + e := jx.Encoder{} + s.Encode(&e) + return e.Bytes(), nil +} + +// UnmarshalJSON implements stdjson.Unmarshaler. +func (s *EnvironmentWithOverridesType) UnmarshalJSON(data []byte) error { d := jx.DecodeBytes(data) return s.Decode(d) } // Encode implements json.Marshaler. -func (s *ListSessionResourcesResponse) Encode(e *jx.Encoder) { +func (s *FileDocumentSource) Encode(e *jx.Encoder) { e.ObjStart() s.encodeFields(e) e.ObjEnd() } // encodeFields encodes fields. -func (s *ListSessionResourcesResponse) encodeFields(e *jx.Encoder) { +func (s *FileDocumentSource) encodeFields(e *jx.Encoder) { { - e.FieldStart("data") - e.ArrStart() - for _, elem := range s.Data { - elem.Encode(e) - } - e.ArrEnd() + e.FieldStart("file_id") + e.Str(s.FileID) } } -var jsonFieldsNameOfListSessionResourcesResponse = [1]string{ - 0: "data", +var jsonFieldsNameOfFileDocumentSource = [1]string{ + 0: "file_id", } -// Decode decodes ListSessionResourcesResponse from json. -func (s *ListSessionResourcesResponse) Decode(d *jx.Decoder) error { +// Decode decodes FileDocumentSource from json. +func (s *FileDocumentSource) Decode(d *jx.Decoder) error { if s == nil { - return errors.New("invalid: unable to decode ListSessionResourcesResponse to nil") + return errors.New("invalid: unable to decode FileDocumentSource to nil") } var requiredBitSet [1]uint8 if err := d.ObjBytes(func(d *jx.Decoder, k []byte) error { switch string(k) { - case "data": + case "file_id": requiredBitSet[0] |= 1 << 0 if err := func() error { - s.Data = make([]SessionResource, 0) - if err := d.Arr(func(d *jx.Decoder) error { - var elem SessionResource - if err := elem.Decode(d); err != nil { - return err - } - s.Data = append(s.Data, elem) - return nil - }); err != nil { + v, err := d.Str() + s.FileID = string(v) + if err != nil { return err } return nil }(); err != nil { - return errors.Wrap(err, "decode field \"data\"") + return errors.Wrap(err, "decode field \"file_id\"") } default: return d.Skip() } return nil }); err != nil { - return errors.Wrap(err, "decode ListSessionResourcesResponse") + return errors.Wrap(err, "decode FileDocumentSource") } // Validate required fields. var failures []validate.FieldError @@ -1822,8 +2556,8 @@ func (s *ListSessionResourcesResponse) Decode(d *jx.Decoder) error { bitIdx := bits.TrailingZeros8(result) fieldIdx := i*8 + bitIdx var name string - if fieldIdx < len(jsonFieldsNameOfListSessionResourcesResponse) { - name = jsonFieldsNameOfListSessionResourcesResponse[fieldIdx] + if fieldIdx < len(jsonFieldsNameOfFileDocumentSource) { + name = jsonFieldsNameOfFileDocumentSource[fieldIdx] } else { name = strconv.Itoa(fieldIdx) } @@ -1844,91 +2578,64 @@ func (s *ListSessionResourcesResponse) Decode(d *jx.Decoder) error { } // MarshalJSON implements stdjson.Marshaler. -func (s *ListSessionResourcesResponse) MarshalJSON() ([]byte, error) { +func (s *FileDocumentSource) MarshalJSON() ([]byte, error) { e := jx.Encoder{} s.Encode(&e) return e.Bytes(), nil } // UnmarshalJSON implements stdjson.Unmarshaler. -func (s *ListSessionResourcesResponse) UnmarshalJSON(data []byte) error { +func (s *FileDocumentSource) UnmarshalJSON(data []byte) error { d := jx.DecodeBytes(data) return s.Decode(d) } // Encode implements json.Marshaler. -func (s *ListSessionThreadsResponse) Encode(e *jx.Encoder) { +func (s *FileImageSource) Encode(e *jx.Encoder) { e.ObjStart() s.encodeFields(e) e.ObjEnd() } // encodeFields encodes fields. -func (s *ListSessionThreadsResponse) encodeFields(e *jx.Encoder) { - { - e.FieldStart("data") - e.ArrStart() - for _, elem := range s.Data { - elem.Encode(e) - } - e.ArrEnd() - } +func (s *FileImageSource) encodeFields(e *jx.Encoder) { { - if s.NextPage.Set { - e.FieldStart("next_page") - s.NextPage.Encode(e) - } + e.FieldStart("file_id") + e.Str(s.FileID) } } -var jsonFieldsNameOfListSessionThreadsResponse = [2]string{ - 0: "data", - 1: "next_page", +var jsonFieldsNameOfFileImageSource = [1]string{ + 0: "file_id", } -// Decode decodes ListSessionThreadsResponse from json. -func (s *ListSessionThreadsResponse) Decode(d *jx.Decoder) error { +// Decode decodes FileImageSource from json. +func (s *FileImageSource) Decode(d *jx.Decoder) error { if s == nil { - return errors.New("invalid: unable to decode ListSessionThreadsResponse to nil") + return errors.New("invalid: unable to decode FileImageSource to nil") } var requiredBitSet [1]uint8 if err := d.ObjBytes(func(d *jx.Decoder, k []byte) error { switch string(k) { - case "data": + case "file_id": requiredBitSet[0] |= 1 << 0 if err := func() error { - s.Data = make([]SessionThread, 0) - if err := d.Arr(func(d *jx.Decoder) error { - var elem SessionThread - if err := elem.Decode(d); err != nil { - return err - } - s.Data = append(s.Data, elem) - return nil - }); err != nil { - return err - } - return nil - }(); err != nil { - return errors.Wrap(err, "decode field \"data\"") - } - case "next_page": - if err := func() error { - s.NextPage.Reset() - if err := s.NextPage.Decode(d); err != nil { + v, err := d.Str() + s.FileID = string(v) + if err != nil { return err } return nil }(); err != nil { - return errors.Wrap(err, "decode field \"next_page\"") + return errors.Wrap(err, "decode field \"file_id\"") } default: return d.Skip() } return nil }); err != nil { - return errors.Wrap(err, "decode ListSessionThreadsResponse") + return errors.Wrap(err, "decode FileImageSource") } // Validate required fields. var failures []validate.FieldError @@ -1945,8 +2652,8 @@ func (s *ListSessionThreadsResponse) Decode(d *jx.Decoder) error { bitIdx := bits.TrailingZeros8(result) fieldIdx := i*8 + bitIdx var name string - if fieldIdx < len(jsonFieldsNameOfListSessionThreadsResponse) { - name = jsonFieldsNameOfListSessionThreadsResponse[fieldIdx] + if fieldIdx < len(jsonFieldsNameOfFileImageSource) { + name = jsonFieldsNameOfFileImageSource[fieldIdx] } else { name = strconv.Itoa(fieldIdx) } @@ -1967,52 +2674,225 @@ func (s *ListSessionThreadsResponse) Decode(d *jx.Decoder) error { } // MarshalJSON implements stdjson.Marshaler. -func (s *ListSessionThreadsResponse) MarshalJSON() ([]byte, error) { +func (s *FileImageSource) MarshalJSON() ([]byte, error) { e := jx.Encoder{} s.Encode(&e) return e.Bytes(), nil } // UnmarshalJSON implements stdjson.Unmarshaler. -func (s *ListSessionThreadsResponse) UnmarshalJSON(data []byte) error { +func (s *FileImageSource) UnmarshalJSON(data []byte) error { d := jx.DecodeBytes(data) return s.Decode(d) } // Encode implements json.Marshaler. -func (s *ListSessionsResponse) Encode(e *jx.Encoder) { +func (s *ImageSource) Encode(e *jx.Encoder) { e.ObjStart() s.encodeFields(e) e.ObjEnd() } // encodeFields encodes fields. -func (s *ListSessionsResponse) encodeFields(e *jx.Encoder) { - { - e.FieldStart("data") - e.ArrStart() - for _, elem := range s.Data { - elem.Encode(e) - } - e.ArrEnd() - } - { - if s.NextPage.Set { - e.FieldStart("next_page") - s.NextPage.Encode(e) +func (s *ImageSource) encodeFields(e *jx.Encoder) { + s.OneOf.encodeFields(e) +} + +var jsonFieldsNameOfImageSource = [0]string{} + +// Decode decodes ImageSource from json. +func (s *ImageSource) Decode(d *jx.Decoder) error { + if s == nil { + return errors.New("invalid: unable to decode ImageSource to nil") + } + if err := d.Capture(func(d *jx.Decoder) error { + return s.OneOf.Decode(d) + }); err != nil { + return errors.Wrap(err, "decode field OneOf") + } + + if err := d.ObjBytes(func(d *jx.Decoder, k []byte) error { + switch string(k) { + default: + return d.Skip() } + }); err != nil { + return errors.Wrap(err, "decode ImageSource") } + + return nil } -var jsonFieldsNameOfListSessionsResponse = [2]string{ +// MarshalJSON implements stdjson.Marshaler. +func (s *ImageSource) MarshalJSON() ([]byte, error) { + e := jx.Encoder{} + s.Encode(&e) + return e.Bytes(), nil +} + +// UnmarshalJSON implements stdjson.Unmarshaler. +func (s *ImageSource) UnmarshalJSON(data []byte) error { + d := jx.DecodeBytes(data) + return s.Decode(d) +} + +// Encode encodes ImageSourceSum as json. +func (s ImageSourceSum) Encode(e *jx.Encoder) { + e.ObjStart() + s.encodeFields(e) + e.ObjEnd() +} + +func (s ImageSourceSum) encodeFields(e *jx.Encoder) { + switch s.Type { + case Base64ImageSourceImageSourceSum: + e.FieldStart("type") + e.Str("base64") + { + s := s.Base64ImageSource + { + e.FieldStart("data") + e.Str(s.Data) + } + { + e.FieldStart("media_type") + e.Str(s.MediaType) + } + } + case UrlImageSourceImageSourceSum: + e.FieldStart("type") + e.Str("url") + { + s := s.UrlImageSource + { + e.FieldStart("url") + e.Str(s.URL) + } + } + case FileImageSourceImageSourceSum: + e.FieldStart("type") + e.Str("file") + { + s := s.FileImageSource + { + e.FieldStart("file_id") + e.Str(s.FileID) + } + } + } +} + +// Decode decodes ImageSourceSum from json. +func (s *ImageSourceSum) Decode(d *jx.Decoder) error { + if s == nil { + return errors.New("invalid: unable to decode ImageSourceSum to nil") + } + // Sum type discriminator. + if typ := d.Next(); typ != jx.Object { + return errors.Errorf("unexpected json type %q", typ) + } + + var found bool + if err := d.Capture(func(d *jx.Decoder) error { + return d.ObjBytes(func(d *jx.Decoder, key []byte) error { + if found { + return d.Skip() + } + switch string(key) { + case "type": + typ, err := d.Str() + if err != nil { + return err + } + switch typ { + case "base64": + s.Type = Base64ImageSourceImageSourceSum + found = true + case "url": + s.Type = UrlImageSourceImageSourceSum + found = true + case "file": + s.Type = FileImageSourceImageSourceSum + found = true + default: + return errors.Errorf("unknown type %s", typ) + } + return nil + } + return d.Skip() + }) + }); err != nil { + return errors.Wrap(err, "capture") + } + if !found { + return errors.New("unable to detect sum type variant") + } + switch s.Type { + case Base64ImageSourceImageSourceSum: + if err := s.Base64ImageSource.Decode(d); err != nil { + return err + } + case UrlImageSourceImageSourceSum: + if err := s.UrlImageSource.Decode(d); err != nil { + return err + } + case FileImageSourceImageSourceSum: + if err := s.FileImageSource.Decode(d); err != nil { + return err + } + default: + return errors.Errorf("inferred invalid type: %s", s.Type) + } + return nil +} + +// MarshalJSON implements stdjson.Marshaler. +func (s ImageSourceSum) MarshalJSON() ([]byte, error) { + e := jx.Encoder{} + s.Encode(&e) + return e.Bytes(), nil +} + +// UnmarshalJSON implements stdjson.Unmarshaler. +func (s *ImageSourceSum) UnmarshalJSON(data []byte) error { + d := jx.DecodeBytes(data) + return s.Decode(d) +} + +// Encode implements json.Marshaler. +func (s *ListSessionEventsResponseWire) Encode(e *jx.Encoder) { + e.ObjStart() + s.encodeFields(e) + e.ObjEnd() +} + +// encodeFields encodes fields. +func (s *ListSessionEventsResponseWire) encodeFields(e *jx.Encoder) { + { + e.FieldStart("data") + e.ArrStart() + for _, elem := range s.Data { + elem.Encode(e) + } + e.ArrEnd() + } + { + if s.NextPage.Set { + e.FieldStart("next_page") + s.NextPage.Encode(e) + } + } +} + +var jsonFieldsNameOfListSessionEventsResponseWire = [2]string{ 0: "data", 1: "next_page", } -// Decode decodes ListSessionsResponse from json. -func (s *ListSessionsResponse) Decode(d *jx.Decoder) error { +// Decode decodes ListSessionEventsResponseWire from json. +func (s *ListSessionEventsResponseWire) Decode(d *jx.Decoder) error { if s == nil { - return errors.New("invalid: unable to decode ListSessionsResponse to nil") + return errors.New("invalid: unable to decode ListSessionEventsResponseWire to nil") } var requiredBitSet [1]uint8 @@ -2021,9 +2901,9 @@ func (s *ListSessionsResponse) Decode(d *jx.Decoder) error { case "data": requiredBitSet[0] |= 1 << 0 if err := func() error { - s.Data = make([]Session, 0) + s.Data = make([]ListSessionEventsResponseWireDataItem, 0) if err := d.Arr(func(d *jx.Decoder) error { - var elem Session + var elem ListSessionEventsResponseWireDataItem if err := elem.Decode(d); err != nil { return err } @@ -2051,7 +2931,7 @@ func (s *ListSessionsResponse) Decode(d *jx.Decoder) error { } return nil }); err != nil { - return errors.Wrap(err, "decode ListSessionsResponse") + return errors.Wrap(err, "decode ListSessionEventsResponseWire") } // Validate required fields. var failures []validate.FieldError @@ -2068,8 +2948,8 @@ func (s *ListSessionsResponse) Decode(d *jx.Decoder) error { bitIdx := bits.TrailingZeros8(result) fieldIdx := i*8 + bitIdx var name string - if fieldIdx < len(jsonFieldsNameOfListSessionsResponse) { - name = jsonFieldsNameOfListSessionsResponse[fieldIdx] + if fieldIdx < len(jsonFieldsNameOfListSessionEventsResponseWire) { + name = jsonFieldsNameOfListSessionEventsResponseWire[fieldIdx] } else { name = strconv.Itoa(fieldIdx) } @@ -2090,96 +2970,132 @@ func (s *ListSessionsResponse) Decode(d *jx.Decoder) error { } // MarshalJSON implements stdjson.Marshaler. -func (s *ListSessionsResponse) MarshalJSON() ([]byte, error) { +func (s *ListSessionEventsResponseWire) MarshalJSON() ([]byte, error) { e := jx.Encoder{} s.Encode(&e) return e.Bytes(), nil } // UnmarshalJSON implements stdjson.Unmarshaler. -func (s *ListSessionsResponse) UnmarshalJSON(data []byte) error { +func (s *ListSessionEventsResponseWire) UnmarshalJSON(data []byte) error { d := jx.DecodeBytes(data) return s.Decode(d) } // Encode implements json.Marshaler. -func (s *ManagedAgentsDocumentBlock) Encode(e *jx.Encoder) { +func (s ListSessionEventsResponseWireDataItem) Encode(e *jx.Encoder) { e.ObjStart() s.encodeFields(e) e.ObjEnd() } -// encodeFields encodes fields. -func (s *ManagedAgentsDocumentBlock) encodeFields(e *jx.Encoder) { - { - e.FieldStart("source") - s.Source.Encode(e) - } - { - if s.Title.Set { - e.FieldStart("title") - s.Title.Encode(e) - } - } - { - if s.Context.Set { - e.FieldStart("context") - s.Context.Encode(e) +// encodeFields implements json.Marshaler. +func (s ListSessionEventsResponseWireDataItem) encodeFields(e *jx.Encoder) { + for k, elem := range s { + e.FieldStart(k) + + if len(elem) != 0 { + e.Raw(elem) } } } -var jsonFieldsNameOfManagedAgentsDocumentBlock = [3]string{ - 0: "source", - 1: "title", - 2: "context", -} - -// Decode decodes ManagedAgentsDocumentBlock from json. -func (s *ManagedAgentsDocumentBlock) Decode(d *jx.Decoder) error { +// Decode decodes ListSessionEventsResponseWireDataItem from json. +func (s *ListSessionEventsResponseWireDataItem) Decode(d *jx.Decoder) error { if s == nil { - return errors.New("invalid: unable to decode ManagedAgentsDocumentBlock to nil") + return errors.New("invalid: unable to decode ListSessionEventsResponseWireDataItem to nil") } - var requiredBitSet [1]uint8 - + m := s.init() if err := d.ObjBytes(func(d *jx.Decoder, k []byte) error { - switch string(k) { - case "source": - requiredBitSet[0] |= 1 << 0 - if err := func() error { - if err := s.Source.Decode(d); err != nil { - return err - } - return nil - }(); err != nil { - return errors.Wrap(err, "decode field \"source\"") - } - case "title": - if err := func() error { - s.Title.Reset() - if err := s.Title.Decode(d); err != nil { - return err - } - return nil - }(); err != nil { - return errors.Wrap(err, "decode field \"title\"") + var elem jx.Raw + if err := func() error { + v, err := d.RawAppend(nil) + elem = jx.Raw(v) + if err != nil { + return err } - case "context": + return nil + }(); err != nil { + return errors.Wrapf(err, "decode field %q", k) + } + m[string(k)] = elem + return nil + }); err != nil { + return errors.Wrap(err, "decode ListSessionEventsResponseWireDataItem") + } + + return nil +} + +// MarshalJSON implements stdjson.Marshaler. +func (s ListSessionEventsResponseWireDataItem) MarshalJSON() ([]byte, error) { + e := jx.Encoder{} + s.Encode(&e) + return e.Bytes(), nil +} + +// UnmarshalJSON implements stdjson.Unmarshaler. +func (s *ListSessionEventsResponseWireDataItem) UnmarshalJSON(data []byte) error { + d := jx.DecodeBytes(data) + return s.Decode(d) +} + +// Encode implements json.Marshaler. +func (s *ListSessionResourcesResponse) Encode(e *jx.Encoder) { + e.ObjStart() + s.encodeFields(e) + e.ObjEnd() +} + +// encodeFields encodes fields. +func (s *ListSessionResourcesResponse) encodeFields(e *jx.Encoder) { + { + e.FieldStart("data") + e.ArrStart() + for _, elem := range s.Data { + elem.Encode(e) + } + e.ArrEnd() + } +} + +var jsonFieldsNameOfListSessionResourcesResponse = [1]string{ + 0: "data", +} + +// Decode decodes ListSessionResourcesResponse from json. +func (s *ListSessionResourcesResponse) Decode(d *jx.Decoder) error { + if s == nil { + return errors.New("invalid: unable to decode ListSessionResourcesResponse to nil") + } + var requiredBitSet [1]uint8 + + if err := d.ObjBytes(func(d *jx.Decoder, k []byte) error { + switch string(k) { + case "data": + requiredBitSet[0] |= 1 << 0 if err := func() error { - s.Context.Reset() - if err := s.Context.Decode(d); err != nil { + s.Data = make([]SessionResource, 0) + if err := d.Arr(func(d *jx.Decoder) error { + var elem SessionResource + if err := elem.Decode(d); err != nil { + return err + } + s.Data = append(s.Data, elem) + return nil + }); err != nil { return err } return nil }(); err != nil { - return errors.Wrap(err, "decode field \"context\"") + return errors.Wrap(err, "decode field \"data\"") } default: return d.Skip() } return nil }); err != nil { - return errors.Wrap(err, "decode ManagedAgentsDocumentBlock") + return errors.Wrap(err, "decode ListSessionResourcesResponse") } // Validate required fields. var failures []validate.FieldError @@ -2196,8 +3112,8 @@ func (s *ManagedAgentsDocumentBlock) Decode(d *jx.Decoder) error { bitIdx := bits.TrailingZeros8(result) fieldIdx := i*8 + bitIdx var name string - if fieldIdx < len(jsonFieldsNameOfManagedAgentsDocumentBlock) { - name = jsonFieldsNameOfManagedAgentsDocumentBlock[fieldIdx] + if fieldIdx < len(jsonFieldsNameOfListSessionResourcesResponse) { + name = jsonFieldsNameOfListSessionResourcesResponse[fieldIdx] } else { name = strconv.Itoa(fieldIdx) } @@ -2218,116 +3134,490 @@ func (s *ManagedAgentsDocumentBlock) Decode(d *jx.Decoder) error { } // MarshalJSON implements stdjson.Marshaler. -func (s *ManagedAgentsDocumentBlock) MarshalJSON() ([]byte, error) { +func (s *ListSessionResourcesResponse) MarshalJSON() ([]byte, error) { e := jx.Encoder{} s.Encode(&e) return e.Bytes(), nil } // UnmarshalJSON implements stdjson.Unmarshaler. -func (s *ManagedAgentsDocumentBlock) UnmarshalJSON(data []byte) error { +func (s *ListSessionResourcesResponse) UnmarshalJSON(data []byte) error { d := jx.DecodeBytes(data) return s.Decode(d) } // Encode implements json.Marshaler. -func (s *ManagedAgentsEventParams) Encode(e *jx.Encoder) { +func (s *ListSessionThreadsResponse) Encode(e *jx.Encoder) { e.ObjStart() s.encodeFields(e) e.ObjEnd() } // encodeFields encodes fields. -func (s *ManagedAgentsEventParams) encodeFields(e *jx.Encoder) { - s.OneOf.encodeFields(e) +func (s *ListSessionThreadsResponse) encodeFields(e *jx.Encoder) { + { + e.FieldStart("data") + e.ArrStart() + for _, elem := range s.Data { + elem.Encode(e) + } + e.ArrEnd() + } + { + if s.NextPage.Set { + e.FieldStart("next_page") + s.NextPage.Encode(e) + } + } } -var jsonFieldsNameOfManagedAgentsEventParams = [0]string{} +var jsonFieldsNameOfListSessionThreadsResponse = [2]string{ + 0: "data", + 1: "next_page", +} -// Decode decodes ManagedAgentsEventParams from json. -func (s *ManagedAgentsEventParams) Decode(d *jx.Decoder) error { +// Decode decodes ListSessionThreadsResponse from json. +func (s *ListSessionThreadsResponse) Decode(d *jx.Decoder) error { if s == nil { - return errors.New("invalid: unable to decode ManagedAgentsEventParams to nil") - } - if err := d.Capture(func(d *jx.Decoder) error { - return s.OneOf.Decode(d) - }); err != nil { - return errors.Wrap(err, "decode field OneOf") + return errors.New("invalid: unable to decode ListSessionThreadsResponse to nil") } + var requiredBitSet [1]uint8 if err := d.ObjBytes(func(d *jx.Decoder, k []byte) error { switch string(k) { + case "data": + requiredBitSet[0] |= 1 << 0 + if err := func() error { + s.Data = make([]SessionThread, 0) + if err := d.Arr(func(d *jx.Decoder) error { + var elem SessionThread + if err := elem.Decode(d); err != nil { + return err + } + s.Data = append(s.Data, elem) + return nil + }); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"data\"") + } + case "next_page": + if err := func() error { + s.NextPage.Reset() + if err := s.NextPage.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"next_page\"") + } default: return d.Skip() } + return nil }); err != nil { - return errors.Wrap(err, "decode ManagedAgentsEventParams") + return errors.Wrap(err, "decode ListSessionThreadsResponse") + } + // Validate required fields. + var failures []validate.FieldError + for i, mask := range [1]uint8{ + 0b00000001, + } { + if result := (requiredBitSet[i] & mask) ^ mask; result != 0 { + // Mask only required fields and check equality to mask using XOR. + // + // If XOR result is not zero, result is not equal to expected, so some fields are missed. + // Bits of fields which would be set are actually bits of missed fields. + missed := bits.OnesCount8(result) + for bitN := 0; bitN < missed; bitN++ { + bitIdx := bits.TrailingZeros8(result) + fieldIdx := i*8 + bitIdx + var name string + if fieldIdx < len(jsonFieldsNameOfListSessionThreadsResponse) { + name = jsonFieldsNameOfListSessionThreadsResponse[fieldIdx] + } else { + name = strconv.Itoa(fieldIdx) + } + failures = append(failures, validate.FieldError{ + Name: name, + Error: validate.ErrFieldRequired, + }) + // Reset bit. + result &^= 1 << bitIdx + } + } + } + if len(failures) > 0 { + return &validate.Error{Fields: failures} } return nil } // MarshalJSON implements stdjson.Marshaler. -func (s *ManagedAgentsEventParams) MarshalJSON() ([]byte, error) { +func (s *ListSessionThreadsResponse) MarshalJSON() ([]byte, error) { e := jx.Encoder{} s.Encode(&e) return e.Bytes(), nil } // UnmarshalJSON implements stdjson.Unmarshaler. -func (s *ManagedAgentsEventParams) UnmarshalJSON(data []byte) error { +func (s *ListSessionThreadsResponse) UnmarshalJSON(data []byte) error { d := jx.DecodeBytes(data) return s.Decode(d) } -// Encode encodes ManagedAgentsEventParamsSum as json. -func (s ManagedAgentsEventParamsSum) Encode(e *jx.Encoder) { +// Encode implements json.Marshaler. +func (s *ListSessionsResponse) Encode(e *jx.Encoder) { e.ObjStart() s.encodeFields(e) e.ObjEnd() } -func (s ManagedAgentsEventParamsSum) encodeFields(e *jx.Encoder) { - switch s.Type { - case ManagedAgentsUserMessageEventParamsManagedAgentsEventParamsSum: - e.FieldStart("type") - e.Str("user.message") - { - s := s.ManagedAgentsUserMessageEventParams - { - if s.Content != nil { - e.FieldStart("content") - e.ArrStart() - for _, elem := range s.Content { - elem.Encode(e) - } - e.ArrEnd() - } - } - { - if s.SessionID.Set { - e.FieldStart("session_id") - s.SessionID.Encode(e) - } - } +// encodeFields encodes fields. +func (s *ListSessionsResponse) encodeFields(e *jx.Encoder) { + { + e.FieldStart("data") + e.ArrStart() + for _, elem := range s.Data { + elem.Encode(e) } - case ManagedAgentsSystemMessageEventParamsManagedAgentsEventParamsSum: - e.FieldStart("type") - e.Str("system.message") - { - s := s.ManagedAgentsSystemMessageEventParams - { - if s.Content != nil { - e.FieldStart("content") - e.ArrStart() - for _, elem := range s.Content { - elem.Encode(e) - } - e.ArrEnd() - } - } - { - if s.SessionID.Set { + e.ArrEnd() + } + { + if s.NextPage.Set { + e.FieldStart("next_page") + s.NextPage.Encode(e) + } + } +} + +var jsonFieldsNameOfListSessionsResponse = [2]string{ + 0: "data", + 1: "next_page", +} + +// Decode decodes ListSessionsResponse from json. +func (s *ListSessionsResponse) Decode(d *jx.Decoder) error { + if s == nil { + return errors.New("invalid: unable to decode ListSessionsResponse to nil") + } + var requiredBitSet [1]uint8 + + if err := d.ObjBytes(func(d *jx.Decoder, k []byte) error { + switch string(k) { + case "data": + requiredBitSet[0] |= 1 << 0 + if err := func() error { + s.Data = make([]Session, 0) + if err := d.Arr(func(d *jx.Decoder) error { + var elem Session + if err := elem.Decode(d); err != nil { + return err + } + s.Data = append(s.Data, elem) + return nil + }); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"data\"") + } + case "next_page": + if err := func() error { + s.NextPage.Reset() + if err := s.NextPage.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"next_page\"") + } + default: + return d.Skip() + } + return nil + }); err != nil { + return errors.Wrap(err, "decode ListSessionsResponse") + } + // Validate required fields. + var failures []validate.FieldError + for i, mask := range [1]uint8{ + 0b00000001, + } { + if result := (requiredBitSet[i] & mask) ^ mask; result != 0 { + // Mask only required fields and check equality to mask using XOR. + // + // If XOR result is not zero, result is not equal to expected, so some fields are missed. + // Bits of fields which would be set are actually bits of missed fields. + missed := bits.OnesCount8(result) + for bitN := 0; bitN < missed; bitN++ { + bitIdx := bits.TrailingZeros8(result) + fieldIdx := i*8 + bitIdx + var name string + if fieldIdx < len(jsonFieldsNameOfListSessionsResponse) { + name = jsonFieldsNameOfListSessionsResponse[fieldIdx] + } else { + name = strconv.Itoa(fieldIdx) + } + failures = append(failures, validate.FieldError{ + Name: name, + Error: validate.ErrFieldRequired, + }) + // Reset bit. + result &^= 1 << bitIdx + } + } + } + if len(failures) > 0 { + return &validate.Error{Fields: failures} + } + + return nil +} + +// MarshalJSON implements stdjson.Marshaler. +func (s *ListSessionsResponse) MarshalJSON() ([]byte, error) { + e := jx.Encoder{} + s.Encode(&e) + return e.Bytes(), nil +} + +// UnmarshalJSON implements stdjson.Unmarshaler. +func (s *ListSessionsResponse) UnmarshalJSON(data []byte) error { + d := jx.DecodeBytes(data) + return s.Decode(d) +} + +// Encode implements json.Marshaler. +func (s *ManagedAgentsDocumentBlock) Encode(e *jx.Encoder) { + e.ObjStart() + s.encodeFields(e) + e.ObjEnd() +} + +// encodeFields encodes fields. +func (s *ManagedAgentsDocumentBlock) encodeFields(e *jx.Encoder) { + { + e.FieldStart("source") + s.Source.Encode(e) + } + { + if s.Title.Set { + e.FieldStart("title") + s.Title.Encode(e) + } + } + { + if s.Context.Set { + e.FieldStart("context") + s.Context.Encode(e) + } + } +} + +var jsonFieldsNameOfManagedAgentsDocumentBlock = [3]string{ + 0: "source", + 1: "title", + 2: "context", +} + +// Decode decodes ManagedAgentsDocumentBlock from json. +func (s *ManagedAgentsDocumentBlock) Decode(d *jx.Decoder) error { + if s == nil { + return errors.New("invalid: unable to decode ManagedAgentsDocumentBlock to nil") + } + var requiredBitSet [1]uint8 + + if err := d.ObjBytes(func(d *jx.Decoder, k []byte) error { + switch string(k) { + case "source": + requiredBitSet[0] |= 1 << 0 + if err := func() error { + if err := s.Source.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"source\"") + } + case "title": + if err := func() error { + s.Title.Reset() + if err := s.Title.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"title\"") + } + case "context": + if err := func() error { + s.Context.Reset() + if err := s.Context.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"context\"") + } + default: + return d.Skip() + } + return nil + }); err != nil { + return errors.Wrap(err, "decode ManagedAgentsDocumentBlock") + } + // Validate required fields. + var failures []validate.FieldError + for i, mask := range [1]uint8{ + 0b00000001, + } { + if result := (requiredBitSet[i] & mask) ^ mask; result != 0 { + // Mask only required fields and check equality to mask using XOR. + // + // If XOR result is not zero, result is not equal to expected, so some fields are missed. + // Bits of fields which would be set are actually bits of missed fields. + missed := bits.OnesCount8(result) + for bitN := 0; bitN < missed; bitN++ { + bitIdx := bits.TrailingZeros8(result) + fieldIdx := i*8 + bitIdx + var name string + if fieldIdx < len(jsonFieldsNameOfManagedAgentsDocumentBlock) { + name = jsonFieldsNameOfManagedAgentsDocumentBlock[fieldIdx] + } else { + name = strconv.Itoa(fieldIdx) + } + failures = append(failures, validate.FieldError{ + Name: name, + Error: validate.ErrFieldRequired, + }) + // Reset bit. + result &^= 1 << bitIdx + } + } + } + if len(failures) > 0 { + return &validate.Error{Fields: failures} + } + + return nil +} + +// MarshalJSON implements stdjson.Marshaler. +func (s *ManagedAgentsDocumentBlock) MarshalJSON() ([]byte, error) { + e := jx.Encoder{} + s.Encode(&e) + return e.Bytes(), nil +} + +// UnmarshalJSON implements stdjson.Unmarshaler. +func (s *ManagedAgentsDocumentBlock) UnmarshalJSON(data []byte) error { + d := jx.DecodeBytes(data) + return s.Decode(d) +} + +// Encode implements json.Marshaler. +func (s *ManagedAgentsEventParams) Encode(e *jx.Encoder) { + e.ObjStart() + s.encodeFields(e) + e.ObjEnd() +} + +// encodeFields encodes fields. +func (s *ManagedAgentsEventParams) encodeFields(e *jx.Encoder) { + s.OneOf.encodeFields(e) +} + +var jsonFieldsNameOfManagedAgentsEventParams = [0]string{} + +// Decode decodes ManagedAgentsEventParams from json. +func (s *ManagedAgentsEventParams) Decode(d *jx.Decoder) error { + if s == nil { + return errors.New("invalid: unable to decode ManagedAgentsEventParams to nil") + } + if err := d.Capture(func(d *jx.Decoder) error { + return s.OneOf.Decode(d) + }); err != nil { + return errors.Wrap(err, "decode field OneOf") + } + + if err := d.ObjBytes(func(d *jx.Decoder, k []byte) error { + switch string(k) { + default: + return d.Skip() + } + }); err != nil { + return errors.Wrap(err, "decode ManagedAgentsEventParams") + } + + return nil +} + +// MarshalJSON implements stdjson.Marshaler. +func (s *ManagedAgentsEventParams) MarshalJSON() ([]byte, error) { + e := jx.Encoder{} + s.Encode(&e) + return e.Bytes(), nil +} + +// UnmarshalJSON implements stdjson.Unmarshaler. +func (s *ManagedAgentsEventParams) UnmarshalJSON(data []byte) error { + d := jx.DecodeBytes(data) + return s.Decode(d) +} + +// Encode encodes ManagedAgentsEventParamsSum as json. +func (s ManagedAgentsEventParamsSum) Encode(e *jx.Encoder) { + e.ObjStart() + s.encodeFields(e) + e.ObjEnd() +} + +func (s ManagedAgentsEventParamsSum) encodeFields(e *jx.Encoder) { + switch s.Type { + case ManagedAgentsUserMessageEventParamsManagedAgentsEventParamsSum: + e.FieldStart("type") + e.Str("user.message") + { + s := s.ManagedAgentsUserMessageEventParams + { + if s.Content != nil { + e.FieldStart("content") + e.ArrStart() + for _, elem := range s.Content { + elem.Encode(e) + } + e.ArrEnd() + } + } + { + if s.SessionID.Set { + e.FieldStart("session_id") + s.SessionID.Encode(e) + } + } + } + case ManagedAgentsSystemMessageEventParamsManagedAgentsEventParamsSum: + e.FieldStart("type") + e.Str("system.message") + { + s := s.ManagedAgentsSystemMessageEventParams + { + if s.Content != nil { + e.FieldStart("content") + e.ArrStart() + for _, elem := range s.Content { + elem.Encode(e) + } + e.ArrEnd() + } + } + { + if s.SessionID.Set { e.FieldStart("session_id") s.SessionID.Encode(e) } @@ -4445,51 +5735,236 @@ func (s *ManagedAgentsUserMessageEventParams) Decode(d *jx.Decoder) error { }(); err != nil { return errors.Wrap(err, "decode field \"session_id\"") } - default: - return d.Skip() + default: + return d.Skip() + } + return nil + }); err != nil { + return errors.Wrap(err, "decode ManagedAgentsUserMessageEventParams") + } + + return nil +} + +// MarshalJSON implements stdjson.Marshaler. +func (s *ManagedAgentsUserMessageEventParams) MarshalJSON() ([]byte, error) { + e := jx.Encoder{} + s.Encode(&e) + return e.Bytes(), nil +} + +// UnmarshalJSON implements stdjson.Unmarshaler. +func (s *ManagedAgentsUserMessageEventParams) UnmarshalJSON(data []byte) error { + d := jx.DecodeBytes(data) + return s.Decode(d) +} + +// Encode implements json.Marshaler. +func (s *ManagedAgentsUserToolConfirmationEventParams) Encode(e *jx.Encoder) { + e.ObjStart() + s.encodeFields(e) + e.ObjEnd() +} + +// encodeFields encodes fields. +func (s *ManagedAgentsUserToolConfirmationEventParams) encodeFields(e *jx.Encoder) { + { + e.FieldStart("result") + s.Result.Encode(e) + } + { + e.FieldStart("tool_use_id") + e.Str(s.ToolUseID) + } + { + if s.DenyMessage.Set { + e.FieldStart("deny_message") + s.DenyMessage.Encode(e) + } + } + { + if s.SessionID.Set { + e.FieldStart("session_id") + s.SessionID.Encode(e) + } + } + { + if s.SessionThreadID.Set { + e.FieldStart("session_thread_id") + s.SessionThreadID.Encode(e) + } + } + { + if s.TurnID.Set { + e.FieldStart("turn_id") + s.TurnID.Encode(e) + } + } +} + +var jsonFieldsNameOfManagedAgentsUserToolConfirmationEventParams = [6]string{ + 0: "result", + 1: "tool_use_id", + 2: "deny_message", + 3: "session_id", + 4: "session_thread_id", + 5: "turn_id", +} + +// Decode decodes ManagedAgentsUserToolConfirmationEventParams from json. +func (s *ManagedAgentsUserToolConfirmationEventParams) Decode(d *jx.Decoder) error { + if s == nil { + return errors.New("invalid: unable to decode ManagedAgentsUserToolConfirmationEventParams to nil") + } + var requiredBitSet [1]uint8 + + if err := d.ObjBytes(func(d *jx.Decoder, k []byte) error { + switch string(k) { + case "result": + requiredBitSet[0] |= 1 << 0 + if err := func() error { + if err := s.Result.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"result\"") + } + case "tool_use_id": + requiredBitSet[0] |= 1 << 1 + if err := func() error { + v, err := d.Str() + s.ToolUseID = string(v) + if err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"tool_use_id\"") + } + case "deny_message": + if err := func() error { + s.DenyMessage.Reset() + if err := s.DenyMessage.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"deny_message\"") + } + case "session_id": + if err := func() error { + s.SessionID.Reset() + if err := s.SessionID.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"session_id\"") + } + case "session_thread_id": + if err := func() error { + s.SessionThreadID.Reset() + if err := s.SessionThreadID.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"session_thread_id\"") + } + case "turn_id": + if err := func() error { + s.TurnID.Reset() + if err := s.TurnID.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"turn_id\"") + } + default: + return d.Skip() + } + return nil + }); err != nil { + return errors.Wrap(err, "decode ManagedAgentsUserToolConfirmationEventParams") + } + // Validate required fields. + var failures []validate.FieldError + for i, mask := range [1]uint8{ + 0b00000011, + } { + if result := (requiredBitSet[i] & mask) ^ mask; result != 0 { + // Mask only required fields and check equality to mask using XOR. + // + // If XOR result is not zero, result is not equal to expected, so some fields are missed. + // Bits of fields which would be set are actually bits of missed fields. + missed := bits.OnesCount8(result) + for bitN := 0; bitN < missed; bitN++ { + bitIdx := bits.TrailingZeros8(result) + fieldIdx := i*8 + bitIdx + var name string + if fieldIdx < len(jsonFieldsNameOfManagedAgentsUserToolConfirmationEventParams) { + name = jsonFieldsNameOfManagedAgentsUserToolConfirmationEventParams[fieldIdx] + } else { + name = strconv.Itoa(fieldIdx) + } + failures = append(failures, validate.FieldError{ + Name: name, + Error: validate.ErrFieldRequired, + }) + // Reset bit. + result &^= 1 << bitIdx + } } - return nil - }); err != nil { - return errors.Wrap(err, "decode ManagedAgentsUserMessageEventParams") + } + if len(failures) > 0 { + return &validate.Error{Fields: failures} } return nil } // MarshalJSON implements stdjson.Marshaler. -func (s *ManagedAgentsUserMessageEventParams) MarshalJSON() ([]byte, error) { +func (s *ManagedAgentsUserToolConfirmationEventParams) MarshalJSON() ([]byte, error) { e := jx.Encoder{} s.Encode(&e) return e.Bytes(), nil } // UnmarshalJSON implements stdjson.Unmarshaler. -func (s *ManagedAgentsUserMessageEventParams) UnmarshalJSON(data []byte) error { +func (s *ManagedAgentsUserToolConfirmationEventParams) UnmarshalJSON(data []byte) error { d := jx.DecodeBytes(data) return s.Decode(d) } // Encode implements json.Marshaler. -func (s *ManagedAgentsUserToolConfirmationEventParams) Encode(e *jx.Encoder) { +func (s *ManagedAgentsUserToolResultEventParams) Encode(e *jx.Encoder) { e.ObjStart() s.encodeFields(e) e.ObjEnd() } // encodeFields encodes fields. -func (s *ManagedAgentsUserToolConfirmationEventParams) encodeFields(e *jx.Encoder) { - { - e.FieldStart("result") - s.Result.Encode(e) - } +func (s *ManagedAgentsUserToolResultEventParams) encodeFields(e *jx.Encoder) { { e.FieldStart("tool_use_id") e.Str(s.ToolUseID) } { - if s.DenyMessage.Set { - e.FieldStart("deny_message") - s.DenyMessage.Encode(e) + if s.Content != nil { + e.FieldStart("content") + e.ArrStart() + for _, elem := range s.Content { + elem.Encode(e) + } + e.ArrEnd() + } + } + { + if s.IsError.Set { + e.FieldStart("is_error") + s.IsError.Encode(e) } } { @@ -4504,63 +5979,63 @@ func (s *ManagedAgentsUserToolConfirmationEventParams) encodeFields(e *jx.Encode s.SessionThreadID.Encode(e) } } - { - if s.TurnID.Set { - e.FieldStart("turn_id") - s.TurnID.Encode(e) - } - } } -var jsonFieldsNameOfManagedAgentsUserToolConfirmationEventParams = [6]string{ - 0: "result", - 1: "tool_use_id", - 2: "deny_message", +var jsonFieldsNameOfManagedAgentsUserToolResultEventParams = [5]string{ + 0: "tool_use_id", + 1: "content", + 2: "is_error", 3: "session_id", 4: "session_thread_id", - 5: "turn_id", } -// Decode decodes ManagedAgentsUserToolConfirmationEventParams from json. -func (s *ManagedAgentsUserToolConfirmationEventParams) Decode(d *jx.Decoder) error { +// Decode decodes ManagedAgentsUserToolResultEventParams from json. +func (s *ManagedAgentsUserToolResultEventParams) Decode(d *jx.Decoder) error { if s == nil { - return errors.New("invalid: unable to decode ManagedAgentsUserToolConfirmationEventParams to nil") + return errors.New("invalid: unable to decode ManagedAgentsUserToolResultEventParams to nil") } var requiredBitSet [1]uint8 if err := d.ObjBytes(func(d *jx.Decoder, k []byte) error { switch string(k) { - case "result": + case "tool_use_id": requiredBitSet[0] |= 1 << 0 if err := func() error { - if err := s.Result.Decode(d); err != nil { + v, err := d.Str() + s.ToolUseID = string(v) + if err != nil { return err } return nil }(); err != nil { - return errors.Wrap(err, "decode field \"result\"") + return errors.Wrap(err, "decode field \"tool_use_id\"") } - case "tool_use_id": - requiredBitSet[0] |= 1 << 1 + case "content": if err := func() error { - v, err := d.Str() - s.ToolUseID = string(v) - if err != nil { + s.Content = make([]ManagedAgentsToolResultContentBlock, 0) + if err := d.Arr(func(d *jx.Decoder) error { + var elem ManagedAgentsToolResultContentBlock + if err := elem.Decode(d); err != nil { + return err + } + s.Content = append(s.Content, elem) + return nil + }); err != nil { return err } return nil }(); err != nil { - return errors.Wrap(err, "decode field \"tool_use_id\"") + return errors.Wrap(err, "decode field \"content\"") } - case "deny_message": + case "is_error": if err := func() error { - s.DenyMessage.Reset() - if err := s.DenyMessage.Decode(d); err != nil { + s.IsError.Reset() + if err := s.IsError.Decode(d); err != nil { return err } return nil }(); err != nil { - return errors.Wrap(err, "decode field \"deny_message\"") + return errors.Wrap(err, "decode field \"is_error\"") } case "session_id": if err := func() error { @@ -4582,27 +6057,17 @@ func (s *ManagedAgentsUserToolConfirmationEventParams) Decode(d *jx.Decoder) err }(); err != nil { return errors.Wrap(err, "decode field \"session_thread_id\"") } - case "turn_id": - if err := func() error { - s.TurnID.Reset() - if err := s.TurnID.Decode(d); err != nil { - return err - } - return nil - }(); err != nil { - return errors.Wrap(err, "decode field \"turn_id\"") - } default: return d.Skip() } return nil }); err != nil { - return errors.Wrap(err, "decode ManagedAgentsUserToolConfirmationEventParams") + return errors.Wrap(err, "decode ManagedAgentsUserToolResultEventParams") } // Validate required fields. var failures []validate.FieldError for i, mask := range [1]uint8{ - 0b00000011, + 0b00000001, } { if result := (requiredBitSet[i] & mask) ^ mask; result != 0 { // Mask only required fields and check equality to mask using XOR. @@ -4614,8 +6079,8 @@ func (s *ManagedAgentsUserToolConfirmationEventParams) Decode(d *jx.Decoder) err bitIdx := bits.TrailingZeros8(result) fieldIdx := i*8 + bitIdx var name string - if fieldIdx < len(jsonFieldsNameOfManagedAgentsUserToolConfirmationEventParams) { - name = jsonFieldsNameOfManagedAgentsUserToolConfirmationEventParams[fieldIdx] + if fieldIdx < len(jsonFieldsNameOfManagedAgentsUserToolResultEventParams) { + name = jsonFieldsNameOfManagedAgentsUserToolResultEventParams[fieldIdx] } else { name = strconv.Itoa(fieldIdx) } @@ -4628,248 +6093,436 @@ func (s *ManagedAgentsUserToolConfirmationEventParams) Decode(d *jx.Decoder) err } } } - if len(failures) > 0 { - return &validate.Error{Fields: failures} + if len(failures) > 0 { + return &validate.Error{Fields: failures} + } + + return nil +} + +// MarshalJSON implements stdjson.Marshaler. +func (s *ManagedAgentsUserToolResultEventParams) MarshalJSON() ([]byte, error) { + e := jx.Encoder{} + s.Encode(&e) + return e.Bytes(), nil +} + +// UnmarshalJSON implements stdjson.Unmarshaler. +func (s *ManagedAgentsUserToolResultEventParams) UnmarshalJSON(data []byte) error { + d := jx.DecodeBytes(data) + return s.Decode(d) +} + +// Encode implements json.Marshaler. +func (s *ModelOverrides) Encode(e *jx.Encoder) { + e.ObjStart() + s.encodeFields(e) + e.ObjEnd() +} + +// encodeFields encodes fields. +func (s *ModelOverrides) encodeFields(e *jx.Encoder) { + { + if s.Speed.Set { + e.FieldStart("speed") + s.Speed.Encode(e) + } + } + { + if s.Thinking.Set { + e.FieldStart("thinking") + s.Thinking.Encode(e) + } + } + { + if s.ReasoningEffort.Set { + e.FieldStart("reasoning_effort") + s.ReasoningEffort.Encode(e) + } + } +} + +var jsonFieldsNameOfModelOverrides = [3]string{ + 0: "speed", + 1: "thinking", + 2: "reasoning_effort", +} + +// Decode decodes ModelOverrides from json. +func (s *ModelOverrides) Decode(d *jx.Decoder) error { + if s == nil { + return errors.New("invalid: unable to decode ModelOverrides to nil") + } + + if err := d.ObjBytes(func(d *jx.Decoder, k []byte) error { + switch string(k) { + case "speed": + if err := func() error { + s.Speed.Reset() + if err := s.Speed.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"speed\"") + } + case "thinking": + if err := func() error { + s.Thinking.Reset() + if err := s.Thinking.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"thinking\"") + } + case "reasoning_effort": + if err := func() error { + s.ReasoningEffort.Reset() + if err := s.ReasoningEffort.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"reasoning_effort\"") + } + default: + return d.Skip() + } + return nil + }); err != nil { + return errors.Wrap(err, "decode ModelOverrides") + } + + return nil +} + +// MarshalJSON implements stdjson.Marshaler. +func (s *ModelOverrides) MarshalJSON() ([]byte, error) { + e := jx.Encoder{} + s.Encode(&e) + return e.Bytes(), nil +} + +// UnmarshalJSON implements stdjson.Unmarshaler. +func (s *ModelOverrides) UnmarshalJSON(data []byte) error { + d := jx.DecodeBytes(data) + return s.Decode(d) +} + +// Encode encodes AgentRefMultiagent as json. +func (o OptAgentRefMultiagent) Encode(e *jx.Encoder) { + if !o.Set { + return + } + o.Value.Encode(e) +} + +// Decode decodes AgentRefMultiagent from json. +func (o *OptAgentRefMultiagent) Decode(d *jx.Decoder) error { + if o == nil { + return errors.New("invalid: unable to decode OptAgentRefMultiagent to nil") + } + o.Set = true + o.Value = make(AgentRefMultiagent) + if err := o.Value.Decode(d); err != nil { + return err + } + return nil +} + +// MarshalJSON implements stdjson.Marshaler. +func (s OptAgentRefMultiagent) MarshalJSON() ([]byte, error) { + e := jx.Encoder{} + s.Encode(&e) + return e.Bytes(), nil +} + +// UnmarshalJSON implements stdjson.Unmarshaler. +func (s *OptAgentRefMultiagent) UnmarshalJSON(data []byte) error { + d := jx.DecodeBytes(data) + return s.Decode(d) +} + +// Encode encodes bool as json. +func (o OptBool) Encode(e *jx.Encoder) { + if !o.Set { + return + } + e.Bool(bool(o.Value)) +} + +// Decode decodes bool from json. +func (o *OptBool) Decode(d *jx.Decoder) error { + if o == nil { + return errors.New("invalid: unable to decode OptBool to nil") + } + o.Set = true + v, err := d.Bool() + if err != nil { + return err + } + o.Value = bool(v) + return nil +} + +// MarshalJSON implements stdjson.Marshaler. +func (s OptBool) MarshalJSON() ([]byte, error) { + e := jx.Encoder{} + s.Encode(&e) + return e.Bytes(), nil +} + +// UnmarshalJSON implements stdjson.Unmarshaler. +func (s *OptBool) UnmarshalJSON(data []byte) error { + d := jx.DecodeBytes(data) + return s.Decode(d) +} + +// Encode encodes CacheCreation as json. +func (o OptCacheCreation) Encode(e *jx.Encoder) { + if !o.Set { + return + } + o.Value.Encode(e) +} + +// Decode decodes CacheCreation from json. +func (o *OptCacheCreation) Decode(d *jx.Decoder) error { + if o == nil { + return errors.New("invalid: unable to decode OptCacheCreation to nil") + } + o.Set = true + if err := o.Value.Decode(d); err != nil { + return err + } + return nil +} + +// MarshalJSON implements stdjson.Marshaler. +func (s OptCacheCreation) MarshalJSON() ([]byte, error) { + e := jx.Encoder{} + s.Encode(&e) + return e.Bytes(), nil +} + +// UnmarshalJSON implements stdjson.Unmarshaler. +func (s *OptCacheCreation) UnmarshalJSON(data []byte) error { + d := jx.DecodeBytes(data) + return s.Decode(d) +} + +// Encode encodes EnvironmentConfigOverride as json. +func (o OptEnvironmentConfigOverride) Encode(e *jx.Encoder) { + if !o.Set { + return } + o.Value.Encode(e) +} +// Decode decodes EnvironmentConfigOverride from json. +func (o *OptEnvironmentConfigOverride) Decode(d *jx.Decoder) error { + if o == nil { + return errors.New("invalid: unable to decode OptEnvironmentConfigOverride to nil") + } + o.Set = true + if err := o.Value.Decode(d); err != nil { + return err + } return nil } // MarshalJSON implements stdjson.Marshaler. -func (s *ManagedAgentsUserToolConfirmationEventParams) MarshalJSON() ([]byte, error) { +func (s OptEnvironmentConfigOverride) MarshalJSON() ([]byte, error) { e := jx.Encoder{} s.Encode(&e) return e.Bytes(), nil } // UnmarshalJSON implements stdjson.Unmarshaler. -func (s *ManagedAgentsUserToolConfirmationEventParams) UnmarshalJSON(data []byte) error { +func (s *OptEnvironmentConfigOverride) UnmarshalJSON(data []byte) error { d := jx.DecodeBytes(data) return s.Decode(d) } -// Encode implements json.Marshaler. -func (s *ManagedAgentsUserToolResultEventParams) Encode(e *jx.Encoder) { - e.ObjStart() - s.encodeFields(e) - e.ObjEnd() +// Encode encodes EnvironmentConfigOverrideEnv as json. +func (o OptEnvironmentConfigOverrideEnv) Encode(e *jx.Encoder) { + if !o.Set { + return + } + o.Value.Encode(e) } -// encodeFields encodes fields. -func (s *ManagedAgentsUserToolResultEventParams) encodeFields(e *jx.Encoder) { - { - e.FieldStart("tool_use_id") - e.Str(s.ToolUseID) +// Decode decodes EnvironmentConfigOverrideEnv from json. +func (o *OptEnvironmentConfigOverrideEnv) Decode(d *jx.Decoder) error { + if o == nil { + return errors.New("invalid: unable to decode OptEnvironmentConfigOverrideEnv to nil") } - { - if s.Content != nil { - e.FieldStart("content") - e.ArrStart() - for _, elem := range s.Content { - elem.Encode(e) - } - e.ArrEnd() - } + o.Set = true + o.Value = make(EnvironmentConfigOverrideEnv) + if err := o.Value.Decode(d); err != nil { + return err } - { - if s.IsError.Set { - e.FieldStart("is_error") - s.IsError.Encode(e) - } + return nil +} + +// MarshalJSON implements stdjson.Marshaler. +func (s OptEnvironmentConfigOverrideEnv) MarshalJSON() ([]byte, error) { + e := jx.Encoder{} + s.Encode(&e) + return e.Bytes(), nil +} + +// UnmarshalJSON implements stdjson.Unmarshaler. +func (s *OptEnvironmentConfigOverrideEnv) UnmarshalJSON(data []byte) error { + d := jx.DecodeBytes(data) + return s.Decode(d) +} + +// Encode encodes EnvironmentNetworkingConfig as json. +func (o OptEnvironmentNetworkingConfig) Encode(e *jx.Encoder) { + if !o.Set { + return } - { - if s.SessionID.Set { - e.FieldStart("session_id") - s.SessionID.Encode(e) - } + o.Value.Encode(e) +} + +// Decode decodes EnvironmentNetworkingConfig from json. +func (o *OptEnvironmentNetworkingConfig) Decode(d *jx.Decoder) error { + if o == nil { + return errors.New("invalid: unable to decode OptEnvironmentNetworkingConfig to nil") } - { - if s.SessionThreadID.Set { - e.FieldStart("session_thread_id") - s.SessionThreadID.Encode(e) - } + o.Set = true + if err := o.Value.Decode(d); err != nil { + return err } + return nil } -var jsonFieldsNameOfManagedAgentsUserToolResultEventParams = [5]string{ - 0: "tool_use_id", - 1: "content", - 2: "is_error", - 3: "session_id", - 4: "session_thread_id", +// MarshalJSON implements stdjson.Marshaler. +func (s OptEnvironmentNetworkingConfig) MarshalJSON() ([]byte, error) { + e := jx.Encoder{} + s.Encode(&e) + return e.Bytes(), nil } -// Decode decodes ManagedAgentsUserToolResultEventParams from json. -func (s *ManagedAgentsUserToolResultEventParams) Decode(d *jx.Decoder) error { - if s == nil { - return errors.New("invalid: unable to decode ManagedAgentsUserToolResultEventParams to nil") +// UnmarshalJSON implements stdjson.Unmarshaler. +func (s *OptEnvironmentNetworkingConfig) UnmarshalJSON(data []byte) error { + d := jx.DecodeBytes(data) + return s.Decode(d) +} + +// Encode encodes EnvironmentPackagesConfig as json. +func (o OptEnvironmentPackagesConfig) Encode(e *jx.Encoder) { + if !o.Set { + return } - var requiredBitSet [1]uint8 + o.Value.Encode(e) +} - if err := d.ObjBytes(func(d *jx.Decoder, k []byte) error { - switch string(k) { - case "tool_use_id": - requiredBitSet[0] |= 1 << 0 - if err := func() error { - v, err := d.Str() - s.ToolUseID = string(v) - if err != nil { - return err - } - return nil - }(); err != nil { - return errors.Wrap(err, "decode field \"tool_use_id\"") - } - case "content": - if err := func() error { - s.Content = make([]ManagedAgentsToolResultContentBlock, 0) - if err := d.Arr(func(d *jx.Decoder) error { - var elem ManagedAgentsToolResultContentBlock - if err := elem.Decode(d); err != nil { - return err - } - s.Content = append(s.Content, elem) - return nil - }); err != nil { - return err - } - return nil - }(); err != nil { - return errors.Wrap(err, "decode field \"content\"") - } - case "is_error": - if err := func() error { - s.IsError.Reset() - if err := s.IsError.Decode(d); err != nil { - return err - } - return nil - }(); err != nil { - return errors.Wrap(err, "decode field \"is_error\"") - } - case "session_id": - if err := func() error { - s.SessionID.Reset() - if err := s.SessionID.Decode(d); err != nil { - return err - } - return nil - }(); err != nil { - return errors.Wrap(err, "decode field \"session_id\"") - } - case "session_thread_id": - if err := func() error { - s.SessionThreadID.Reset() - if err := s.SessionThreadID.Decode(d); err != nil { - return err - } - return nil - }(); err != nil { - return errors.Wrap(err, "decode field \"session_thread_id\"") - } - default: - return d.Skip() - } - return nil - }); err != nil { - return errors.Wrap(err, "decode ManagedAgentsUserToolResultEventParams") +// Decode decodes EnvironmentPackagesConfig from json. +func (o *OptEnvironmentPackagesConfig) Decode(d *jx.Decoder) error { + if o == nil { + return errors.New("invalid: unable to decode OptEnvironmentPackagesConfig to nil") } - // Validate required fields. - var failures []validate.FieldError - for i, mask := range [1]uint8{ - 0b00000001, - } { - if result := (requiredBitSet[i] & mask) ^ mask; result != 0 { - // Mask only required fields and check equality to mask using XOR. - // - // If XOR result is not zero, result is not equal to expected, so some fields are missed. - // Bits of fields which would be set are actually bits of missed fields. - missed := bits.OnesCount8(result) - for bitN := 0; bitN < missed; bitN++ { - bitIdx := bits.TrailingZeros8(result) - fieldIdx := i*8 + bitIdx - var name string - if fieldIdx < len(jsonFieldsNameOfManagedAgentsUserToolResultEventParams) { - name = jsonFieldsNameOfManagedAgentsUserToolResultEventParams[fieldIdx] - } else { - name = strconv.Itoa(fieldIdx) - } - failures = append(failures, validate.FieldError{ - Name: name, - Error: validate.ErrFieldRequired, - }) - // Reset bit. - result &^= 1 << bitIdx - } - } + o.Set = true + if err := o.Value.Decode(d); err != nil { + return err } - if len(failures) > 0 { - return &validate.Error{Fields: failures} + return nil +} + +// MarshalJSON implements stdjson.Marshaler. +func (s OptEnvironmentPackagesConfig) MarshalJSON() ([]byte, error) { + e := jx.Encoder{} + s.Encode(&e) + return e.Bytes(), nil +} + +// UnmarshalJSON implements stdjson.Unmarshaler. +func (s *OptEnvironmentPackagesConfig) UnmarshalJSON(data []byte) error { + d := jx.DecodeBytes(data) + return s.Decode(d) +} + +// Encode encodes EnvironmentPackagesConfigType as json. +func (o OptEnvironmentPackagesConfigType) Encode(e *jx.Encoder) { + if !o.Set { + return } + e.Str(string(o.Value)) +} +// Decode decodes EnvironmentPackagesConfigType from json. +func (o *OptEnvironmentPackagesConfigType) Decode(d *jx.Decoder) error { + if o == nil { + return errors.New("invalid: unable to decode OptEnvironmentPackagesConfigType to nil") + } + o.Set = true + if err := o.Value.Decode(d); err != nil { + return err + } return nil } // MarshalJSON implements stdjson.Marshaler. -func (s *ManagedAgentsUserToolResultEventParams) MarshalJSON() ([]byte, error) { +func (s OptEnvironmentPackagesConfigType) MarshalJSON() ([]byte, error) { e := jx.Encoder{} s.Encode(&e) return e.Bytes(), nil } // UnmarshalJSON implements stdjson.Unmarshaler. -func (s *ManagedAgentsUserToolResultEventParams) UnmarshalJSON(data []byte) error { +func (s *OptEnvironmentPackagesConfigType) UnmarshalJSON(data []byte) error { d := jx.DecodeBytes(data) return s.Decode(d) } -// Encode encodes bool as json. -func (o OptBool) Encode(e *jx.Encoder) { +// Encode encodes EnvironmentTosConfig as json. +func (o OptEnvironmentTosConfig) Encode(e *jx.Encoder) { if !o.Set { return } - e.Bool(bool(o.Value)) + o.Value.Encode(e) } -// Decode decodes bool from json. -func (o *OptBool) Decode(d *jx.Decoder) error { +// Decode decodes EnvironmentTosConfig from json. +func (o *OptEnvironmentTosConfig) Decode(d *jx.Decoder) error { if o == nil { - return errors.New("invalid: unable to decode OptBool to nil") + return errors.New("invalid: unable to decode OptEnvironmentTosConfig to nil") } o.Set = true - v, err := d.Bool() - if err != nil { + if err := o.Value.Decode(d); err != nil { return err } - o.Value = bool(v) return nil } // MarshalJSON implements stdjson.Marshaler. -func (s OptBool) MarshalJSON() ([]byte, error) { +func (s OptEnvironmentTosConfig) MarshalJSON() ([]byte, error) { e := jx.Encoder{} s.Encode(&e) return e.Bytes(), nil } // UnmarshalJSON implements stdjson.Unmarshaler. -func (s *OptBool) UnmarshalJSON(data []byte) error { +func (s *OptEnvironmentTosConfig) UnmarshalJSON(data []byte) error { d := jx.DecodeBytes(data) return s.Decode(d) } -// Encode encodes CacheCreation as json. -func (o OptCacheCreation) Encode(e *jx.Encoder) { +// Encode encodes EnvironmentWithOverrides as json. +func (o OptEnvironmentWithOverrides) Encode(e *jx.Encoder) { if !o.Set { return } o.Value.Encode(e) } -// Decode decodes CacheCreation from json. -func (o *OptCacheCreation) Decode(d *jx.Decoder) error { +// Decode decodes EnvironmentWithOverrides from json. +func (o *OptEnvironmentWithOverrides) Decode(d *jx.Decoder) error { if o == nil { - return errors.New("invalid: unable to decode OptCacheCreation to nil") + return errors.New("invalid: unable to decode OptEnvironmentWithOverrides to nil") } o.Set = true if err := o.Value.Decode(d); err != nil { @@ -4879,14 +6532,14 @@ func (o *OptCacheCreation) Decode(d *jx.Decoder) error { } // MarshalJSON implements stdjson.Marshaler. -func (s OptCacheCreation) MarshalJSON() ([]byte, error) { +func (s OptEnvironmentWithOverrides) MarshalJSON() ([]byte, error) { e := jx.Encoder{} s.Encode(&e) return e.Bytes(), nil } // UnmarshalJSON implements stdjson.Unmarshaler. -func (s *OptCacheCreation) UnmarshalJSON(data []byte) error { +func (s *OptEnvironmentWithOverrides) UnmarshalJSON(data []byte) error { d := jx.DecodeBytes(data) return s.Decode(d) } @@ -4961,6 +6614,89 @@ func (s *OptInt64) UnmarshalJSON(data []byte) error { return s.Decode(d) } +// Encode encodes ModelOverrides as json. +func (o OptModelOverrides) Encode(e *jx.Encoder) { + if !o.Set { + return + } + o.Value.Encode(e) +} + +// Decode decodes ModelOverrides from json. +func (o *OptModelOverrides) Decode(d *jx.Decoder) error { + if o == nil { + return errors.New("invalid: unable to decode OptModelOverrides to nil") + } + o.Set = true + if err := o.Value.Decode(d); err != nil { + return err + } + return nil +} + +// MarshalJSON implements stdjson.Marshaler. +func (s OptModelOverrides) MarshalJSON() ([]byte, error) { + e := jx.Encoder{} + s.Encode(&e) + return e.Bytes(), nil +} + +// UnmarshalJSON implements stdjson.Unmarshaler. +func (s *OptModelOverrides) UnmarshalJSON(data []byte) error { + d := jx.DecodeBytes(data) + return s.Decode(d) +} + +// Encode encodes SessionEnvironment as json. +func (o OptNilSessionEnvironment) Encode(e *jx.Encoder) { + if !o.Set { + return + } + if o.Null { + e.Null() + return + } + o.Value.Encode(e) +} + +// Decode decodes SessionEnvironment from json. +func (o *OptNilSessionEnvironment) Decode(d *jx.Decoder) error { + if o == nil { + return errors.New("invalid: unable to decode OptNilSessionEnvironment to nil") + } + if d.Next() == jx.Null { + if err := d.Null(); err != nil { + return err + } + + var v SessionEnvironment + o.Value = v + o.Set = true + o.Null = true + return nil + } + o.Set = true + o.Null = false + o.Value = make(SessionEnvironment) + if err := o.Value.Decode(d); err != nil { + return err + } + return nil +} + +// MarshalJSON implements stdjson.Marshaler. +func (s OptNilSessionEnvironment) MarshalJSON() ([]byte, error) { + e := jx.Encoder{} + s.Encode(&e) + return e.Bytes(), nil +} + +// UnmarshalJSON implements stdjson.Unmarshaler. +func (s *OptNilSessionEnvironment) UnmarshalJSON(data []byte) error { + d := jx.DecodeBytes(data) + return s.Decode(d) +} + // Encode encodes []SessionResource as json. func (o OptNilSessionResourceArray) Encode(e *jx.Encoder) { if !o.Set { @@ -5735,15 +7471,17 @@ func (s *SendSessionEventsResponse) Encode(e *jx.Encoder) { // encodeFields encodes fields. func (s *SendSessionEventsResponse) encodeFields(e *jx.Encoder) { { - if s.Success.Set { - e.FieldStart("success") - s.Success.Encode(e) + e.FieldStart("data") + e.ArrStart() + for _, elem := range s.Data { + elem.Encode(e) } + e.ArrEnd() } } var jsonFieldsNameOfSendSessionEventsResponse = [1]string{ - 0: "success", + 0: "data", } // Decode decodes SendSessionEventsResponse from json. @@ -5751,18 +7489,27 @@ func (s *SendSessionEventsResponse) Decode(d *jx.Decoder) error { if s == nil { return errors.New("invalid: unable to decode SendSessionEventsResponse to nil") } + var requiredBitSet [1]uint8 if err := d.ObjBytes(func(d *jx.Decoder, k []byte) error { switch string(k) { - case "success": + case "data": + requiredBitSet[0] |= 1 << 0 if err := func() error { - s.Success.Reset() - if err := s.Success.Decode(d); err != nil { + s.Data = make([]SendSessionEventsResponseDataItem, 0) + if err := d.Arr(func(d *jx.Decoder) error { + var elem SendSessionEventsResponseDataItem + if err := elem.Decode(d); err != nil { + return err + } + s.Data = append(s.Data, elem) + return nil + }); err != nil { return err } return nil }(); err != nil { - return errors.Wrap(err, "decode field \"success\"") + return errors.Wrap(err, "decode field \"data\"") } default: return d.Skip() @@ -5771,6 +7518,38 @@ func (s *SendSessionEventsResponse) Decode(d *jx.Decoder) error { }); err != nil { return errors.Wrap(err, "decode SendSessionEventsResponse") } + // Validate required fields. + var failures []validate.FieldError + for i, mask := range [1]uint8{ + 0b00000001, + } { + if result := (requiredBitSet[i] & mask) ^ mask; result != 0 { + // Mask only required fields and check equality to mask using XOR. + // + // If XOR result is not zero, result is not equal to expected, so some fields are missed. + // Bits of fields which would be set are actually bits of missed fields. + missed := bits.OnesCount8(result) + for bitN := 0; bitN < missed; bitN++ { + bitIdx := bits.TrailingZeros8(result) + fieldIdx := i*8 + bitIdx + var name string + if fieldIdx < len(jsonFieldsNameOfSendSessionEventsResponse) { + name = jsonFieldsNameOfSendSessionEventsResponse[fieldIdx] + } else { + name = strconv.Itoa(fieldIdx) + } + failures = append(failures, validate.FieldError{ + Name: name, + Error: validate.ErrFieldRequired, + }) + // Reset bit. + result &^= 1 << bitIdx + } + } + } + if len(failures) > 0 { + return &validate.Error{Fields: failures} + } return nil } @@ -5788,6 +7567,64 @@ func (s *SendSessionEventsResponse) UnmarshalJSON(data []byte) error { return s.Decode(d) } +// Encode implements json.Marshaler. +func (s SendSessionEventsResponseDataItem) Encode(e *jx.Encoder) { + e.ObjStart() + s.encodeFields(e) + e.ObjEnd() +} + +// encodeFields implements json.Marshaler. +func (s SendSessionEventsResponseDataItem) encodeFields(e *jx.Encoder) { + for k, elem := range s { + e.FieldStart(k) + + if len(elem) != 0 { + e.Raw(elem) + } + } +} + +// Decode decodes SendSessionEventsResponseDataItem from json. +func (s *SendSessionEventsResponseDataItem) Decode(d *jx.Decoder) error { + if s == nil { + return errors.New("invalid: unable to decode SendSessionEventsResponseDataItem to nil") + } + m := s.init() + if err := d.ObjBytes(func(d *jx.Decoder, k []byte) error { + var elem jx.Raw + if err := func() error { + v, err := d.RawAppend(nil) + elem = jx.Raw(v) + if err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrapf(err, "decode field %q", k) + } + m[string(k)] = elem + return nil + }); err != nil { + return errors.Wrap(err, "decode SendSessionEventsResponseDataItem") + } + + return nil +} + +// MarshalJSON implements stdjson.Marshaler. +func (s SendSessionEventsResponseDataItem) MarshalJSON() ([]byte, error) { + e := jx.Encoder{} + s.Encode(&e) + return e.Bytes(), nil +} + +// UnmarshalJSON implements stdjson.Unmarshaler. +func (s *SendSessionEventsResponseDataItem) UnmarshalJSON(data []byte) error { + d := jx.DecodeBytes(data) + return s.Decode(d) +} + // Encode implements json.Marshaler. func (s *Session) Encode(e *jx.Encoder) { e.ObjStart() @@ -5861,9 +7698,15 @@ func (s *Session) encodeFields(e *jx.Encoder) { s.Tags.Encode(e) } } + { + if s.Environment.Set { + e.FieldStart("environment") + s.Environment.Encode(e) + } + } } -var jsonFieldsNameOfSession = [13]string{ +var jsonFieldsNameOfSession = [14]string{ 0: "id", 1: "type", 2: "status", @@ -5877,6 +7720,7 @@ var jsonFieldsNameOfSession = [13]string{ 10: "vault_ids", 11: "usage", 12: "tags", + 13: "environment", } // Decode decodes Session from json. @@ -6026,6 +7870,16 @@ func (s *Session) Decode(d *jx.Decoder) error { }(); err != nil { return errors.Wrap(err, "decode field \"tags\"") } + case "environment": + if err := func() error { + s.Environment.Reset() + if err := s.Environment.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"environment\"") + } default: return d.Skip() } @@ -6141,6 +7995,64 @@ func (s *SessionAgent) UnmarshalJSON(data []byte) error { return s.Decode(d) } +// Encode implements json.Marshaler. +func (s SessionEnvironment) Encode(e *jx.Encoder) { + e.ObjStart() + s.encodeFields(e) + e.ObjEnd() +} + +// encodeFields implements json.Marshaler. +func (s SessionEnvironment) encodeFields(e *jx.Encoder) { + for k, elem := range s { + e.FieldStart(k) + + if len(elem) != 0 { + e.Raw(elem) + } + } +} + +// Decode decodes SessionEnvironment from json. +func (s *SessionEnvironment) Decode(d *jx.Decoder) error { + if s == nil { + return errors.New("invalid: unable to decode SessionEnvironment to nil") + } + m := s.init() + if err := d.ObjBytes(func(d *jx.Decoder, k []byte) error { + var elem jx.Raw + if err := func() error { + v, err := d.RawAppend(nil) + elem = jx.Raw(v) + if err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrapf(err, "decode field %q", k) + } + m[string(k)] = elem + return nil + }); err != nil { + return errors.Wrap(err, "decode SessionEnvironment") + } + + return nil +} + +// MarshalJSON implements stdjson.Marshaler. +func (s SessionEnvironment) MarshalJSON() ([]byte, error) { + e := jx.Encoder{} + s.Encode(&e) + return e.Bytes(), nil +} + +// UnmarshalJSON implements stdjson.Unmarshaler. +func (s *SessionEnvironment) UnmarshalJSON(data []byte) error { + d := jx.DecodeBytes(data) + return s.Decode(d) +} + // Encode implements json.Marshaler. func (s *SessionResource) Encode(e *jx.Encoder) { e.ObjStart() @@ -6765,8 +8677,8 @@ func (s *SessionThreadStatus) Decode(d *jx.Decoder) error { *s = SessionThreadStatusRunning case SessionThreadStatusTerminated: *s = SessionThreadStatusTerminated - case SessionThreadStatusArchived: - *s = SessionThreadStatusArchived + case SessionThreadStatusRescheduling: + *s = SessionThreadStatusRescheduling default: *s = SessionThreadStatus(v) } diff --git a/arkruntime/model/session/oas_schemas_gen.go b/arkruntime/model/session/oas_schemas_gen.go index 39d480c..2b3c747 100644 --- a/arkruntime/model/session/oas_schemas_gen.go +++ b/arkruntime/model/session/oas_schemas_gen.go @@ -77,25 +77,39 @@ func NewAgentRefAgentIdentifier(v AgentRef) AgentIdentifier { return s } -// Agent 引用(对象形态):`type: "agent"` + id + optional version。 -// 与 CreateSessionRequest.agent 联合使用。. +// Agent 引用(对象形态):`type: "agent"` 或 `"agent_with_overrides"`。 +// MA wire 上两种对象形态都走同一个 JSON object 承载,避免 SDK 生成复杂 union。. // Ref: #/components/schemas/AgentRef type AgentRef struct { - // 固定 `"agent"`。. - Type AgentRefType `json:"type"` + // `"agent"` 或 `"agent_with_overrides"`。. + Type string `json:"type"` // Agent ID。. - ID string `json:"id"` + ID OptString `json:"id"` // Agent 版本号;不传走最新。. Version OptInt32 `json:"version"` + // System prompt 覆写。. + System OptString `json:"system"` + // 工具配置覆写。. + Tools []AgentRefToolsItem `json:"tools"` + // MCP server 配置覆写。. + McpServers []AgentRefMcpServersItem `json:"mcp_servers"` + // Skill 配置覆写。. + Skills []AgentRefSkillsItem `json:"skills"` + // 多 Agent 配置覆写。. + Multiagent OptAgentRefMultiagent `json:"multiagent"` + // Session 响应中冻结的 Agent 展示名。. + DisplayName OptString `json:"display_name"` + // 模型运行参数覆写。. + Model OptModelOverrides `json:"model"` } // GetType returns the value of Type. -func (s *AgentRef) GetType() AgentRefType { +func (s *AgentRef) GetType() string { return s.Type } // GetID returns the value of ID. -func (s *AgentRef) GetID() string { +func (s *AgentRef) GetID() OptString { return s.ID } @@ -104,13 +118,48 @@ func (s *AgentRef) GetVersion() OptInt32 { return s.Version } +// GetSystem returns the value of System. +func (s *AgentRef) GetSystem() OptString { + return s.System +} + +// GetTools returns the value of Tools. +func (s *AgentRef) GetTools() []AgentRefToolsItem { + return s.Tools +} + +// GetMcpServers returns the value of McpServers. +func (s *AgentRef) GetMcpServers() []AgentRefMcpServersItem { + return s.McpServers +} + +// GetSkills returns the value of Skills. +func (s *AgentRef) GetSkills() []AgentRefSkillsItem { + return s.Skills +} + +// GetMultiagent returns the value of Multiagent. +func (s *AgentRef) GetMultiagent() OptAgentRefMultiagent { + return s.Multiagent +} + +// GetDisplayName returns the value of DisplayName. +func (s *AgentRef) GetDisplayName() OptString { + return s.DisplayName +} + +// GetModel returns the value of Model. +func (s *AgentRef) GetModel() OptModelOverrides { + return s.Model +} + // SetType sets the value of Type. -func (s *AgentRef) SetType(val AgentRefType) { +func (s *AgentRef) SetType(val string) { s.Type = val } // SetID sets the value of ID. -func (s *AgentRef) SetID(val string) { +func (s *AgentRef) SetID(val OptString) { s.ID = val } @@ -119,39 +168,84 @@ func (s *AgentRef) SetVersion(val OptInt32) { s.Version = val } -// 固定 `"agent"`。. -type AgentRefType string +// SetSystem sets the value of System. +func (s *AgentRef) SetSystem(val OptString) { + s.System = val +} -const ( - AgentRefTypeAgent AgentRefType = "agent" -) +// SetTools sets the value of Tools. +func (s *AgentRef) SetTools(val []AgentRefToolsItem) { + s.Tools = val +} + +// SetMcpServers sets the value of McpServers. +func (s *AgentRef) SetMcpServers(val []AgentRefMcpServersItem) { + s.McpServers = val +} + +// SetSkills sets the value of Skills. +func (s *AgentRef) SetSkills(val []AgentRefSkillsItem) { + s.Skills = val +} + +// SetMultiagent sets the value of Multiagent. +func (s *AgentRef) SetMultiagent(val OptAgentRefMultiagent) { + s.Multiagent = val +} + +// SetDisplayName sets the value of DisplayName. +func (s *AgentRef) SetDisplayName(val OptString) { + s.DisplayName = val +} + +// SetModel sets the value of Model. +func (s *AgentRef) SetModel(val OptModelOverrides) { + s.Model = val +} + +type AgentRefMcpServersItem map[string]jx.Raw -// AllValues returns all AgentRefType values. -func (AgentRefType) AllValues() []AgentRefType { - return []AgentRefType{ - AgentRefTypeAgent, +func (s *AgentRefMcpServersItem) init() AgentRefMcpServersItem { + m := *s + if m == nil { + m = map[string]jx.Raw{} + *s = m } + return m } -// MarshalText implements encoding.TextMarshaler. -func (s AgentRefType) MarshalText() ([]byte, error) { - switch s { - case AgentRefTypeAgent: - return []byte(s), nil - default: - return nil, errors.Errorf("invalid value: %q", s) +// 多 Agent 配置覆写。. +type AgentRefMultiagent map[string]jx.Raw + +func (s *AgentRefMultiagent) init() AgentRefMultiagent { + m := *s + if m == nil { + m = map[string]jx.Raw{} + *s = m + } + return m +} + +type AgentRefSkillsItem map[string]jx.Raw + +func (s *AgentRefSkillsItem) init() AgentRefSkillsItem { + m := *s + if m == nil { + m = map[string]jx.Raw{} + *s = m } + return m } -// UnmarshalText implements encoding.TextUnmarshaler. -func (s *AgentRefType) UnmarshalText(data []byte) error { - switch AgentRefType(data) { - case AgentRefTypeAgent: - *s = AgentRefTypeAgent - return nil - default: - return errors.Errorf("invalid value: %q", data) +type AgentRefToolsItem map[string]jx.Raw + +func (s *AgentRefToolsItem) init() AgentRefToolsItem { + m := *s + if m == nil { + m = map[string]jx.Raw{} + *s = m } + return m } // 直接以 base64 携带的文档数据。. @@ -249,8 +343,10 @@ func (s *CacheCreation) SetEphemeral5mInputTokens(val OptInt64) { type CreateSessionRequest struct { // Agent 标识。. Agent AgentIdentifier `json:"agent"` - // 关联的 Environment ID。. - EnvironmentID string `json:"environment_id"` + // 关联的 Environment ID。与 `environment` 二选一。. + EnvironmentID OptString `json:"environment_id"` + // 关联 Environment 的覆写引用。与 `environment_id` 二选一。. + Environment OptEnvironmentWithOverrides `json:"environment"` // 资源标签。. Tags []Tag `json:"tags"` // 挂载资源列表。. @@ -267,10 +363,15 @@ func (s *CreateSessionRequest) GetAgent() AgentIdentifier { } // GetEnvironmentID returns the value of EnvironmentID. -func (s *CreateSessionRequest) GetEnvironmentID() string { +func (s *CreateSessionRequest) GetEnvironmentID() OptString { return s.EnvironmentID } +// GetEnvironment returns the value of Environment. +func (s *CreateSessionRequest) GetEnvironment() OptEnvironmentWithOverrides { + return s.Environment +} + // GetTags returns the value of Tags. func (s *CreateSessionRequest) GetTags() []Tag { return s.Tags @@ -297,10 +398,15 @@ func (s *CreateSessionRequest) SetAgent(val AgentIdentifier) { } // SetEnvironmentID sets the value of EnvironmentID. -func (s *CreateSessionRequest) SetEnvironmentID(val string) { +func (s *CreateSessionRequest) SetEnvironmentID(val OptString) { s.EnvironmentID = val } +// SetEnvironment sets the value of Environment. +func (s *CreateSessionRequest) SetEnvironment(val OptEnvironmentWithOverrides) { + s.Environment = val +} + // SetTags sets the value of Tags. func (s *CreateSessionRequest) SetTags(val []Tag) { s.Tags = val @@ -557,6 +663,377 @@ func NewFileDocumentSourceDocumentSourceSum(v FileDocumentSource) DocumentSource return s } +// Environment 覆写时使用的运行环境配置。. +// Ref: #/components/schemas/EnvironmentConfigOverride +type EnvironmentConfigOverride struct { + // 运行环境类型。. + Type string `json:"type"` + // 容器出网策略。. + Networking OptEnvironmentNetworkingConfig `json:"networking"` + // 启动时预装的依赖包。. + Packages OptEnvironmentPackagesConfig `json:"packages"` + // 容器环境变量。. + Env OptEnvironmentConfigOverrideEnv `json:"env"` + // 沙箱启动脚本。. + SetupScript OptString `json:"setup_script"` + // Environment outputs 的 TOS 存储配置。. + Tos OptEnvironmentTosConfig `json:"tos"` +} + +// GetType returns the value of Type. +func (s *EnvironmentConfigOverride) GetType() string { + return s.Type +} + +// GetNetworking returns the value of Networking. +func (s *EnvironmentConfigOverride) GetNetworking() OptEnvironmentNetworkingConfig { + return s.Networking +} + +// GetPackages returns the value of Packages. +func (s *EnvironmentConfigOverride) GetPackages() OptEnvironmentPackagesConfig { + return s.Packages +} + +// GetEnv returns the value of Env. +func (s *EnvironmentConfigOverride) GetEnv() OptEnvironmentConfigOverrideEnv { + return s.Env +} + +// GetSetupScript returns the value of SetupScript. +func (s *EnvironmentConfigOverride) GetSetupScript() OptString { + return s.SetupScript +} + +// GetTos returns the value of Tos. +func (s *EnvironmentConfigOverride) GetTos() OptEnvironmentTosConfig { + return s.Tos +} + +// SetType sets the value of Type. +func (s *EnvironmentConfigOverride) SetType(val string) { + s.Type = val +} + +// SetNetworking sets the value of Networking. +func (s *EnvironmentConfigOverride) SetNetworking(val OptEnvironmentNetworkingConfig) { + s.Networking = val +} + +// SetPackages sets the value of Packages. +func (s *EnvironmentConfigOverride) SetPackages(val OptEnvironmentPackagesConfig) { + s.Packages = val +} + +// SetEnv sets the value of Env. +func (s *EnvironmentConfigOverride) SetEnv(val OptEnvironmentConfigOverrideEnv) { + s.Env = val +} + +// SetSetupScript sets the value of SetupScript. +func (s *EnvironmentConfigOverride) SetSetupScript(val OptString) { + s.SetupScript = val +} + +// SetTos sets the value of Tos. +func (s *EnvironmentConfigOverride) SetTos(val OptEnvironmentTosConfig) { + s.Tos = val +} + +// 容器环境变量。. +type EnvironmentConfigOverrideEnv map[string]string + +func (s *EnvironmentConfigOverrideEnv) init() EnvironmentConfigOverrideEnv { + m := *s + if m == nil { + m = map[string]string{} + *s = m + } + return m +} + +// Environment 覆写时使用的容器出网策略。. +// Ref: #/components/schemas/EnvironmentNetworkingConfig +type EnvironmentNetworkingConfig struct { + // 出网策略类型。. + Type string `json:"type"` + // 是否允许出网到 MCP servers。. + AllowMcpServers OptBool `json:"allow_mcp_servers"` + // 是否允许访问包管理器。. + AllowPackageManagers OptBool `json:"allow_package_managers"` + // 显式允许的出网域名。. + AllowedHosts []string `json:"allowed_hosts"` +} + +// GetType returns the value of Type. +func (s *EnvironmentNetworkingConfig) GetType() string { + return s.Type +} + +// GetAllowMcpServers returns the value of AllowMcpServers. +func (s *EnvironmentNetworkingConfig) GetAllowMcpServers() OptBool { + return s.AllowMcpServers +} + +// GetAllowPackageManagers returns the value of AllowPackageManagers. +func (s *EnvironmentNetworkingConfig) GetAllowPackageManagers() OptBool { + return s.AllowPackageManagers +} + +// GetAllowedHosts returns the value of AllowedHosts. +func (s *EnvironmentNetworkingConfig) GetAllowedHosts() []string { + return s.AllowedHosts +} + +// SetType sets the value of Type. +func (s *EnvironmentNetworkingConfig) SetType(val string) { + s.Type = val +} + +// SetAllowMcpServers sets the value of AllowMcpServers. +func (s *EnvironmentNetworkingConfig) SetAllowMcpServers(val OptBool) { + s.AllowMcpServers = val +} + +// SetAllowPackageManagers sets the value of AllowPackageManagers. +func (s *EnvironmentNetworkingConfig) SetAllowPackageManagers(val OptBool) { + s.AllowPackageManagers = val +} + +// SetAllowedHosts sets the value of AllowedHosts. +func (s *EnvironmentNetworkingConfig) SetAllowedHosts(val []string) { + s.AllowedHosts = val +} + +// Environment 覆写时使用的预装依赖包配置。. +// Ref: #/components/schemas/EnvironmentPackagesConfig +type EnvironmentPackagesConfig struct { + // 固定 `"packages"`。. + Type OptEnvironmentPackagesConfigType `json:"type"` + // Pip 依赖。. + Pip []string `json:"pip"` + // Apt 依赖。. + Apt []string `json:"apt"` + // Npm 依赖。. + Npm []string `json:"npm"` + // Cargo 依赖。. + Cargo []string `json:"cargo"` + // Gem 依赖。. + Gem []string `json:"gem"` + // Go module 依赖。. + Go []string `json:"go"` +} + +// GetType returns the value of Type. +func (s *EnvironmentPackagesConfig) GetType() OptEnvironmentPackagesConfigType { + return s.Type +} + +// GetPip returns the value of Pip. +func (s *EnvironmentPackagesConfig) GetPip() []string { + return s.Pip +} + +// GetApt returns the value of Apt. +func (s *EnvironmentPackagesConfig) GetApt() []string { + return s.Apt +} + +// GetNpm returns the value of Npm. +func (s *EnvironmentPackagesConfig) GetNpm() []string { + return s.Npm +} + +// GetCargo returns the value of Cargo. +func (s *EnvironmentPackagesConfig) GetCargo() []string { + return s.Cargo +} + +// GetGem returns the value of Gem. +func (s *EnvironmentPackagesConfig) GetGem() []string { + return s.Gem +} + +// GetGo returns the value of Go. +func (s *EnvironmentPackagesConfig) GetGo() []string { + return s.Go +} + +// SetType sets the value of Type. +func (s *EnvironmentPackagesConfig) SetType(val OptEnvironmentPackagesConfigType) { + s.Type = val +} + +// SetPip sets the value of Pip. +func (s *EnvironmentPackagesConfig) SetPip(val []string) { + s.Pip = val +} + +// SetApt sets the value of Apt. +func (s *EnvironmentPackagesConfig) SetApt(val []string) { + s.Apt = val +} + +// SetNpm sets the value of Npm. +func (s *EnvironmentPackagesConfig) SetNpm(val []string) { + s.Npm = val +} + +// SetCargo sets the value of Cargo. +func (s *EnvironmentPackagesConfig) SetCargo(val []string) { + s.Cargo = val +} + +// SetGem sets the value of Gem. +func (s *EnvironmentPackagesConfig) SetGem(val []string) { + s.Gem = val +} + +// SetGo sets the value of Go. +func (s *EnvironmentPackagesConfig) SetGo(val []string) { + s.Go = val +} + +// 固定 `"packages"`。. +type EnvironmentPackagesConfigType string + +const ( + EnvironmentPackagesConfigTypePackages EnvironmentPackagesConfigType = "packages" +) + +// AllValues returns all EnvironmentPackagesConfigType values. +func (EnvironmentPackagesConfigType) AllValues() []EnvironmentPackagesConfigType { + return []EnvironmentPackagesConfigType{ + EnvironmentPackagesConfigTypePackages, + } +} + +// MarshalText implements encoding.TextMarshaler. +func (s EnvironmentPackagesConfigType) MarshalText() ([]byte, error) { + switch s { + case EnvironmentPackagesConfigTypePackages: + return []byte(s), nil + default: + return nil, errors.Errorf("invalid value: %q", s) + } +} + +// UnmarshalText implements encoding.TextUnmarshaler. +func (s *EnvironmentPackagesConfigType) UnmarshalText(data []byte) error { + switch EnvironmentPackagesConfigType(data) { + case EnvironmentPackagesConfigTypePackages: + *s = EnvironmentPackagesConfigTypePackages + return nil + default: + return errors.Errorf("invalid value: %q", data) + } +} + +// Environment 覆写时使用的 TOS 配置。. +// Ref: #/components/schemas/EnvironmentTosConfig +type EnvironmentTosConfig struct { + // TOS bucket 名称。. + Bucket OptString `json:"bucket"` + // TOS 前缀。. + Prefix OptString `json:"prefix"` +} + +// GetBucket returns the value of Bucket. +func (s *EnvironmentTosConfig) GetBucket() OptString { + return s.Bucket +} + +// GetPrefix returns the value of Prefix. +func (s *EnvironmentTosConfig) GetPrefix() OptString { + return s.Prefix +} + +// SetBucket sets the value of Bucket. +func (s *EnvironmentTosConfig) SetBucket(val OptString) { + s.Bucket = val +} + +// SetPrefix sets the value of Prefix. +func (s *EnvironmentTosConfig) SetPrefix(val OptString) { + s.Prefix = val +} + +// CreateSession 时的 Environment 覆写引用。. +// Ref: #/components/schemas/EnvironmentWithOverrides +type EnvironmentWithOverrides struct { + // 固定 `"environment_with_overrides"`。. + Type EnvironmentWithOverridesType `json:"type"` + // Base Environment ID。. + ID string `json:"id"` + // 运行时配置覆写。. + Config OptEnvironmentConfigOverride `json:"config"` +} + +// GetType returns the value of Type. +func (s *EnvironmentWithOverrides) GetType() EnvironmentWithOverridesType { + return s.Type +} + +// GetID returns the value of ID. +func (s *EnvironmentWithOverrides) GetID() string { + return s.ID +} + +// GetConfig returns the value of Config. +func (s *EnvironmentWithOverrides) GetConfig() OptEnvironmentConfigOverride { + return s.Config +} + +// SetType sets the value of Type. +func (s *EnvironmentWithOverrides) SetType(val EnvironmentWithOverridesType) { + s.Type = val +} + +// SetID sets the value of ID. +func (s *EnvironmentWithOverrides) SetID(val string) { + s.ID = val +} + +// SetConfig sets the value of Config. +func (s *EnvironmentWithOverrides) SetConfig(val OptEnvironmentConfigOverride) { + s.Config = val +} + +// 固定 `"environment_with_overrides"`。. +type EnvironmentWithOverridesType string + +const ( + EnvironmentWithOverridesTypeEnvironmentWithOverrides EnvironmentWithOverridesType = "environment_with_overrides" +) + +// AllValues returns all EnvironmentWithOverridesType values. +func (EnvironmentWithOverridesType) AllValues() []EnvironmentWithOverridesType { + return []EnvironmentWithOverridesType{ + EnvironmentWithOverridesTypeEnvironmentWithOverrides, + } +} + +// MarshalText implements encoding.TextMarshaler. +func (s EnvironmentWithOverridesType) MarshalText() ([]byte, error) { + switch s { + case EnvironmentWithOverridesTypeEnvironmentWithOverrides: + return []byte(s), nil + default: + return nil, errors.Errorf("invalid value: %q", s) + } +} + +// UnmarshalText implements encoding.TextUnmarshaler. +func (s *EnvironmentWithOverridesType) UnmarshalText(data []byte) error { + switch EnvironmentWithOverridesType(data) { + case EnvironmentWithOverridesTypeEnvironmentWithOverrides: + *s = EnvironmentWithOverridesTypeEnvironmentWithOverrides + return nil + default: + return errors.Errorf("invalid value: %q", data) + } +} + // 通过 Files API 上传后引用的文档。. // Ref: #/components/schemas/FileDocumentSource type FileDocumentSource struct { @@ -2040,18 +2517,105 @@ func (s *ManagedAgentsUserToolResultEventParams) SetSessionThreadID(val OptStrin s.SessionThreadID = val } -// NewOptBool returns new OptBool with value set to v. -func NewOptBool(v bool) OptBool { - return OptBool{ - Value: v, - Set: true, - } +// Session 创建时允许临时覆写的模型运行参数。. +// Ref: #/components/schemas/ModelOverrides +type ModelOverrides struct { + // 模型速度档位。. + Speed OptString `json:"speed"` + // Thinking 配置。. + Thinking OptString `json:"thinking"` + // 推理努力程度。. + ReasoningEffort OptString `json:"reasoning_effort"` } -// OptBool is optional bool. -type OptBool struct { - Value bool - Set bool +// GetSpeed returns the value of Speed. +func (s *ModelOverrides) GetSpeed() OptString { + return s.Speed +} + +// GetThinking returns the value of Thinking. +func (s *ModelOverrides) GetThinking() OptString { + return s.Thinking +} + +// GetReasoningEffort returns the value of ReasoningEffort. +func (s *ModelOverrides) GetReasoningEffort() OptString { + return s.ReasoningEffort +} + +// SetSpeed sets the value of Speed. +func (s *ModelOverrides) SetSpeed(val OptString) { + s.Speed = val +} + +// SetThinking sets the value of Thinking. +func (s *ModelOverrides) SetThinking(val OptString) { + s.Thinking = val +} + +// SetReasoningEffort sets the value of ReasoningEffort. +func (s *ModelOverrides) SetReasoningEffort(val OptString) { + s.ReasoningEffort = val +} + +// NewOptAgentRefMultiagent returns new OptAgentRefMultiagent with value set to v. +func NewOptAgentRefMultiagent(v AgentRefMultiagent) OptAgentRefMultiagent { + return OptAgentRefMultiagent{ + Value: v, + Set: true, + } +} + +// OptAgentRefMultiagent is optional AgentRefMultiagent. +type OptAgentRefMultiagent struct { + Value AgentRefMultiagent + Set bool +} + +// IsSet returns true if OptAgentRefMultiagent was set. +func (o OptAgentRefMultiagent) IsSet() bool { return o.Set } + +// Reset unsets value. +func (o *OptAgentRefMultiagent) Reset() { + var v AgentRefMultiagent + o.Value = v + o.Set = false +} + +// SetTo sets value to v. +func (o *OptAgentRefMultiagent) SetTo(v AgentRefMultiagent) { + o.Set = true + o.Value = v +} + +// Get returns value and boolean that denotes whether value was set. +func (o OptAgentRefMultiagent) Get() (v AgentRefMultiagent, ok bool) { + if !o.Set { + return v, false + } + return o.Value, true +} + +// Or returns value if set, or given parameter if does not. +func (o OptAgentRefMultiagent) Or(d AgentRefMultiagent) AgentRefMultiagent { + if v, ok := o.Get(); ok { + return v + } + return d +} + +// NewOptBool returns new OptBool with value set to v. +func NewOptBool(v bool) OptBool { + return OptBool{ + Value: v, + Set: true, + } +} + +// OptBool is optional bool. +type OptBool struct { + Value bool + Set bool } // IsSet returns true if OptBool was set. @@ -2132,6 +2696,328 @@ func (o OptCacheCreation) Or(d CacheCreation) CacheCreation { return d } +// NewOptEnvironmentConfigOverride returns new OptEnvironmentConfigOverride with value set to v. +func NewOptEnvironmentConfigOverride(v EnvironmentConfigOverride) OptEnvironmentConfigOverride { + return OptEnvironmentConfigOverride{ + Value: v, + Set: true, + } +} + +// OptEnvironmentConfigOverride is optional EnvironmentConfigOverride. +type OptEnvironmentConfigOverride struct { + Value EnvironmentConfigOverride + Set bool +} + +// IsSet returns true if OptEnvironmentConfigOverride was set. +func (o OptEnvironmentConfigOverride) IsSet() bool { return o.Set } + +// Reset unsets value. +func (o *OptEnvironmentConfigOverride) Reset() { + var v EnvironmentConfigOverride + o.Value = v + o.Set = false +} + +// SetTo sets value to v. +func (o *OptEnvironmentConfigOverride) SetTo(v EnvironmentConfigOverride) { + o.Set = true + o.Value = v +} + +// Get returns value and boolean that denotes whether value was set. +func (o OptEnvironmentConfigOverride) Get() (v EnvironmentConfigOverride, ok bool) { + if !o.Set { + return v, false + } + return o.Value, true +} + +// Or returns value if set, or given parameter if does not. +func (o OptEnvironmentConfigOverride) Or(d EnvironmentConfigOverride) EnvironmentConfigOverride { + if v, ok := o.Get(); ok { + return v + } + return d +} + +// NewOptEnvironmentConfigOverrideEnv returns new OptEnvironmentConfigOverrideEnv with value set to v. +func NewOptEnvironmentConfigOverrideEnv(v EnvironmentConfigOverrideEnv) OptEnvironmentConfigOverrideEnv { + return OptEnvironmentConfigOverrideEnv{ + Value: v, + Set: true, + } +} + +// OptEnvironmentConfigOverrideEnv is optional EnvironmentConfigOverrideEnv. +type OptEnvironmentConfigOverrideEnv struct { + Value EnvironmentConfigOverrideEnv + Set bool +} + +// IsSet returns true if OptEnvironmentConfigOverrideEnv was set. +func (o OptEnvironmentConfigOverrideEnv) IsSet() bool { return o.Set } + +// Reset unsets value. +func (o *OptEnvironmentConfigOverrideEnv) Reset() { + var v EnvironmentConfigOverrideEnv + o.Value = v + o.Set = false +} + +// SetTo sets value to v. +func (o *OptEnvironmentConfigOverrideEnv) SetTo(v EnvironmentConfigOverrideEnv) { + o.Set = true + o.Value = v +} + +// Get returns value and boolean that denotes whether value was set. +func (o OptEnvironmentConfigOverrideEnv) Get() (v EnvironmentConfigOverrideEnv, ok bool) { + if !o.Set { + return v, false + } + return o.Value, true +} + +// Or returns value if set, or given parameter if does not. +func (o OptEnvironmentConfigOverrideEnv) Or(d EnvironmentConfigOverrideEnv) EnvironmentConfigOverrideEnv { + if v, ok := o.Get(); ok { + return v + } + return d +} + +// NewOptEnvironmentNetworkingConfig returns new OptEnvironmentNetworkingConfig with value set to v. +func NewOptEnvironmentNetworkingConfig(v EnvironmentNetworkingConfig) OptEnvironmentNetworkingConfig { + return OptEnvironmentNetworkingConfig{ + Value: v, + Set: true, + } +} + +// OptEnvironmentNetworkingConfig is optional EnvironmentNetworkingConfig. +type OptEnvironmentNetworkingConfig struct { + Value EnvironmentNetworkingConfig + Set bool +} + +// IsSet returns true if OptEnvironmentNetworkingConfig was set. +func (o OptEnvironmentNetworkingConfig) IsSet() bool { return o.Set } + +// Reset unsets value. +func (o *OptEnvironmentNetworkingConfig) Reset() { + var v EnvironmentNetworkingConfig + o.Value = v + o.Set = false +} + +// SetTo sets value to v. +func (o *OptEnvironmentNetworkingConfig) SetTo(v EnvironmentNetworkingConfig) { + o.Set = true + o.Value = v +} + +// Get returns value and boolean that denotes whether value was set. +func (o OptEnvironmentNetworkingConfig) Get() (v EnvironmentNetworkingConfig, ok bool) { + if !o.Set { + return v, false + } + return o.Value, true +} + +// Or returns value if set, or given parameter if does not. +func (o OptEnvironmentNetworkingConfig) Or(d EnvironmentNetworkingConfig) EnvironmentNetworkingConfig { + if v, ok := o.Get(); ok { + return v + } + return d +} + +// NewOptEnvironmentPackagesConfig returns new OptEnvironmentPackagesConfig with value set to v. +func NewOptEnvironmentPackagesConfig(v EnvironmentPackagesConfig) OptEnvironmentPackagesConfig { + return OptEnvironmentPackagesConfig{ + Value: v, + Set: true, + } +} + +// OptEnvironmentPackagesConfig is optional EnvironmentPackagesConfig. +type OptEnvironmentPackagesConfig struct { + Value EnvironmentPackagesConfig + Set bool +} + +// IsSet returns true if OptEnvironmentPackagesConfig was set. +func (o OptEnvironmentPackagesConfig) IsSet() bool { return o.Set } + +// Reset unsets value. +func (o *OptEnvironmentPackagesConfig) Reset() { + var v EnvironmentPackagesConfig + o.Value = v + o.Set = false +} + +// SetTo sets value to v. +func (o *OptEnvironmentPackagesConfig) SetTo(v EnvironmentPackagesConfig) { + o.Set = true + o.Value = v +} + +// Get returns value and boolean that denotes whether value was set. +func (o OptEnvironmentPackagesConfig) Get() (v EnvironmentPackagesConfig, ok bool) { + if !o.Set { + return v, false + } + return o.Value, true +} + +// Or returns value if set, or given parameter if does not. +func (o OptEnvironmentPackagesConfig) Or(d EnvironmentPackagesConfig) EnvironmentPackagesConfig { + if v, ok := o.Get(); ok { + return v + } + return d +} + +// NewOptEnvironmentPackagesConfigType returns new OptEnvironmentPackagesConfigType with value set to v. +func NewOptEnvironmentPackagesConfigType(v EnvironmentPackagesConfigType) OptEnvironmentPackagesConfigType { + return OptEnvironmentPackagesConfigType{ + Value: v, + Set: true, + } +} + +// OptEnvironmentPackagesConfigType is optional EnvironmentPackagesConfigType. +type OptEnvironmentPackagesConfigType struct { + Value EnvironmentPackagesConfigType + Set bool +} + +// IsSet returns true if OptEnvironmentPackagesConfigType was set. +func (o OptEnvironmentPackagesConfigType) IsSet() bool { return o.Set } + +// Reset unsets value. +func (o *OptEnvironmentPackagesConfigType) Reset() { + var v EnvironmentPackagesConfigType + o.Value = v + o.Set = false +} + +// SetTo sets value to v. +func (o *OptEnvironmentPackagesConfigType) SetTo(v EnvironmentPackagesConfigType) { + o.Set = true + o.Value = v +} + +// Get returns value and boolean that denotes whether value was set. +func (o OptEnvironmentPackagesConfigType) Get() (v EnvironmentPackagesConfigType, ok bool) { + if !o.Set { + return v, false + } + return o.Value, true +} + +// Or returns value if set, or given parameter if does not. +func (o OptEnvironmentPackagesConfigType) Or(d EnvironmentPackagesConfigType) EnvironmentPackagesConfigType { + if v, ok := o.Get(); ok { + return v + } + return d +} + +// NewOptEnvironmentTosConfig returns new OptEnvironmentTosConfig with value set to v. +func NewOptEnvironmentTosConfig(v EnvironmentTosConfig) OptEnvironmentTosConfig { + return OptEnvironmentTosConfig{ + Value: v, + Set: true, + } +} + +// OptEnvironmentTosConfig is optional EnvironmentTosConfig. +type OptEnvironmentTosConfig struct { + Value EnvironmentTosConfig + Set bool +} + +// IsSet returns true if OptEnvironmentTosConfig was set. +func (o OptEnvironmentTosConfig) IsSet() bool { return o.Set } + +// Reset unsets value. +func (o *OptEnvironmentTosConfig) Reset() { + var v EnvironmentTosConfig + o.Value = v + o.Set = false +} + +// SetTo sets value to v. +func (o *OptEnvironmentTosConfig) SetTo(v EnvironmentTosConfig) { + o.Set = true + o.Value = v +} + +// Get returns value and boolean that denotes whether value was set. +func (o OptEnvironmentTosConfig) Get() (v EnvironmentTosConfig, ok bool) { + if !o.Set { + return v, false + } + return o.Value, true +} + +// Or returns value if set, or given parameter if does not. +func (o OptEnvironmentTosConfig) Or(d EnvironmentTosConfig) EnvironmentTosConfig { + if v, ok := o.Get(); ok { + return v + } + return d +} + +// NewOptEnvironmentWithOverrides returns new OptEnvironmentWithOverrides with value set to v. +func NewOptEnvironmentWithOverrides(v EnvironmentWithOverrides) OptEnvironmentWithOverrides { + return OptEnvironmentWithOverrides{ + Value: v, + Set: true, + } +} + +// OptEnvironmentWithOverrides is optional EnvironmentWithOverrides. +type OptEnvironmentWithOverrides struct { + Value EnvironmentWithOverrides + Set bool +} + +// IsSet returns true if OptEnvironmentWithOverrides was set. +func (o OptEnvironmentWithOverrides) IsSet() bool { return o.Set } + +// Reset unsets value. +func (o *OptEnvironmentWithOverrides) Reset() { + var v EnvironmentWithOverrides + o.Value = v + o.Set = false +} + +// SetTo sets value to v. +func (o *OptEnvironmentWithOverrides) SetTo(v EnvironmentWithOverrides) { + o.Set = true + o.Value = v +} + +// Get returns value and boolean that denotes whether value was set. +func (o OptEnvironmentWithOverrides) Get() (v EnvironmentWithOverrides, ok bool) { + if !o.Set { + return v, false + } + return o.Value, true +} + +// Or returns value if set, or given parameter if does not. +func (o OptEnvironmentWithOverrides) Or(d EnvironmentWithOverrides) EnvironmentWithOverrides { + if v, ok := o.Get(); ok { + return v + } + return d +} + // NewOptInt32 returns new OptInt32 with value set to v. func NewOptInt32(v int32) OptInt32 { return OptInt32{ @@ -2270,6 +3156,115 @@ func (o OptListSessionsOrder) Or(d ListSessionsOrder) ListSessionsOrder { return d } +// NewOptModelOverrides returns new OptModelOverrides with value set to v. +func NewOptModelOverrides(v ModelOverrides) OptModelOverrides { + return OptModelOverrides{ + Value: v, + Set: true, + } +} + +// OptModelOverrides is optional ModelOverrides. +type OptModelOverrides struct { + Value ModelOverrides + Set bool +} + +// IsSet returns true if OptModelOverrides was set. +func (o OptModelOverrides) IsSet() bool { return o.Set } + +// Reset unsets value. +func (o *OptModelOverrides) Reset() { + var v ModelOverrides + o.Value = v + o.Set = false +} + +// SetTo sets value to v. +func (o *OptModelOverrides) SetTo(v ModelOverrides) { + o.Set = true + o.Value = v +} + +// Get returns value and boolean that denotes whether value was set. +func (o OptModelOverrides) Get() (v ModelOverrides, ok bool) { + if !o.Set { + return v, false + } + return o.Value, true +} + +// Or returns value if set, or given parameter if does not. +func (o OptModelOverrides) Or(d ModelOverrides) ModelOverrides { + if v, ok := o.Get(); ok { + return v + } + return d +} + +// NewOptNilSessionEnvironment returns new OptNilSessionEnvironment with value set to v. +func NewOptNilSessionEnvironment(v SessionEnvironment) OptNilSessionEnvironment { + return OptNilSessionEnvironment{ + Value: v, + Set: true, + } +} + +// OptNilSessionEnvironment is optional nullable SessionEnvironment. +type OptNilSessionEnvironment struct { + Value SessionEnvironment + Set bool + Null bool +} + +// IsSet returns true if OptNilSessionEnvironment was set. +func (o OptNilSessionEnvironment) IsSet() bool { return o.Set } + +// Reset unsets value. +func (o *OptNilSessionEnvironment) Reset() { + var v SessionEnvironment + o.Value = v + o.Set = false + o.Null = false +} + +// SetTo sets value to v. +func (o *OptNilSessionEnvironment) SetTo(v SessionEnvironment) { + o.Set = true + o.Null = false + o.Value = v +} + +// IsNull returns true if value is Null. +func (o OptNilSessionEnvironment) IsNull() bool { return o.Null } + +// SetToNull sets value to null. +func (o *OptNilSessionEnvironment) SetToNull() { + o.Set = true + o.Null = true + var v SessionEnvironment + o.Value = v +} + +// Get returns value and boolean that denotes whether value was set. +func (o OptNilSessionEnvironment) Get() (v SessionEnvironment, ok bool) { + if o.Null { + return v, false + } + if !o.Set { + return v, false + } + return o.Value, true +} + +// Or returns value if set, or given parameter if does not. +func (o OptNilSessionEnvironment) Or(d SessionEnvironment) SessionEnvironment { + if v, ok := o.Get(); ok { + return v + } + return d +} + // NewOptNilSessionResourceArray returns new OptNilSessionResourceArray with value set to v. func NewOptNilSessionResourceArray(v []SessionResource) OptNilSessionResourceArray { return OptNilSessionResourceArray{ @@ -2760,18 +3755,29 @@ func (s *SendSessionEventsRequest) SetEvents(val []ManagedAgentsEventParams) { // SendSessionEvents 响应体(回执)。. // Ref: #/components/schemas/SendSessionEventsResponse type SendSessionEventsResponse struct { - // 是否成功接收(server 决定语义)。. - Success OptBool `json:"success"` + // 服务端落库 / 转发完成后的事件回声。. + Data []SendSessionEventsResponseDataItem `json:"data"` } -// GetSuccess returns the value of Success. -func (s *SendSessionEventsResponse) GetSuccess() OptBool { - return s.Success +// GetData returns the value of Data. +func (s *SendSessionEventsResponse) GetData() []SendSessionEventsResponseDataItem { + return s.Data } -// SetSuccess sets the value of Success. -func (s *SendSessionEventsResponse) SetSuccess(val OptBool) { - s.Success = val +// SetData sets the value of Data. +func (s *SendSessionEventsResponse) SetData(val []SendSessionEventsResponseDataItem) { + s.Data = val +} + +type SendSessionEventsResponseDataItem map[string]jx.Raw + +func (s *SendSessionEventsResponseDataItem) init() SendSessionEventsResponseDataItem { + m := *s + if m == nil { + m = map[string]jx.Raw{} + *s = m + } + return m } // 一次 Agent 会话。 @@ -2807,6 +3813,8 @@ type Session struct { Usage OptNilSessionUsage `json:"usage"` // 资源标签。. Tags OptNilTagArray `json:"tags"` + // Session 创建时冻结的 Environment 快照。. + Environment OptNilSessionEnvironment `json:"environment"` } // GetID returns the value of ID. @@ -2874,6 +3882,11 @@ func (s *Session) GetTags() OptNilTagArray { return s.Tags } +// GetEnvironment returns the value of Environment. +func (s *Session) GetEnvironment() OptNilSessionEnvironment { + return s.Environment +} + // SetID sets the value of ID. func (s *Session) SetID(val string) { s.ID = val @@ -2939,6 +3952,11 @@ func (s *Session) SetTags(val OptNilTagArray) { s.Tags = val } +// SetEnvironment sets the value of Environment. +func (s *Session) SetEnvironment(val OptNilSessionEnvironment) { + s.Environment = val +} + // Agent 快照。反映创建瞬间的 Agent wire 形态;shape 与 Ark.Agent 的 // Agent 一致。. type SessionAgent map[string]jx.Raw @@ -2952,6 +3970,17 @@ func (s *SessionAgent) init() SessionAgent { return m } +type SessionEnvironment map[string]jx.Raw + +func (s *SessionEnvironment) init() SessionEnvironment { + m := *s + if m == nil { + m = map[string]jx.Raw{} + *s = m + } + return m +} + // 会话挂载的资源。按 `type` 决定字段: // - `file` → `file_id` // - `memory_store` → `memory_store_id` + `access` @@ -3352,10 +4381,10 @@ func (s *SessionThread) SetUpdatedAt(val string) { type SessionThreadStatus string const ( - SessionThreadStatusIdle SessionThreadStatus = "idle" - SessionThreadStatusRunning SessionThreadStatus = "running" - SessionThreadStatusTerminated SessionThreadStatus = "terminated" - SessionThreadStatusArchived SessionThreadStatus = "archived" + SessionThreadStatusIdle SessionThreadStatus = "idle" + SessionThreadStatusRunning SessionThreadStatus = "running" + SessionThreadStatusTerminated SessionThreadStatus = "terminated" + SessionThreadStatusRescheduling SessionThreadStatus = "rescheduling" ) // AllValues returns all SessionThreadStatus values. @@ -3364,7 +4393,7 @@ func (SessionThreadStatus) AllValues() []SessionThreadStatus { SessionThreadStatusIdle, SessionThreadStatusRunning, SessionThreadStatusTerminated, - SessionThreadStatusArchived, + SessionThreadStatusRescheduling, } } @@ -3377,7 +4406,7 @@ func (s SessionThreadStatus) MarshalText() ([]byte, error) { return []byte(s), nil case SessionThreadStatusTerminated: return []byte(s), nil - case SessionThreadStatusArchived: + case SessionThreadStatusRescheduling: return []byte(s), nil default: return nil, errors.Errorf("invalid value: %q", s) @@ -3396,8 +4425,8 @@ func (s *SessionThreadStatus) UnmarshalText(data []byte) error { case SessionThreadStatusTerminated: *s = SessionThreadStatusTerminated return nil - case SessionThreadStatusArchived: - *s = SessionThreadStatusArchived + case SessionThreadStatusRescheduling: + *s = SessionThreadStatusRescheduling return nil default: return errors.Errorf("invalid value: %q", data) diff --git a/arkruntime/model/session/oas_validators_gen.go b/arkruntime/model/session/oas_validators_gen.go index 7b3daa7..6cbbccd 100644 --- a/arkruntime/model/session/oas_validators_gen.go +++ b/arkruntime/model/session/oas_validators_gen.go @@ -12,21 +12,62 @@ import ( "github.com/volcengine/ark-runtime-go/arkruntime/internal/validate" ) -func (s AgentIdentifier) Validate() error { - switch s.Type { - case StringAgentIdentifier: - return nil // no validation needed - case AgentRefAgentIdentifier: - if err := s.AgentRef.Validate(); err != nil { - return err +func (s *CreateSessionRequest) Validate() error { + if s == nil { + return validate.ErrNilPointer + } + + var failures []validate.FieldError + if err := func() error { + if value, ok := s.Environment.Get(); ok { + if err := func() error { + if err := value.Validate(); err != nil { + return err + } + return nil + }(); err != nil { + return err + } } return nil - default: - return errors.Errorf("invalid type %q", s.Type) + }(); err != nil { + failures = append(failures, validate.FieldError{ + Name: "environment", + Error: err, + }) + } + if err := func() error { + var failures []validate.FieldError + for i, elem := range s.Resources { + if err := func() error { + if err := elem.Validate(); err != nil { + return err + } + return nil + }(); err != nil { + failures = append(failures, validate.FieldError{ + Name: fmt.Sprintf("[%d]", i), + Error: err, + }) + } + } + if len(failures) > 0 { + return &validate.Error{Fields: failures} + } + return nil + }(); err != nil { + failures = append(failures, validate.FieldError{ + Name: "resources", + Error: err, + }) } + if len(failures) > 0 { + return &validate.Error{Fields: failures} + } + return nil } -func (s *AgentRef) Validate() error { +func (s *CreateSessionResourceRequest) Validate() error { if s == nil { return validate.ErrNilPointer } @@ -49,54 +90,66 @@ func (s *AgentRef) Validate() error { return nil } -func (s AgentRefType) Validate() error { +func (s CreateSessionResourceRequestType) Validate() error { switch s { - case "agent": + case "file": return nil default: return errors.Errorf("invalid value: %v", s) } } -func (s *CreateSessionRequest) Validate() error { +func (s *EnvironmentConfigOverride) Validate() error { if s == nil { return validate.ErrNilPointer } var failures []validate.FieldError if err := func() error { - if err := s.Agent.Validate(); err != nil { - return err + if value, ok := s.Packages.Get(); ok { + if err := func() error { + if err := value.Validate(); err != nil { + return err + } + return nil + }(); err != nil { + return err + } } return nil }(); err != nil { failures = append(failures, validate.FieldError{ - Name: "agent", + Name: "packages", Error: err, }) } + if len(failures) > 0 { + return &validate.Error{Fields: failures} + } + return nil +} + +func (s *EnvironmentPackagesConfig) Validate() error { + if s == nil { + return validate.ErrNilPointer + } + + var failures []validate.FieldError if err := func() error { - var failures []validate.FieldError - for i, elem := range s.Resources { + if value, ok := s.Type.Get(); ok { if err := func() error { - if err := elem.Validate(); err != nil { + if err := value.Validate(); err != nil { return err } return nil }(); err != nil { - failures = append(failures, validate.FieldError{ - Name: fmt.Sprintf("[%d]", i), - Error: err, - }) + return err } } - if len(failures) > 0 { - return &validate.Error{Fields: failures} - } return nil }(); err != nil { failures = append(failures, validate.FieldError{ - Name: "resources", + Name: "type", Error: err, }) } @@ -106,7 +159,16 @@ func (s *CreateSessionRequest) Validate() error { return nil } -func (s *CreateSessionResourceRequest) Validate() error { +func (s EnvironmentPackagesConfigType) Validate() error { + switch s { + case "packages": + return nil + default: + return errors.Errorf("invalid value: %v", s) + } +} + +func (s *EnvironmentWithOverrides) Validate() error { if s == nil { return validate.ErrNilPointer } @@ -123,15 +185,33 @@ func (s *CreateSessionResourceRequest) Validate() error { Error: err, }) } + if err := func() error { + if value, ok := s.Config.Get(); ok { + if err := func() error { + if err := value.Validate(); err != nil { + return err + } + return nil + }(); err != nil { + return err + } + } + return nil + }(); err != nil { + failures = append(failures, validate.FieldError{ + Name: "config", + Error: err, + }) + } if len(failures) > 0 { return &validate.Error{Fields: failures} } return nil } -func (s CreateSessionResourceRequestType) Validate() error { +func (s EnvironmentWithOverridesType) Validate() error { switch s { - case "file": + case "environment_with_overrides": return nil default: return errors.Errorf("invalid value: %v", s) @@ -546,6 +626,29 @@ func (s *SendSessionEventsRequest) Validate() error { return nil } +func (s *SendSessionEventsResponse) Validate() error { + if s == nil { + return validate.ErrNilPointer + } + + var failures []validate.FieldError + if err := func() error { + if s.Data == nil { + return errors.New("nil is invalid value") + } + return nil + }(); err != nil { + failures = append(failures, validate.FieldError{ + Name: "data", + Error: err, + }) + } + if len(failures) > 0 { + return &validate.Error{Fields: failures} + } + return nil +} + func (s *Session) Validate() error { if s == nil { return validate.ErrNilPointer @@ -773,7 +876,7 @@ func (s SessionThreadStatus) Validate() error { return nil case "terminated": return nil - case "archived": + case "rescheduling": return nil default: return errors.Errorf("invalid value: %v", s) diff --git a/arkruntime/model/session/session_stream_shim.go b/arkruntime/model/session/session_stream_shim.go index c5db6ff..9c8dba3 100644 --- a/arkruntime/model/session/session_stream_shim.go +++ b/arkruntime/model/session/session_stream_shim.go @@ -254,10 +254,11 @@ type ManagedAgentsAgentThinkingEvent struct { // `Input` is the tool's argument object (schema per tool). type ManagedAgentsAgentToolUseEvent struct { sessionEventBase - SessionThreadID string `json:"session_thread_id,omitempty"` - ToolUseID string `json:"tool_use_id"` - Name string `json:"name"` - Input json.RawMessage `json:"input,omitempty"` + SessionThreadID string `json:"session_thread_id,omitempty"` + ToolUseID string `json:"tool_use_id"` + Name string `json:"name"` + Input json.RawMessage `json:"input,omitempty"` + EvaluatedPermission string `json:"evaluated_permission,omitempty"` } // ManagedAgentsAgentToolResultEvent — the tool's result (post @@ -271,6 +272,16 @@ type ManagedAgentsAgentToolResultEvent struct { IsError bool `json:"is_error,omitempty"` } +// ManagedAgentsAgentCustomToolUseEvent — the agent invoked a user-defined +// custom tool handled by the self-hosted worker. +type ManagedAgentsAgentCustomToolUseEvent struct { + sessionEventBase + SessionThreadID string `json:"session_thread_id,omitempty"` + CustomToolUseID string `json:"custom_tool_use_id,omitempty"` + Name string `json:"name"` + Input json.RawMessage `json:"input,omitempty"` +} + // ManagedAgentsAgentMCPToolUseEvent — the agent invoked an MCP tool // (registered via a Vault credential). Distinct wire event from // agent.tool_use to signal the extra `mcp_server_name` routing metadata. @@ -430,6 +441,8 @@ func decodeSessionEvent(raw []byte) ManagedAgentsSessionEvent { out = &ManagedAgentsAgentToolUseEvent{} case "agent.tool_result": out = &ManagedAgentsAgentToolResultEvent{} + case "agent.custom_tool_use": + out = &ManagedAgentsAgentCustomToolUseEvent{} case "agent.mcp_tool_use": out = &ManagedAgentsAgentMCPToolUseEvent{} case "agent.mcp_tool_result": diff --git a/arkruntime/model/session/session_thread_shim.go b/arkruntime/model/session/session_thread_shim.go index 61919e5..3780b01 100644 --- a/arkruntime/model/session/session_thread_shim.go +++ b/arkruntime/model/session/session_thread_shim.go @@ -108,6 +108,14 @@ func URLQuerySessionEventsList(req *SessionEventsListEventsParams) (url.Values, return apiquery.Marshal(req) } +// URLQuerySessionEventsStream encodes SessionEventsStreamEventsParams. +func URLQuerySessionEventsStream(req *SessionEventsStreamEventsParams) (url.Values, error) { + if req == nil { + return url.Values{}, nil + } + return apiquery.Marshal(req) +} + // URLQuerySessionThreadsList encodes SessionThreadsListParams. func URLQuerySessionThreadsList(req *SessionThreadsListParams) (url.Values, error) { if req == nil { @@ -123,3 +131,11 @@ func URLQuerySessionThreadEventsList(req *SessionThreadsListEventsParams) (url.V } return apiquery.Marshal(req) } + +// URLQuerySessionThreadEventsStream encodes SessionThreadsStreamEventsParams. +func URLQuerySessionThreadEventsStream(req *SessionThreadsStreamEventsParams) (url.Values, error) { + if req == nil { + return url.Values{}, nil + } + return apiquery.Marshal(req) +} diff --git a/arkruntime/model/skill/oas_json_gen.go b/arkruntime/model/skill/oas_json_gen.go index ca19f0f..3422a94 100644 --- a/arkruntime/model/skill/oas_json_gen.go +++ b/arkruntime/model/skill/oas_json_gen.go @@ -29,10 +29,17 @@ func (s *CreateSkillRequest) encodeFields(e *jx.Encoder) { s.DisplayTitle.Encode(e) } } + { + if s.ProtectionEnabled.Set { + e.FieldStart("protection_enabled") + s.ProtectionEnabled.Encode(e) + } + } } -var jsonFieldsNameOfCreateSkillRequest = [1]string{ +var jsonFieldsNameOfCreateSkillRequest = [2]string{ 0: "display_title", + 1: "protection_enabled", } // Decode decodes CreateSkillRequest from json. @@ -53,6 +60,16 @@ func (s *CreateSkillRequest) Decode(d *jx.Decoder) error { }(); err != nil { return errors.Wrap(err, "decode field \"display_title\"") } + case "protection_enabled": + if err := func() error { + s.ProtectionEnabled.Reset() + if err := s.ProtectionEnabled.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"protection_enabled\"") + } default: return d.Skip() } @@ -77,6 +94,41 @@ func (s *CreateSkillRequest) UnmarshalJSON(data []byte) error { return s.Decode(d) } +// Encode encodes bool as json. +func (o OptBool) Encode(e *jx.Encoder) { + if !o.Set { + return + } + e.Bool(bool(o.Value)) +} + +// Decode decodes bool from json. +func (o *OptBool) Decode(d *jx.Decoder) error { + if o == nil { + return errors.New("invalid: unable to decode OptBool to nil") + } + o.Set = true + v, err := d.Bool() + if err != nil { + return err + } + o.Value = bool(v) + return nil +} + +// MarshalJSON implements stdjson.Marshaler. +func (s OptBool) MarshalJSON() ([]byte, error) { + e := jx.Encoder{} + s.Encode(&e) + return e.Bytes(), nil +} + +// UnmarshalJSON implements stdjson.Unmarshaler. +func (s *OptBool) UnmarshalJSON(data []byte) error { + d := jx.DecodeBytes(data) + return s.Decode(d) +} + // Encode encodes string as json. func (o OptString) Encode(e *jx.Encoder) { if !o.Set { @@ -144,18 +196,42 @@ func (s *Skill) encodeFields(e *jx.Encoder) { e.Str(s.LatestVersion) } { - e.FieldStart("name") - e.Str(s.Name) + e.FieldStart("display_title") + e.Str(s.DisplayTitle) + } + { + e.FieldStart("source") + e.Str(s.Source) + } + { + e.FieldStart("updated_at") + e.Int64(s.UpdatedAt) + } + { + if s.Name.Set { + e.FieldStart("name") + s.Name.Encode(e) + } + } + { + if s.ProtectionEnabled.Set { + e.FieldStart("protection_enabled") + s.ProtectionEnabled.Encode(e) + } } } -var jsonFieldsNameOfSkill = [6]string{ +var jsonFieldsNameOfSkill = [10]string{ 0: "id", 1: "object", 2: "created_at", 3: "description", 4: "latest_version", - 5: "name", + 5: "display_title", + 6: "source", + 7: "updated_at", + 8: "name", + 9: "protection_enabled", } // Decode decodes Skill from json. @@ -163,7 +239,7 @@ func (s *Skill) Decode(d *jx.Decoder) error { if s == nil { return errors.New("invalid: unable to decode Skill to nil") } - var requiredBitSet [1]uint8 + var requiredBitSet [2]uint8 if err := d.ObjBytes(func(d *jx.Decoder, k []byte) error { switch string(k) { @@ -223,18 +299,62 @@ func (s *Skill) Decode(d *jx.Decoder) error { }(); err != nil { return errors.Wrap(err, "decode field \"latest_version\"") } - case "name": + case "display_title": requiredBitSet[0] |= 1 << 5 if err := func() error { v, err := d.Str() - s.Name = string(v) + s.DisplayTitle = string(v) + if err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"display_title\"") + } + case "source": + requiredBitSet[0] |= 1 << 6 + if err := func() error { + v, err := d.Str() + s.Source = string(v) + if err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"source\"") + } + case "updated_at": + requiredBitSet[0] |= 1 << 7 + if err := func() error { + v, err := d.Int64() + s.UpdatedAt = int64(v) if err != nil { return err } return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"updated_at\"") + } + case "name": + if err := func() error { + s.Name.Reset() + if err := s.Name.Decode(d); err != nil { + return err + } + return nil }(); err != nil { return errors.Wrap(err, "decode field \"name\"") } + case "protection_enabled": + if err := func() error { + s.ProtectionEnabled.Reset() + if err := s.ProtectionEnabled.Decode(d); err != nil { + return err + } + return nil + }(); err != nil { + return errors.Wrap(err, "decode field \"protection_enabled\"") + } default: return d.Skip() } @@ -244,8 +364,9 @@ func (s *Skill) Decode(d *jx.Decoder) error { } // Validate required fields. var failures []validate.FieldError - for i, mask := range [1]uint8{ - 0b00110111, + for i, mask := range [2]uint8{ + 0b11110111, + 0b00000000, } { if result := (requiredBitSet[i] & mask) ^ mask; result != 0 { // Mask only required fields and check equality to mask using XOR. diff --git a/arkruntime/model/skill/oas_parameters_gen.go b/arkruntime/model/skill/oas_parameters_gen.go index 2591d96..ca7e792 100644 --- a/arkruntime/model/skill/oas_parameters_gen.go +++ b/arkruntime/model/skill/oas_parameters_gen.go @@ -5,6 +5,12 @@ package skill +// SkillVersionContentDownloadParams is parameters of SkillVersionContent_download operation. +type SkillVersionContentDownloadParams struct { + SkillId string + Version string +} + // SkillsRetrieveParams is parameters of Skills_retrieve operation. type SkillsRetrieveParams struct { SkillId string diff --git a/arkruntime/model/skill/oas_schemas_gen.go b/arkruntime/model/skill/oas_schemas_gen.go index 3351213..391327b 100644 --- a/arkruntime/model/skill/oas_schemas_gen.go +++ b/arkruntime/model/skill/oas_schemas_gen.go @@ -13,6 +13,8 @@ import ( type CreateSkillRequest struct { // Skill 展示名。. DisplayTitle OptString `json:"display_title" form:"display_title"` + // 是否启用 Skill 内容保护。. + ProtectionEnabled OptBool `json:"protection_enabled" form:"protection_enabled"` } // GetDisplayTitle returns the value of DisplayTitle. @@ -20,11 +22,67 @@ func (s *CreateSkillRequest) GetDisplayTitle() OptString { return s.DisplayTitle } +// GetProtectionEnabled returns the value of ProtectionEnabled. +func (s *CreateSkillRequest) GetProtectionEnabled() OptBool { + return s.ProtectionEnabled +} + // SetDisplayTitle sets the value of DisplayTitle. func (s *CreateSkillRequest) SetDisplayTitle(val OptString) { s.DisplayTitle = val } +// SetProtectionEnabled sets the value of ProtectionEnabled. +func (s *CreateSkillRequest) SetProtectionEnabled(val OptBool) { + s.ProtectionEnabled = val +} + +// NewOptBool returns new OptBool with value set to v. +func NewOptBool(v bool) OptBool { + return OptBool{ + Value: v, + Set: true, + } +} + +// OptBool is optional bool. +type OptBool struct { + Value bool + Set bool +} + +// IsSet returns true if OptBool was set. +func (o OptBool) IsSet() bool { return o.Set } + +// Reset unsets value. +func (o *OptBool) Reset() { + var v bool + o.Value = v + o.Set = false +} + +// SetTo sets value to v. +func (o *OptBool) SetTo(v bool) { + o.Set = true + o.Value = v +} + +// Get returns value and boolean that denotes whether value was set. +func (o OptBool) Get() (v bool, ok bool) { + if !o.Set { + return v, false + } + return o.Value, true +} + +// Or returns value if set, or given parameter if does not. +func (o OptBool) Or(d bool) bool { + if v, ok := o.Get(); ok { + return v + } + return d +} + // NewOptString returns new OptString with value set to v. func NewOptString(v string) OptString { return OptString{ @@ -87,8 +145,16 @@ type Skill struct { Description OptString `json:"description"` // 最新版本号。. LatestVersion string `json:"latest_version"` - // 人类可读名称。. - Name string `json:"name"` + // Skill 展示名。. + DisplayTitle string `json:"display_title"` + // Skill 来源,例如 `custom` / `skill_hub` / `ark`。. + Source string `json:"source"` + // 更新时间(Unix 秒)。. + UpdatedAt int64 `json:"updated_at"` + // SKILL.md 中解析出的 name。. + Name OptString `json:"name"` + // 是否启用内容保护。. + ProtectionEnabled OptBool `json:"protection_enabled"` } // GetID returns the value of ID. @@ -116,11 +182,31 @@ func (s *Skill) GetLatestVersion() string { return s.LatestVersion } +// GetDisplayTitle returns the value of DisplayTitle. +func (s *Skill) GetDisplayTitle() string { + return s.DisplayTitle +} + +// GetSource returns the value of Source. +func (s *Skill) GetSource() string { + return s.Source +} + +// GetUpdatedAt returns the value of UpdatedAt. +func (s *Skill) GetUpdatedAt() int64 { + return s.UpdatedAt +} + // GetName returns the value of Name. -func (s *Skill) GetName() string { +func (s *Skill) GetName() OptString { return s.Name } +// GetProtectionEnabled returns the value of ProtectionEnabled. +func (s *Skill) GetProtectionEnabled() OptBool { + return s.ProtectionEnabled +} + // SetID sets the value of ID. func (s *Skill) SetID(val string) { s.ID = val @@ -146,11 +232,31 @@ func (s *Skill) SetLatestVersion(val string) { s.LatestVersion = val } +// SetDisplayTitle sets the value of DisplayTitle. +func (s *Skill) SetDisplayTitle(val string) { + s.DisplayTitle = val +} + +// SetSource sets the value of Source. +func (s *Skill) SetSource(val string) { + s.Source = val +} + +// SetUpdatedAt sets the value of UpdatedAt. +func (s *Skill) SetUpdatedAt(val int64) { + s.UpdatedAt = val +} + // SetName sets the value of Name. -func (s *Skill) SetName(val string) { +func (s *Skill) SetName(val OptString) { s.Name = val } +// SetProtectionEnabled sets the value of ProtectionEnabled. +func (s *Skill) SetProtectionEnabled(val OptBool) { + s.ProtectionEnabled = val +} + // 固定 `"skill"`。. type SkillObject string @@ -185,3 +291,18 @@ func (s *SkillObject) UnmarshalText(data []byte) error { return errors.Errorf("invalid value: %q", data) } } + +// SkillVersionContentDownloadFound is response for SkillVersionContentDownload operation. +type SkillVersionContentDownloadFound struct { + Location string +} + +// GetLocation returns the value of Location. +func (s *SkillVersionContentDownloadFound) GetLocation() string { + return s.Location +} + +// SetLocation sets the value of Location. +func (s *SkillVersionContentDownloadFound) SetLocation(val string) { + s.Location = val +} diff --git a/arkruntime/model/skill/skill_shim.go b/arkruntime/model/skill/skill_shim.go index 52251c1..5c063ab 100644 --- a/arkruntime/model/skill/skill_shim.go +++ b/arkruntime/model/skill/skill_shim.go @@ -32,10 +32,13 @@ type UploadForm struct { // DisplayTitle is the optional user-visible title for this skill. DisplayTitle string + + // ProtectionEnabled sets the optional skill protection flag. + ProtectionEnabled *bool } -// MarshalMultipart writes DisplayTitle as a form part and appends the -// zip File as the `files` binary part. +// MarshalMultipart writes metadata form parts and appends the zip File as the +// `files` binary part. func (u *UploadForm) MarshalMultipart() (data []byte, contentType string, err error) { buf := bytes.NewBuffer(nil) writer := multipart.NewWriter(buf) @@ -46,6 +49,12 @@ func (u *UploadForm) MarshalMultipart() (data []byte, contentType string, err er return nil, "", err } } + if u.ProtectionEnabled != nil { + if err = writer.WriteField("protection_enabled", boolString(*u.ProtectionEnabled)); err != nil { + _ = writer.Close() + return nil, "", err + } + } if u.File == nil { _ = writer.Close() @@ -81,3 +90,10 @@ var errNilFile = errNilFileValue{} type errNilFileValue struct{} func (errNilFileValue) Error() string { return "missing required file part" } + +func boolString(v bool) string { + if v { + return "true" + } + return "false" +} diff --git a/arkruntime/self_hosted_client_test.go b/arkruntime/self_hosted_client_test.go new file mode 100644 index 0000000..a9523b3 --- /dev/null +++ b/arkruntime/self_hosted_client_test.go @@ -0,0 +1,412 @@ +// Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +// SPDX-License-Identifier: Apache-2.0 + +package arkruntime + +import ( + "context" + "encoding/json" + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/volcengine/ark-runtime-go/arkruntime/model/environment" + "github.com/volcengine/ark-runtime-go/arkruntime/model/session" +) + +const testBearerToken = "Bearer test-api-key" + +func TestEnvironmentWorkRequests(t *testing.T) { + var seen []string + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + seen = append(seen, r.Method+" "+r.URL.Path) + if got := r.Header.Get("Authorization"); got != testBearerToken { + t.Fatalf("Authorization = %q", got) + } + if got := r.Header.Get("X-Ark-Environment-Key"); got != "" { + t.Fatalf("unexpected environment key header = %q", got) + } + w.Header().Set("Content-Type", "application/json") + + switch r.Method + " " + r.URL.Path { + case "GET /environments/env-1/work/poll": + if got := r.Header.Get(environmentWorkWorkerIDHeader); got != "worker-1" { + t.Fatalf("Ark-Worker-ID = %q", got) + } + if got := r.URL.Query().Get("block_ms"); got != "999" { + t.Fatalf("block_ms = %q", got) + } + if got := r.URL.Query().Get("worker_id"); got != "" { + t.Fatalf("unexpected worker_id query = %q", got) + } + if got := r.URL.Query().Get("max_items"); got != "" { + t.Fatalf("unexpected max_items query = %q", got) + } + _, _ = w.Write([]byte(`{"id":"work-1","created_at":"2026-08-10T00:00:00Z","data":{"id":"sess-1","type":"session"},"environment_id":"env-1","latest_heartbeat_at":"2026-08-10T00:00:00Z","state":"queued","type":"work"}`)) + case "POST /environments/env-1/work/work-1/ack": + assertNoBody(t, r) + if got := r.Header.Get(environmentWorkWorkerIDHeader); got != "worker-1" { + t.Fatalf("Ark-Worker-ID = %q", got) + } + _, _ = w.Write([]byte(`{"id":"work-1","created_at":"2026-08-10T00:00:00Z","data":{"id":"sess-1","type":"session"},"environment_id":"env-1","state":"starting","type":"work"}`)) + case "POST /environments/env-1/work/work-1/heartbeat": + assertNoBody(t, r) + if got := r.URL.Query().Get("desired_ttl_seconds"); got != "30" { + t.Fatalf("desired_ttl_seconds = %q", got) + } + if got := r.URL.Query().Get("expected_last_heartbeat"); got != "2026-08-10T00:00:00Z" { + t.Fatalf("expected_last_heartbeat = %q", got) + } + _, _ = w.Write([]byte(`{"last_heartbeat":"2026-08-10T00:00:00Z","lease_extended":true,"state":"active","ttl_seconds":30,"type":"work_heartbeat"}`)) + case "POST /environments/env-1/work/work-1/stop": + var body map[string]bool + decodeJSONBody(t, r, &body) + if len(body) != 1 || !body["force"] { + t.Fatalf("stop body = %+v", body) + } + _, _ = w.Write([]byte(`{"id":"work-1","created_at":"2026-08-10T00:00:00Z","data":{"id":"sess-1","type":"session"},"environment_id":"env-1","state":"stopping","type":"work"}`)) + case "POST /environments/env-1/work/work-2/stop": + assertNoBody(t, r) + _, _ = w.Write([]byte(`{"id":"work-2","created_at":"2026-08-10T00:00:00Z","data":{"id":"sess-2","type":"session"},"environment_id":"env-1","state":"stopping","type":"work"}`)) + default: + http.NotFound(w, r) + } + })) + defer server.Close() + + client := NewClientWithApiKey("test-api-key", WithBaseUrl(server.URL)) + ctx := context.Background() + + work, err := client.PollWork(ctx, &environment.PollWorkRequest{ + EnvironmentID: "env-1", + WorkerID: "worker-1", + BlockMS: 999, + }) + if err != nil { + t.Fatalf("PollWork() error = %v", err) + } + if work == nil || work.SessionIDValue() != "sess-1" { + t.Fatalf("PollWork() = %+v", work) + } + if got := work.LatestHeartbeatValue(); got != "2026-08-10T00:00:00Z" { + t.Fatalf("PollWork().LatestHeartbeatValue() = %q", got) + } + if err := client.AckWork(ctx, &environment.AckWorkRequest{ + EnvironmentID: "env-1", + WorkID: "work-1", + WorkerID: environment.NewOptString("worker-1"), + }); err != nil { + t.Fatalf("AckWork() error = %v", err) + } + heartbeat, err := client.HeartbeatWork(ctx, &environment.HeartbeatWorkRequest{ + EnvironmentID: "env-1", + WorkID: "work-1", + DesiredTTLSeconds: environment.NewOptInt64(30), + ExpectedLastHeartbeat: environment.NewOptString("2026-08-10T00:00:00Z"), + }) + if err != nil { + t.Fatalf("HeartbeatWork() error = %v", err) + } + if heartbeat == nil { + t.Fatalf("HeartbeatWork() = %+v", heartbeat) + } + if heartbeat.TTLSeconds != 30 { + t.Fatalf("HeartbeatWork() = %+v", heartbeat) + } + if err := client.StopWork(ctx, &environment.StopWorkRequest{ + EnvironmentID: "env-1", + WorkID: "work-1", + Force: environment.NewOptBool(true), + }); err != nil { + t.Fatalf("StopWork() error = %v", err) + } + if err := client.StopWork(ctx, &environment.StopWorkRequest{ + EnvironmentID: "env-1", + WorkID: "work-2", + }); err != nil { + t.Fatalf("StopWork() error = %v", err) + } + + want := []string{ + "GET /environments/env-1/work/poll", + "POST /environments/env-1/work/work-1/ack", + "POST /environments/env-1/work/work-1/heartbeat", + "POST /environments/env-1/work/work-1/stop", + "POST /environments/env-1/work/work-2/stop", + } + if strings.Join(seen, "\n") != strings.Join(want, "\n") { + t.Fatalf("requests = %v, want %v", seen, want) + } +} + +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" { + http.NotFound(w, r) + return + } + assertNoBody(t, r) + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{}`)) + })) + defer server.Close() + + client := NewClientWithApiKey("test-api-key", WithBaseUrl(server.URL)) + work, err := client.PollWork(context.Background(), &environment.PollWorkRequest{ + EnvironmentID: "env-empty", + }) + if err != nil { + t.Fatalf("PollWork() error = %v", err) + } + if work != nil { + t.Fatalf("PollWork() = %+v, want nil", work) + } +} + +func TestSkillContentDownload(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodGet && r.URL.Path == "/zip/skill-1-v1.zip" { + w.Header().Set("Content-Type", "application/zip") + _, _ = w.Write([]byte("zip-bytes")) + return + } + if r.Method != http.MethodGet || r.URL.Path != "/skills/skill-1/versions/v1/content" { + http.NotFound(w, r) + return + } + if got := r.Header.Get("Authorization"); got != testBearerToken { + t.Fatalf("Authorization = %q", got) + } + http.Redirect(w, r, "/zip/skill-1-v1.zip", http.StatusFound) + })) + defer server.Close() + + client := NewClientWithApiKey("test-api-key", WithBaseUrl(server.URL)) + content, err := client.OpenSkillContent(context.Background(), "skill-1", "v1") + if err != nil { + t.Fatalf("OpenSkillContent() error = %v", err) + } + defer content.Body.Close() + data, err := io.ReadAll(content.Body) + if err != nil { + t.Fatalf("ReadAll() error = %v", err) + } + if string(data) != "zip-bytes" || content.ContentType != "application/zip" { + t.Fatalf("content = %q %q", data, content.ContentType) + } +} + +func TestCreateSkillWithOptionsMultipartContract(t *testing.T) { + protectionEnabled := true + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost || r.URL.Path != "/skills" { + http.NotFound(w, r) + return + } + if err := r.ParseMultipartForm(1024); err != nil { + t.Fatalf("ParseMultipartForm() error = %v", err) + } + if got := r.FormValue("display_title"); got != "Readiness Skill" { + t.Fatalf("display_title = %q", got) + } + if got := r.FormValue("protection_enabled"); got != "true" { + t.Fatalf("protection_enabled = %q", got) + } + if files := r.MultipartForm.File["files"]; len(files) != 1 || files[0].Filename != "skill.zip" { + t.Fatalf("files = %+v", files) + } + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"id":"skill-1","object":"skill","created_at":1786430000,"description":"ok","latest_version":"1","display_title":"Readiness Skill","source":"custom","updated_at":1786430001,"name":"readiness","protection_enabled":true}`)) + })) + defer server.Close() + + client := NewClientWithApiKey("test-api-key", WithBaseUrl(server.URL)) + out, err := client.CreateSkillWithOptions( + context.Background(), + strings.NewReader("zip-bytes"), + "skill.zip", + CreateSkillOptions{ + DisplayTitle: "Readiness Skill", + ProtectionEnabled: &protectionEnabled, + }, + ) + if err != nil { + t.Fatalf("CreateSkillWithOptions() error = %v", err) + } + if out == nil || out.DisplayTitle != "Readiness Skill" || out.Source != "custom" { + t.Fatalf("CreateSkillWithOptions() = %+v", out) + } +} + +func TestSendSessionEventRaw(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost || r.URL.Path != "/sessions/sess-1/events" { + http.NotFound(w, r) + return + } + if got := r.Header.Get("Idempotency-Key"); got != "" { + t.Fatalf("Idempotency-Key = %q", got) + } + var body struct { + Events []map[string]any `json:"events"` + } + decodeJSONBody(t, r, &body) + if len(body.Events) != 1 || body.Events[0]["type"] != "user.tool_result" { + t.Fatalf("body = %+v", body) + } + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"data":[{"type":"user.tool_result","tool_use_id":"toolu-1"}]}`)) + })) + defer server.Close() + + client := NewClientWithApiKey("test-api-key", WithBaseUrl(server.URL)) + err := client.SendSessionEventRaw( + context.Background(), + "sess-1", + map[string]any{"type": "user.tool_result", "tool_use_id": "toolu-1"}, + ) + if err != nil { + t.Fatalf("SendSessionEventRaw() error = %v", err) + } +} + +func TestSessionMAContractRequests(t *testing.T) { + var seen []string + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + seen = append(seen, r.Method+" "+r.URL.Path) + if got := r.Header.Get("Authorization"); got != testBearerToken { + t.Fatalf("Authorization = %q", got) + } + w.Header().Set("Content-Type", "application/json") + + switch r.Method + " " + r.URL.Path { + case "GET /sessions/sess-1": + assertNoBody(t, r) + assertQueryValue(t, r, "work_id", "") + _, _ = w.Write([]byte(`{"id":"sess-1","type":"session","status":"idle","environment_id":"env-1","agent":{},"created_at":"2026-08-10T00:00:00Z","updated_at":"2026-08-10T00:00:00Z"}`)) + case "GET /sessions/sess-1/events": + assertNoBody(t, r) + assertQueryValue(t, r, "created_at[gt]", "2026-08-10T00:00:00Z") + assertQueryValue(t, r, "limit", "2") + assertQueryValue(t, r, "order", "asc") + assertQueryValue(t, r, "page", "page-1") + assertQueryValues(t, r, "types", []string{"agent.message", "agent.tool_use"}) + assertQueryValue(t, r, "work_id", "") + assertQueryValue(t, r, "lease_id", "") + _, _ = w.Write([]byte(`{"data":[{"id":"sevt-1","type":"agent.message","processed_at":"2026-08-10T00:00:01Z","content":[]}],"next_page":"page-2"}`)) + case "GET /sessions/sess-1/events/stream": + assertNoBody(t, r) + if got := r.Header.Get("Accept"); got != "text/event-stream" { + t.Fatalf("Accept = %q", got) + } + assertQueryValues(t, r, "event_deltas", []string{"agent.message", "agent.thinking"}) + w.Header().Set("Content-Type", "text/event-stream") + _, _ = w.Write([]byte("data: [DONE]\n\n")) + case "POST /sessions/sess-1/events": + if got := r.Header.Get("Idempotency-Key"); got != "" { + t.Fatalf("Idempotency-Key = %q", got) + } + var body struct { + Events []map[string]any `json:"events"` + } + decodeJSONBody(t, r, &body) + if len(body.Events) != 1 || body.Events[0]["type"] != "user.tool_result" { + t.Fatalf("body = %+v", body) + } + if _, ok := body.Events[0]["work_id"]; ok { + t.Fatalf("unexpected work_id in body = %+v", body) + } + if _, ok := body.Events[0]["lease_id"]; ok { + t.Fatalf("unexpected lease_id in body = %+v", body) + } + _, _ = w.Write([]byte(`{"data":[{"type":"user.tool_result","tool_use_id":"toolu-1"}]}`)) + default: + http.NotFound(w, r) + } + })) + defer server.Close() + + client := NewClientWithApiKey("test-api-key", WithBaseUrl(server.URL)) + ctx := context.Background() + if got, err := client.GetSession(ctx, "sess-1"); err != nil || got == nil || got.ID != "sess-1" { + t.Fatalf("GetSession() = %+v, %v", got, err) + } + resp, err := client.ListSessionEvents(ctx, "sess-1", &session.SessionEventsListEventsParams{ + CreatedAtGt: session.NewOptString("2026-08-10T00:00:00Z"), + Limit: session.NewOptInt32(2), + Order: session.NewOptListSessionsOrder(session.ListSessionsOrderAsc), + Page: session.NewOptString("page-1"), + Types: []string{"agent.message", "agent.tool_use"}, + }) + if err != nil { + t.Fatalf("ListSessionEvents() error = %v", err) + } + if resp == nil || len(resp.Events) != 1 { + t.Fatalf("ListSessionEvents() = %+v", resp) + } + decoder, err := client.StreamSessionEventsWithParams(ctx, "sess-1", &session.SessionEventsStreamEventsParams{ + EventDeltas: []string{"agent.message", "agent.thinking"}, + }) + if err != nil { + t.Fatalf("StreamSessionEventsWithParams() error = %v", err) + } + _ = decoder.Close() + err = client.SendSessionEventRaw( + ctx, + "sess-1", + map[string]any{"type": "user.tool_result", "tool_use_id": "toolu-1"}, + ) + if err != nil { + t.Fatalf("SendSessionEventRaw() error = %v", err) + } + + want := []string{ + "GET /sessions/sess-1", + "GET /sessions/sess-1/events", + "GET /sessions/sess-1/events/stream", + "POST /sessions/sess-1/events", + } + if strings.Join(seen, "\n") != strings.Join(want, "\n") { + t.Fatalf("requests = %v, want %v", seen, want) + } +} + +func decodeJSONBody(t *testing.T, r *http.Request, out any) { + t.Helper() + if err := json.NewDecoder(r.Body).Decode(out); err != nil { + t.Fatalf("decode body: %v", err) + } +} + +func assertNoBody(t *testing.T, r *http.Request) { + t.Helper() + data, err := io.ReadAll(r.Body) + if err != nil { + t.Fatalf("read body: %v", err) + } + if len(data) != 0 { + t.Fatalf("unexpected body = %s", data) + } +} + +func assertQueryValue(t *testing.T, r *http.Request, key, want string) { + t.Helper() + if got := r.URL.Query().Get(key); got != want { + t.Fatalf("query %s = %q, want %q", key, got, want) + } +} + +func assertQueryValues(t *testing.T, r *http.Request, key string, want []string) { + t.Helper() + got := r.URL.Query()[key] + if len(got) != len(want) { + t.Fatalf("query %s = %v, want %v", key, got, want) + } + for i := range got { + if got[i] != want[i] { + t.Fatalf("query %s = %v, want %v", key, got, want) + } + } +} diff --git a/arkruntime/selfhosted/client_api.go b/arkruntime/selfhosted/client_api.go new file mode 100644 index 0000000..99aa054 --- /dev/null +++ b/arkruntime/selfhosted/client_api.go @@ -0,0 +1,481 @@ +// Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +// SPDX-License-Identifier: Apache-2.0 + +package selfhosted + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "io" + "net/http" + "path" + "strings" + "sync" + + "github.com/volcengine/ark-runtime-go/arkruntime" + "github.com/volcengine/ark-runtime-go/arkruntime/model" + "github.com/volcengine/ark-runtime-go/arkruntime/model/environment" + "github.com/volcengine/ark-runtime-go/arkruntime/model/session" +) + +// ClientAPI adapts arkruntime.Client to the self-hosted worker control-plane API. +type ClientAPI struct { + client *arkruntime.Client +} + +// NewClientAPI creates a self-hosted worker API backed by arkruntime.Client. +func NewClientAPI(client *arkruntime.Client) *ClientAPI { + return &ClientAPI{client: client} +} + +// PollWork polls one work item from the Environment queue. +func (a *ClientAPI) PollWork(ctx context.Context, req PollWorkRequest) (*WorkItem, error) { + if a == nil || a.client == nil { + return nil, errors.New("selfhosted: arkruntime client is nil") + } + item, err := a.client.PollWork(ctx, &environment.PollWorkRequest{ + EnvironmentID: req.EnvironmentID, + WorkerID: req.WorkerID, + MaxItems: req.MaxItems, + BlockMS: req.BlockMS, + ReclaimOlderThanMS: req.ReclaimOlderThanMS, + WorkerClientType: firstNonEmpty(req.WorkerClientType, DefaultWorkerClientType), + WorkerClientVersion: firstNonEmpty(req.WorkerClientVersion, DefaultWorkerClientVersion), + }) + if err != nil { + return nil, toWorkerAPIError(err) + } + return fromEnvironmentWorkItem(item), nil +} + +// AckWork acknowledges one claimed work item. +func (a *ClientAPI) AckWork(ctx context.Context, req AckWorkRequest) error { + if a == nil || a.client == nil { + return errors.New("selfhosted: arkruntime client is nil") + } + err := a.client.AckWork(ctx, &environment.AckWorkRequest{ + EnvironmentID: req.EnvironmentID, + WorkID: req.WorkID, + WorkerID: envOptString(req.WorkerID), + }) + return toWorkerAPIError(err) +} + +// HeartbeatWork refreshes a claimed work lease. +func (a *ClientAPI) HeartbeatWork(ctx context.Context, req HeartbeatWorkRequest) (*HeartbeatResponse, error) { + if a == nil || a.client == nil { + return nil, errors.New("selfhosted: arkruntime client is nil") + } + expectedLastHeartbeat := req.ExpectedLastHeartbeat + if expectedLastHeartbeat == "" { + expectedLastHeartbeat = ExpectedLastHeartbeatNoHeartbeat + } + resp, err := a.client.HeartbeatWork(ctx, &environment.HeartbeatWorkRequest{ + EnvironmentID: req.EnvironmentID, + WorkID: req.WorkID, + ExpectedLastHeartbeat: envOptString(expectedLastHeartbeat), + DesiredTTLSeconds: envOptInt64(req.DesiredTTLSeconds), + }) + if err != nil { + return nil, toWorkerAPIError(err) + } + if resp == nil { + return nil, nil + } + return &HeartbeatResponse{ + LastHeartbeat: resp.LastHeartbeat, + LeaseExtended: optionalBoolPtr(resp.LeaseExtended, true), + State: string(resp.State), + TTLSeconds: int(resp.TTLSeconds), + Type: string(resp.Type), + }, nil +} + +// StopWork releases or stops one claimed work item. +func (a *ClientAPI) StopWork(ctx context.Context, req StopWorkRequest) error { + if a == nil || a.client == nil { + return errors.New("selfhosted: arkruntime client is nil") + } + err := a.client.StopWork(ctx, &environment.StopWorkRequest{ + EnvironmentID: req.EnvironmentID, + WorkID: req.WorkID, + Force: envOptBool(req.Force), + }) + return toWorkerAPIError(err) +} + +// GetSession reads a session snapshot. +func (a *ClientAPI) GetSession(ctx context.Context, req GetSessionRequest) (*Session, error) { + if a == nil || a.client == nil { + return nil, errors.New("selfhosted: arkruntime client is nil") + } + out, err := a.client.GetSession(ctx, req.SessionID) + if err != nil { + return nil, toWorkerAPIError(err) + } + return fromRuntimeSession(out), nil +} + +// ListEvents lists session events in ascending order by default. +func (a *ClientAPI) ListEvents(ctx context.Context, req ListEventsRequest) (*ListEventsResponse, error) { + if a == nil || a.client == nil { + return nil, errors.New("selfhosted: arkruntime client is nil") + } + params := &session.SessionEventsListEventsParams{Types: req.Types} + if req.Page != "" { + params.Page = session.NewOptString(req.Page) + } + if req.CreatedAtGt != "" { + params.CreatedAtGt = session.NewOptString(req.CreatedAtGt) + } + if req.Limit > 0 { + params.Limit = session.NewOptInt32(int32(req.Limit)) + } + switch req.Order { + case EventListOrderDesc: + params.Order = session.NewOptListSessionsOrder(session.ListSessionsOrderDesc) + case EventListOrderAsc: + params.Order = session.NewOptListSessionsOrder(session.ListSessionsOrderAsc) + } + resp, err := a.client.ListSessionEvents(ctx, req.SessionID, params) + if err != nil { + return nil, toWorkerAPIError(err) + } + out := &ListEventsResponse{} + if resp == nil { + return out, nil + } + if next, ok := resp.NextPage.Get(); ok { + out.NextPage = next + } + out.Events = make([]Event, 0, len(resp.Events)) + for _, ev := range resp.Events { + converted, err := fromManagedAgentEvent(ev) + if err != nil { + return nil, err + } + out.Events = append(out.Events, converted) + } + return out, nil +} + +// StreamEvents opens a session event SSE stream. +func (a *ClientAPI) StreamEvents(ctx context.Context, req StreamEventsRequest) (*EventStream, error) { + if a == nil || a.client == nil { + return nil, errors.New("selfhosted: arkruntime client is nil") + } + streamCtx, streamCancel := context.WithCancel(ctx) + decoder, err := a.client.StreamSessionEventsWithParams(streamCtx, req.SessionID, &session.SessionEventsStreamEventsParams{ + EventDeltas: req.EventDeltas, + }) + if err != nil { + streamCancel() + return nil, toWorkerAPIError(err) + } + events := make(chan Event) + var mu sync.Mutex + var streamErr error + var closeOnce sync.Once + setErr := func(err error) { + mu.Lock() + defer mu.Unlock() + streamErr = err + } + stream := &EventStream{ + events: events, + close: func() error { + var err error + closeOnce.Do(func() { + streamCancel() + err = decoder.Close() + }) + return err + }, + err: func() error { + mu.Lock() + defer mu.Unlock() + return streamErr + }, + } + go func() { + defer close(events) + defer func() { _ = stream.Close() }() + for decoder.Next() { + event, err := fromStreamEvent(decoder.Event()) + if err != nil { + setErr(err) + return + } + select { + case <-streamCtx.Done(): + if ctx.Err() != nil { + setErr(ctx.Err()) + } + return + case events <- event: + } + } + if err := decoder.Err(); err != nil && streamCtx.Err() == nil { + setErr(toWorkerAPIError(err)) + } + }() + return stream, nil +} + +// SendEvent writes one user-side event back to the session. +func (a *ClientAPI) SendEvent(ctx context.Context, req SendEventRequest) error { + if a == nil || a.client == nil { + return errors.New("selfhosted: arkruntime client is nil") + } + return toWorkerAPIError(a.client.SendSessionEventRaw(ctx, req.SessionID, req.Event)) +} + +// ResolveSkill enriches a session skill reference with control-plane metadata. +func (a *ClientAPI) ResolveSkill(ctx context.Context, ref SkillRef) (SkillRef, error) { + if a == nil || a.client == nil { + return SkillRef{}, errors.New("selfhosted: arkruntime client is nil") + } + skillID := strings.TrimSpace(ref.IDValue()) + if skillID == "" { + return SkillRef{}, errors.New("skill id is required") + } + metadata, err := a.client.GetSkill(ctx, skillID) + if err != nil { + return SkillRef{}, toWorkerAPIError(err) + } + if metadata == nil { + return SkillRef{}, fmt.Errorf("skill metadata is empty: %s", skillID) + } + name, _ := metadata.Name.Get() + name = strings.TrimSpace(name) + if name == "" { + return SkillRef{}, fmt.Errorf("skill name is empty: %s", skillID) + } + + ref.Name = name + if ref.DisplayName == "" { + ref.DisplayName = strings.TrimSpace(metadata.DisplayTitle) + } + if ref.Type == "" { + ref.Type = strings.TrimSpace(metadata.Source) + } + if ref.Version == "" { + ref.Version = strings.TrimSpace(metadata.LatestVersion) + } + return ref, nil +} + +// OpenSkill opens a skill archive stream. +func (a *ClientAPI) OpenSkill(ctx context.Context, req OpenSkillRequest) (*SkillContent, error) { + if a == nil || a.client == nil { + return nil, errors.New("selfhosted: arkruntime client is nil") + } + if req.Skill.DownloadURL != "" { + return openSignedSkillURL(ctx, a.client.HTTPClient(), req.Skill.DownloadURL) + } + if strings.EqualFold(strings.TrimSpace(req.Skill.Type), skillTypeSkillHub) { + return a.openSkillHub(ctx, req.Skill) + } + content, err := a.client.OpenSkillContent(ctx, req.Skill.IDValue(), req.Skill.Version) + if err != nil { + return nil, toWorkerAPIError(err) + } + return &SkillContent{ + Body: content.Body, + ContentLength: content.ContentLength, + FileName: content.FileName, + ContentType: content.ContentType, + }, nil +} + +func fromEnvironmentWorkItem(item *environment.WorkItem) *WorkItem { + if item == nil { + return nil + } + latestHeartbeat := item.LatestHeartbeatValue() + acknowledgedAt, _ := item.AcknowledgedAt.Get() + secret, _ := item.Secret.Get() + startedAt, _ := item.StartedAt.Get() + stopRequestedAt, _ := item.StopRequestedAt.Get() + stoppedAt, _ := item.StoppedAt.Get() + data := WorkData{} + if item.Data.ID != "" || item.Data.Type != "" { + data = WorkData{ + Type: item.Data.Type, + ID: item.Data.ID, + } + } + tags := make([]WorkTag, 0, len(item.Tags)) + for _, tag := range item.Tags { + value, _ := tag.Value.Get() + tags = append(tags, WorkTag{ + Key: tag.Key, + Value: value, + }) + } + return &WorkItem{ + ID: item.ID, + AcknowledgedAt: acknowledgedAt, + CreatedAt: item.CreatedAt, + Data: data, + EnvironmentID: item.EnvironmentID, + LatestHeartbeatAt: latestHeartbeat, + Tags: tags, + Secret: secret, + StartedAt: startedAt, + State: string(item.State), + StopRequestedAt: stopRequestedAt, + StoppedAt: stoppedAt, + Type: string(item.Type), + SessionID: data.ID, + LastHeartbeat: latestHeartbeat, + } +} + +func fromRuntimeSession(in *session.Session) *Session { + if in == nil { + return nil + } + out := &Session{ID: in.ID} + data, err := json.Marshal(in) + if err != nil { + return out + } + if err := json.Unmarshal(data, out); err != nil { + out.ID = in.ID + return out + } + if out.ID == "" { + out.ID = in.ID + } + return out +} + +func fromStreamEvent(frame session.StreamEvent) (Event, error) { + if len(frame.RawPayload) > 0 { + return eventFromRaw(frame.RawPayload, frame.Type, frame.Data.GetID()) + } + return fromManagedAgentEvent(frame.Data) +} + +func fromManagedAgentEvent(ev session.ManagedAgentsSessionEvent) (Event, error) { + if unknown, ok := ev.(*session.ManagedAgentsUnknownSessionEvent); ok { + return eventFromRaw(unknown.RawPayload, unknown.Type(), unknown.GetID()) + } + data, err := json.Marshal(ev) + if err != nil { + return Event{}, err + } + return eventFromRaw(data, ev.Type(), ev.GetID()) +} + +func eventFromRaw(data []byte, eventType, eventID string) (Event, error) { + var out Event + if err := json.Unmarshal(data, &out); err != nil { + return Event{}, err + } + if out.Type == "" { + out.Type = eventType + } + if out.ID == "" { + out.ID = eventID + } + return out, nil +} + +func openSignedSkillURL(ctx context.Context, httpClient *http.Client, rawURL string) (*SkillContent, error) { + httpReq, err := http.NewRequestWithContext(ctx, http.MethodGet, rawURL, nil) + if err != nil { + return nil, err + } + if httpClient == nil { + httpClient = http.DefaultClient + } + resp, err := httpClient.Do(httpReq) + if err != nil { + return nil, err + } + if resp.StatusCode < http.StatusOK || resp.StatusCode >= http.StatusMultipleChoices { + defer resp.Body.Close() //nolint:errcheck // response body close errors are non-actionable + msg, _ := io.ReadAll(io.LimitReader(resp.Body, 4096)) + if len(msg) == 0 { + msg = []byte(http.StatusText(resp.StatusCode)) + } + return nil, &APIError{ + StatusCode: resp.StatusCode, + Message: strings.TrimSpace(string(msg)), + RequestID: resp.Header.Get("X-Request-Id"), + } + } + return &SkillContent{ + Body: resp.Body, + ContentLength: resp.ContentLength, + FileName: path.Base(resp.Request.URL.Path), + ContentType: resp.Header.Get("Content-Type"), + }, nil +} + +func toWorkerAPIError(err error) error { + if err == nil { + return nil + } + var apiErr *model.APIError + if errors.As(err, &apiErr) { + return &APIError{ + StatusCode: apiErr.HTTPStatusCode, + Message: apiErr.Message, + RequestID: apiErr.RequestId, + } + } + var reqErr *model.RequestError + if errors.As(err, &reqErr) && reqErr.HTTPStatusCode >= http.StatusBadRequest { + return &APIError{ + StatusCode: reqErr.HTTPStatusCode, + Message: fmt.Sprint(reqErr.Err), + RequestID: reqErr.RequestId, + } + } + return err +} + +func firstNonEmpty(values ...string) string { + for _, value := range values { + if value != "" { + return value + } + } + return "" +} + +func envOptString(value string) environment.OptString { + if value == "" { + return environment.OptString{} + } + return environment.NewOptString(value) +} + +func envOptInt64(value int) environment.OptInt64 { + if value <= 0 { + return environment.OptInt64{} + } + return environment.NewOptInt64(int64(value)) +} + +func envOptBool(value bool) environment.OptBool { + if !value { + return environment.OptBool{} + } + return environment.NewOptBool(value) +} + +func optionalBoolPtr(value bool, ok bool) *bool { + if !ok { + return nil + } + return &value +} + +var _ API = (*ClientAPI)(nil) +var _ EventStreamer = (*ClientAPI)(nil) diff --git a/arkruntime/selfhosted/client_api_test.go b/arkruntime/selfhosted/client_api_test.go new file mode 100644 index 0000000..447b179 --- /dev/null +++ b/arkruntime/selfhosted/client_api_test.go @@ -0,0 +1,234 @@ +// Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +// SPDX-License-Identifier: Apache-2.0 + +package selfhosted_test + +import ( + "context" + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/volcengine/ark-runtime-go/arkruntime" + "github.com/volcengine/ark-runtime-go/arkruntime/selfhosted" +) + +type roundTripFunc func(*http.Request) (*http.Response, error) + +const skillArchiveContentType = "application/zip" + +func (f roundTripFunc) RoundTrip(req *http.Request) (*http.Response, error) { + return f(req) +} + +func TestClientAPIConvertsSessionAndEvents(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + switch r.Method + " " + r.URL.Path { + case "GET /sessions/sess-1": + _, _ = w.Write([]byte(`{ + "id":"sess-1", + "type":"session", + "status":"idle", + "environment_id":"env-1", + "agent":{"skills":[{"type":"skill_hub","display_name":"demo","skill_id":"skill-1","version":"v1"}]}, + "created_at":"2026-08-10T00:00:00Z", + "updated_at":"2026-08-10T00:00:00Z" + }`)) + case "GET /sessions/sess-1/events": + _, _ = w.Write([]byte(`{ + "data":[{ + "id":"evt-1", + "type":"agent.tool_use", + "session_thread_id":"thread-1", + "tool_use_id":"toolu-1", + "name":"bash", + "input":{"command":"echo hi"}, + "evaluated_permission":"ask" + }], + "next_page":"page-2" + }`)) + default: + http.NotFound(w, r) + } + })) + defer server.Close() + + api := selfhosted.NewClientAPI(arkruntime.NewClientWithApiKey("test-api-key", arkruntime.WithBaseUrl(server.URL))) + session, err := api.GetSession(context.Background(), selfhosted.GetSessionRequest{SessionID: "sess-1"}) + if err != nil { + t.Fatalf("GetSession() error = %v", err) + } + skills := session.SkillRefs() + if len(skills) != 1 || skills[0].IDValue() != "skill-1" || skills[0].Type != "skill_hub" || skills[0].NameValue() != "demo" || skills[0].Version != "v1" { + t.Fatalf("skills = %+v", skills) + } + + events, err := api.ListEvents(context.Background(), selfhosted.ListEventsRequest{ + SessionID: "sess-1", + Limit: 100, + Order: selfhosted.EventListOrderAsc, + }) + if err != nil { + t.Fatalf("ListEvents() error = %v", err) + } + if events.NextPage != "page-2" || len(events.Events) != 1 { + t.Fatalf("events = %+v", events) + } + event := events.Events[0] + if event.ToolUseID != "toolu-1" || event.EvaluatedPermission != selfhosted.PermissionAsk { + t.Fatalf("event = %+v", event) + } +} + +func TestClientAPIResolveSkillUsesControlPlaneMetadata(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodGet || r.URL.Path != "/skills/skill-1" { + http.NotFound(w, r) + return + } + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{ + "id":"skill-1", + "object":"skill", + "created_at":1786506774, + "updated_at":1786506774, + "name":"canonical-skill-name", + "display_title":"Canonical skill", + "source":"skill_hub", + "latest_version":"1.0.0" + }`)) + })) + defer server.Close() + + api := selfhosted.NewClientAPI(arkruntime.NewClientWithApiKey("test-api-key", arkruntime.WithBaseUrl(server.URL))) + resolved, err := api.ResolveSkill(context.Background(), selfhosted.SkillRef{SkillID: "skill-1"}) + if err != nil { + t.Fatalf("ResolveSkill() error = %v", err) + } + if resolved.Name != "canonical-skill-name" || resolved.DisplayName != "Canonical skill" || resolved.Type != "skill_hub" || resolved.Version != "1.0.0" { + t.Fatalf("resolved skill = %+v", resolved) + } +} + +func TestClientAPIOpenSkillHubUsesMetadataAndVersionedDownload(t *testing.T) { + const archive = "skill-hub-zip" + var requests []string + httpClient := &http.Client{ + Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) { + if got := req.Header.Get("Authorization"); got != "" { + t.Fatalf("Authorization = %q", got) + } + requests = append(requests, req.URL.String()) + body := "" + contentType := "application/json" + switch { + case req.URL.Path == "/v1/skills" && req.URL.Query().Get("skillIds") == "skill-1": + body = `{"Skills":[{"Id":"other-skill","Slug":"wrong/slug"},{"Id":"skill-1","Slug":"volcengine/ark/demo"}],"Total":2}` + case req.URL.Path == "/v1/skills/download/volcengine/ark/demo" && req.URL.Query().Get("version") == "1.0.0": + body = archive + contentType = skillArchiveContentType + default: + t.Fatalf("unexpected request: %s", req.URL) + } + return &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{contentType}}, + Body: io.NopCloser(strings.NewReader(body)), + ContentLength: int64(len(body)), + Request: req, + }, nil + }), + } + api := selfhosted.NewClientAPI(arkruntime.NewClientWithApiKey("test-api-key", arkruntime.WithHTTPClient(httpClient))) + + content, err := api.OpenSkill(context.Background(), selfhosted.OpenSkillRequest{ + Skill: selfhosted.SkillRef{Type: "skill_hub", SkillID: "skill-1", Version: "1.0.0"}, + }) + if err != nil { + t.Fatalf("OpenSkill() error = %v", err) + } + defer content.Body.Close() + data, err := io.ReadAll(content.Body) + if err != nil { + t.Fatalf("ReadAll() error = %v", err) + } + if len(requests) != 2 { + t.Fatalf("requests = %v", requests) + } + if string(data) != archive || content.ContentType != skillArchiveContentType || content.FileName != "demo" { + t.Fatalf("content = %q type=%q file=%q", data, content.ContentType, content.FileName) + } +} + +func TestClientAPIHeartbeatDefaultsExpectedLastHeartbeat(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost || r.URL.Path != "/environments/env-1/work/work-1/heartbeat" { + http.NotFound(w, r) + return + } + if got := r.URL.Query().Get("expected_last_heartbeat"); got != selfhosted.ExpectedLastHeartbeatNoHeartbeat { + t.Fatalf("expected_last_heartbeat = %q", got) + } + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"last_heartbeat":"2026-08-11T00:00:00Z","lease_extended":true,"state":"active","ttl_seconds":30,"type":"work_heartbeat"}`)) + })) + defer server.Close() + + api := selfhosted.NewClientAPI(arkruntime.NewClientWithApiKey("test-api-key", arkruntime.WithBaseUrl(server.URL))) + resp, err := api.HeartbeatWork(context.Background(), selfhosted.HeartbeatWorkRequest{ + EnvironmentID: "env-1", + WorkID: "work-1", + DesiredTTLSeconds: 30, + }) + if err != nil { + t.Fatalf("HeartbeatWork() error = %v", err) + } + if resp.LastHeartbeat != "2026-08-11T00:00:00Z" || resp.State != selfhosted.WorkStateActive { + t.Fatalf("HeartbeatWork() = %+v", resp) + } +} + +func TestClientAPIOpenSkillDownloadURLUsesConfiguredHTTPClient(t *testing.T) { + const archive = "zip-bytes" + called := false + httpClient := &http.Client{ + Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) { + called = true + if got := req.URL.String(); got != "https://signed.example.com/skill.zip" { + t.Fatalf("url = %q", got) + } + if got := req.Header.Get("Authorization"); got != "" { + t.Fatalf("Authorization = %q", got) + } + return &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{skillArchiveContentType}}, + Body: io.NopCloser(strings.NewReader(archive)), + ContentLength: int64(len(archive)), + Request: req, + }, nil + }), + } + api := selfhosted.NewClientAPI(arkruntime.NewClientWithApiKey("test-api-key", arkruntime.WithHTTPClient(httpClient))) + + content, err := api.OpenSkill(context.Background(), selfhosted.OpenSkillRequest{ + Skill: selfhosted.SkillRef{DownloadURL: "https://signed.example.com/skill.zip"}, + }) + if err != nil { + t.Fatalf("OpenSkill() error = %v", err) + } + defer content.Body.Close() + data, err := io.ReadAll(content.Body) + if err != nil { + t.Fatalf("ReadAll() error = %v", err) + } + if !called { + t.Fatal("configured HTTP client was not used") + } + if string(data) != archive || content.ContentType != skillArchiveContentType || content.FileName != "skill.zip" { + t.Fatalf("content = %q type=%q file=%q", data, content.ContentType, content.FileName) + } +} diff --git a/arkruntime/selfhosted/defaults.go b/arkruntime/selfhosted/defaults.go new file mode 100644 index 0000000..7a9756b --- /dev/null +++ b/arkruntime/selfhosted/defaults.go @@ -0,0 +1,12 @@ +// Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +// SPDX-License-Identifier: Apache-2.0 +package selfhosted + +import "time" + +const ( + // DefaultMaxIdle 是 session worker 空闲退出的默认时间。 + DefaultMaxIdle = 60 * time.Second + // DefaultToolTimeout 是单次本地工具执行的默认超时。 + DefaultToolTimeout = 120 * time.Second +) diff --git a/arkruntime/selfhosted/doc.go b/arkruntime/selfhosted/doc.go new file mode 100644 index 0000000..133371a --- /dev/null +++ b/arkruntime/selfhosted/doc.go @@ -0,0 +1,4 @@ +// Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +// SPDX-License-Identifier: Apache-2.0 +// Package ark 提供 self-hosted worker 的公共 SDK 风格入口。 +package selfhosted diff --git a/arkruntime/selfhosted/envinit/initializer.go b/arkruntime/selfhosted/envinit/initializer.go new file mode 100644 index 0000000..898ee7d --- /dev/null +++ b/arkruntime/selfhosted/envinit/initializer.go @@ -0,0 +1,394 @@ +// Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +// SPDX-License-Identifier: Apache-2.0 +// Package envinit 准备 self-hosted worker 的 session 执行环境。 +package envinit + +import ( + "archive/tar" + "archive/zip" + "compress/gzip" + "context" + "errors" + "fmt" + "io" + "log" + "os" + "path/filepath" + "strings" + + "github.com/volcengine/ark-runtime-go/arkruntime/internal/selfhostedlog" + selfhosted "github.com/volcengine/ark-runtime-go/arkruntime/selfhosted" +) + +const ( + defaultMaxArchiveBytes = 128 << 20 + defaultMaxExtractedBytes = 512 << 20 + defaultMaxArchiveEntries = 10_000 + parentDirName = ".." +) + +// Options 是 session 环境初始化配置。 +type Options struct { + Workdir string + SkillsDir string + MaxArchiveBytes int64 + MaxExtractedBytes int64 + MaxArchiveEntries int + Logger *log.Logger +} + +// Initializer 执行 workdir 与 skills 初始化。 +type Initializer struct { + api selfhosted.API + opts Options + logger *selfhostedlog.Logger +} + +// New 创建环境初始化器。 +func New(api selfhosted.API, opts Options) *Initializer { + if opts.MaxArchiveBytes <= 0 { + opts.MaxArchiveBytes = defaultMaxArchiveBytes + } + if opts.MaxExtractedBytes <= 0 { + opts.MaxExtractedBytes = defaultMaxExtractedBytes + } + if opts.MaxArchiveEntries <= 0 { + opts.MaxArchiveEntries = defaultMaxArchiveEntries + } + if opts.SkillsDir == "" { + opts.SkillsDir = filepath.Join(opts.Workdir, "skills") + } + return &Initializer{api: api, opts: opts, logger: selfhostedlog.New(opts.Logger)} +} + +// Setup 创建 workdir 并安装 session 绑定的 skills。 +func (i *Initializer) Setup(ctx context.Context, session *selfhosted.Session) error { + if session == nil { + return errors.New("session must not be nil") + } + if i.opts.Workdir == "" { + return errors.New("workdir must not be empty") + } + if err := os.MkdirAll(i.opts.Workdir, 0o755); err != nil { + return fmt.Errorf("create workdir: %w", err) + } + if err := os.MkdirAll(i.opts.SkillsDir, 0o755); err != nil { + return fmt.Errorf("create skills dir: %w", err) + } + for _, skill := range session.SkillRefs() { + if err := i.installSkill(ctx, session.ID, skill); err != nil { + i.logger.Warn("failed to install skill", "session_id", session.ID, "skill", skill.NameValue(), "version", skill.Version, "err", err) + } + } + return nil +} + +func (i *Initializer) installSkill(ctx context.Context, sessionID string, skill selfhosted.SkillRef) error { + if resolver, ok := i.api.(selfhosted.SkillResolver); ok && strings.TrimSpace(skill.IDValue()) != "" { + resolved, err := resolver.ResolveSkill(ctx, skill) + if err != nil { + return fmt.Errorf("resolve skill %s: %w", skill.IDValue(), err) + } + skill = resolved + } + name, err := safeSkillDirName(skill) + if err != nil { + return err + } + i.logger.Info("install skill", "session_id", sessionID, "skill", name, "version", skill.Version) + content, err := i.api.OpenSkill(ctx, selfhosted.OpenSkillRequest{SessionID: sessionID, Skill: skill}) + if err != nil { + return fmt.Errorf("download skill %s: %w", name, err) + } + if content == nil || content.Body == nil { + return fmt.Errorf("download skill %s: empty content", name) + } + defer func() { _ = content.Body.Close() }() + + archivePath, err := i.copyArchive(name, content.Body) + if err != nil { + return err + } + defer func() { _ = os.Remove(archivePath) }() + + tmp, err := os.MkdirTemp(i.opts.SkillsDir, "."+name+"-*") + if err != nil { + return fmt.Errorf("create skill temp dir: %w", err) + } + committed := false + defer func() { + if !committed { + _ = os.RemoveAll(tmp) + } + }() + if err := i.extractArchive(archivePath, tmp); err != nil { + return err + } + source, err := installSourceDir(tmp) + if err != nil { + return err + } + target := filepath.Join(i.opts.SkillsDir, name) + backup, err := replaceSkillDir(source, target) + if err != nil { + return fmt.Errorf("commit skill %s: %w", name, err) + } + committed = true + if source != tmp { + _ = os.RemoveAll(tmp) + } + if backup != "" { + if err := os.RemoveAll(backup); err != nil { + i.logger.Warn("remove old skill backup failed", "session_id", sessionID, "skill", name, "path", backup, "err", err) + } + } + return nil +} + +func replaceSkillDir(source, target string) (string, error) { + if _, err := os.Lstat(target); errors.Is(err, os.ErrNotExist) { + return "", os.Rename(source, target) + } else if err != nil { + return "", err + } + backup, err := os.MkdirTemp(filepath.Dir(target), "."+filepath.Base(target)+"-backup-*") + if err != nil { + return "", err + } + if err := os.Remove(backup); err != nil { + return "", err + } + if err := os.Rename(target, backup); err != nil { + return "", err + } + if err := os.Rename(source, target); err != nil { + if rollbackErr := os.Rename(backup, target); rollbackErr != nil { + return "", fmt.Errorf("replace skill: %w; rollback: %v", err, rollbackErr) + } + return "", err + } + return backup, nil +} + +func safeSkillDirName(skill selfhosted.SkillRef) (string, error) { + for _, candidate := range []string{skill.Name, skill.DisplayName, skill.IDValue()} { + name, err := safeSkillName(candidate) + if err == nil { + return name, nil + } + } + return "", fmt.Errorf("invalid skill name: %q", skill.Name) +} + +func (i *Initializer) copyArchive(name string, body io.Reader) (string, error) { + tmp, err := os.CreateTemp("", "ark-skill-"+name+"-*") + if err != nil { + return "", err + } + tmpName := tmp.Name() + defer func() { _ = tmp.Close() }() + n, err := io.Copy(tmp, io.LimitReader(body, i.opts.MaxArchiveBytes+1)) + if err != nil { + _ = os.Remove(tmpName) + return "", fmt.Errorf("copy skill archive: %w", err) + } + if n > i.opts.MaxArchiveBytes { + _ = os.Remove(tmpName) + return "", fmt.Errorf("skill archive too large: %d bytes", n) + } + return tmpName, nil +} + +func (i *Initializer) extractArchive(archivePath, dst string) error { + f, err := os.Open(archivePath) + if err != nil { + return err + } + defer func() { _ = f.Close() }() + magic := make([]byte, 4) + n, _ := io.ReadFull(f, magic) + if _, err := f.Seek(0, io.SeekStart); err != nil { + return err + } + switch { + case n >= 4 && magic[0] == 'P' && magic[1] == 'K': + return i.extractZip(archivePath, dst) + case n >= 2 && magic[0] == 0x1f && magic[1] == 0x8b: + return i.extractTarGz(f, dst) + default: + return errors.New("unsupported skill archive format") + } +} + +func (i *Initializer) extractZip(archivePath, dst string) error { + zr, err := zip.OpenReader(archivePath) + if err != nil { + return fmt.Errorf("open zip: %w", err) + } + defer func() { _ = zr.Close() }() + if len(zr.File) > i.opts.MaxArchiveEntries { + return fmt.Errorf("skill archive contains too many entries: %d", len(zr.File)) + } + var total int64 + for _, file := range zr.File { + target, err := safeJoin(dst, file.Name) + if err != nil { + return err + } + if file.FileInfo().IsDir() { + if err := os.MkdirAll(target, 0o755); err != nil { + return err + } + continue + } + if file.Mode()&os.ModeType != 0 { + return fmt.Errorf("unsupported zip entry type: %s", file.Name) + } + remaining := i.opts.MaxExtractedBytes - total + if remaining < 0 || file.UncompressedSize64 > uint64(remaining) { + return fmt.Errorf("skill extracted content too large: more than %d bytes", i.opts.MaxExtractedBytes) + } + if err := os.MkdirAll(filepath.Dir(target), 0o755); err != nil { + return err + } + rc, err := file.Open() + if err != nil { + return err + } + written, err := writeFile(target, rc, file.Mode(), remaining) + if err != nil { + _ = rc.Close() + return err + } + _ = rc.Close() + total += written + } + return nil +} + +func (i *Initializer) extractTarGz(r io.Reader, dst string) error { + gz, err := gzip.NewReader(r) + if err != nil { + return fmt.Errorf("open gzip: %w", err) + } + defer func() { _ = gz.Close() }() + tr := tar.NewReader(gz) + var total int64 + entries := 0 + for { + header, err := tr.Next() + if errors.Is(err, io.EOF) { + return nil + } + if err != nil { + return fmt.Errorf("read tar: %w", err) + } + entries++ + if entries > i.opts.MaxArchiveEntries { + return fmt.Errorf("skill archive contains too many entries: %d", entries) + } + target, err := safeJoin(dst, header.Name) + if err != nil { + return err + } + switch header.Typeflag { + case tar.TypeDir: + if err := os.MkdirAll(target, 0o755); err != nil { + return err + } + case tar.TypeReg, 0: + remaining := i.opts.MaxExtractedBytes - total + if remaining < 0 || header.Size < 0 || header.Size > remaining { + return fmt.Errorf("skill extracted content too large: more than %d bytes", i.opts.MaxExtractedBytes) + } + if err := os.MkdirAll(filepath.Dir(target), 0o755); err != nil { + return err + } + written, err := writeFile(target, tr, os.FileMode(header.Mode).Perm(), remaining) + if err != nil { + return err + } + total += written + default: + return fmt.Errorf("unsupported tar entry type: %s", header.Name) + } + } +} + +func writeFile(target string, r io.Reader, mode os.FileMode, maxBytes int64) (int64, error) { + if mode == 0 { + mode = 0o644 + } + f, err := os.OpenFile(target, os.O_CREATE|os.O_TRUNC|os.O_WRONLY, mode.Perm()) + if err != nil { + return 0, err + } + written, err := io.Copy(f, io.LimitReader(r, maxBytes+1)) + if written > maxBytes { + _ = f.Close() + _ = os.Remove(target) + return written, fmt.Errorf("skill extracted content too large: more than %d bytes", maxBytes) + } + if err != nil { + _ = f.Close() + return written, err + } + if err := f.Close(); err != nil { + return written, err + } + return written, os.Chmod(target, mode.Perm()) +} + +func safeSkillName(name string) (string, error) { + name = strings.TrimSpace(name) + if name == "" { + return "", errors.New("skill name must not be empty") + } + if !isSafeNameComponent(name) { + return "", fmt.Errorf("invalid skill name: %q", name) + } + return name, nil +} + +func isSafeNameComponent(name string) bool { + if name == "" || name == "." || name == parentDirName { + return false + } + for _, c := range name { + if c >= 'a' && c <= 'z' || c >= 'A' && c <= 'Z' || c >= '0' && c <= '9' { + continue + } + if c == '.' || c == '_' || c == '-' { + continue + } + return false + } + return true +} + +func safeJoin(root, name string) (string, error) { + if name == "" || filepath.IsAbs(name) { + return "", fmt.Errorf("invalid archive path: %s", name) + } + clean := filepath.Clean(name) + if clean == "." || clean == parentDirName || strings.HasPrefix(clean, parentDirName+string(os.PathSeparator)) { + return "", fmt.Errorf("archive path escapes skill dir: %s", name) + } + target := filepath.Join(root, clean) + rel, err := filepath.Rel(root, target) + if err != nil || rel == parentDirName || strings.HasPrefix(rel, parentDirName+string(os.PathSeparator)) || filepath.IsAbs(rel) { + return "", fmt.Errorf("archive path escapes skill dir: %s", name) + } + return target, nil +} + +func installSourceDir(tmp string) (string, error) { + entries, err := os.ReadDir(tmp) + if err != nil { + return "", err + } + if len(entries) != 1 || !entries[0].IsDir() { + return tmp, nil + } + return filepath.Join(tmp, entries[0].Name()), nil +} diff --git a/arkruntime/selfhosted/envinit/initializer_test.go b/arkruntime/selfhosted/envinit/initializer_test.go new file mode 100644 index 0000000..546bf81 --- /dev/null +++ b/arkruntime/selfhosted/envinit/initializer_test.go @@ -0,0 +1,315 @@ +// Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +// SPDX-License-Identifier: Apache-2.0 +package envinit + +import ( + "archive/zip" + "bytes" + "context" + "io" + "os" + "path/filepath" + "strings" + "testing" + + selfhosted "github.com/volcengine/ark-runtime-go/arkruntime/selfhosted" +) + +type fakeAPI struct { + data []byte +} + +func (f fakeAPI) PollWork(context.Context, selfhosted.PollWorkRequest) (*selfhosted.WorkItem, error) { + return nil, nil +} +func (f fakeAPI) AckWork(context.Context, selfhosted.AckWorkRequest) error { return nil } +func (f fakeAPI) HeartbeatWork(context.Context, selfhosted.HeartbeatWorkRequest) (*selfhosted.HeartbeatResponse, error) { + return &selfhosted.HeartbeatResponse{}, nil +} +func (f fakeAPI) StopWork(context.Context, selfhosted.StopWorkRequest) error { return nil } +func (f fakeAPI) GetSession(context.Context, selfhosted.GetSessionRequest) (*selfhosted.Session, error) { + return nil, nil +} +func (f fakeAPI) ListEvents(context.Context, selfhosted.ListEventsRequest) (*selfhosted.ListEventsResponse, error) { + return &selfhosted.ListEventsResponse{}, nil +} +func (f fakeAPI) SendEvent(context.Context, selfhosted.SendEventRequest) error { return nil } +func (f fakeAPI) OpenSkill(context.Context, selfhosted.OpenSkillRequest) (*selfhosted.SkillContent, error) { + return &selfhosted.SkillContent{Body: io.NopCloser(bytes.NewReader(f.data)), ContentLength: int64(len(f.data))}, nil +} + +type nilSkillAPI struct { + fakeAPI +} + +func (f nilSkillAPI) OpenSkill(context.Context, selfhosted.OpenSkillRequest) (*selfhosted.SkillContent, error) { + return nil, nil +} + +type resolvingAPI struct { + fakeAPI + resolved selfhosted.SkillRef +} + +func (f resolvingAPI) ResolveSkill(context.Context, selfhosted.SkillRef) (selfhosted.SkillRef, error) { + return f.resolved, nil +} + +func TestSetupRejectsNilSession(t *testing.T) { + init := New(fakeAPI{}, Options{Workdir: t.TempDir()}) + if err := init.Setup(context.Background(), nil); err == nil { + t.Fatal("expected nil session error") + } +} + +func TestSetupSkipsEmptySkillContent(t *testing.T) { + root := t.TempDir() + init := New(nilSkillAPI{}, Options{Workdir: root}) + session := &selfhosted.Session{ID: "sess_1", Skills: []selfhosted.SkillRef{{Name: "demo", ID: "sk_1"}}} + if err := init.Setup(context.Background(), session); err != nil { + t.Fatal(err) + } + if _, err := os.Stat(filepath.Join(root, "skills", "demo")); !os.IsNotExist(err) { + t.Fatalf("skill dir should not be installed, err=%v", err) + } +} + +func TestSetupInstallsZipSkill(t *testing.T) { + var buf bytes.Buffer + zw := zip.NewWriter(&buf) + w, err := zw.Create("SKILL.md") + if err != nil { + t.Fatal(err) + } + if _, err := w.Write([]byte("hello")); err != nil { + t.Fatal(err) + } + if err := zw.Close(); err != nil { + t.Fatal(err) + } + + root := t.TempDir() + init := New(fakeAPI{data: buf.Bytes()}, Options{Workdir: root}) + session := &selfhosted.Session{ID: "sess_1", Skills: []selfhosted.SkillRef{{Name: "demo", ID: "sk_1"}}} + if err := init.Setup(context.Background(), session); err != nil { + t.Fatal(err) + } + got, err := os.ReadFile(filepath.Join(root, "skills", "demo", "SKILL.md")) + if err != nil { + t.Fatal(err) + } + if string(got) != "hello" { + t.Fatalf("skill = %q", got) + } +} + +func TestReplaceSkillDirRollsBackWhenCommitFails(t *testing.T) { + root := t.TempDir() + target := filepath.Join(root, "demo") + if err := os.MkdirAll(target, 0o755); err != nil { + t.Fatal(err) + } + oldPath := filepath.Join(target, "SKILL.md") + if err := os.WriteFile(oldPath, []byte("old"), 0o600); err != nil { + t.Fatal(err) + } + if _, err := replaceSkillDir(filepath.Join(root, "missing"), target); err == nil { + t.Fatal("expected commit failure") + } + data, err := os.ReadFile(oldPath) + if err != nil { + t.Fatal(err) + } + if string(data) != "old" { + t.Fatalf("old skill data=%q", data) + } +} + +func TestSetupInstallsSkillUnderResolvedMetadataName(t *testing.T) { + var buf bytes.Buffer + zw := zip.NewWriter(&buf) + w, err := zw.Create("SKILL.md") + if err != nil { + t.Fatal(err) + } + if _, err := w.Write([]byte("hello")); err != nil { + t.Fatal(err) + } + if err := zw.Close(); err != nil { + t.Fatal(err) + } + + root := t.TempDir() + api := resolvingAPI{ + fakeAPI: fakeAPI{data: buf.Bytes()}, + resolved: selfhosted.SkillRef{ + Name: "canonical-skill-name", + SkillID: "skill-1", + Type: "custom", + Version: "1", + }, + } + init := New(api, Options{Workdir: root}) + session := &selfhosted.Session{ID: "sess_1", Skills: []selfhosted.SkillRef{{SkillID: "skill-1", Type: "custom", Version: "1"}}} + if err := init.Setup(context.Background(), session); err != nil { + t.Fatal(err) + } + if _, err := os.Stat(filepath.Join(root, "skills", "canonical-skill-name", "SKILL.md")); err != nil { + t.Fatal(err) + } + if _, err := os.Stat(filepath.Join(root, "skills", "skill-1")); !os.IsNotExist(err) { + t.Fatalf("skill id fallback directory should not exist, err=%v", err) + } +} + +func TestSetupFallsBackToSkillIDForUnsafeName(t *testing.T) { + var buf bytes.Buffer + zw := zip.NewWriter(&buf) + w, err := zw.Create("SKILL.md") + if err != nil { + t.Fatal(err) + } + if _, err := w.Write([]byte("hello")); err != nil { + t.Fatal(err) + } + if err := zw.Close(); err != nil { + t.Fatal(err) + } + + root := t.TempDir() + init := New(fakeAPI{data: buf.Bytes()}, Options{Workdir: root}) + session := &selfhosted.Session{ID: "sess_1", Skills: []selfhosted.SkillRef{{Name: "../demo", ID: "sk_1"}}} + if err := init.Setup(context.Background(), session); err != nil { + t.Fatal(err) + } + if _, err := os.Stat(filepath.Join(root, "skills", "sk_1", "SKILL.md")); err != nil { + t.Fatal(err) + } +} + +func TestSetupFlattensSingleRootZipAndPreservesExecutableMode(t *testing.T) { + var buf bytes.Buffer + zw := zip.NewWriter(&buf) + skillHeader := &zip.FileHeader{Name: "wrapped/SKILL.md", Method: zip.Deflate} + skillHeader.SetMode(0o644) + w, err := zw.CreateHeader(skillHeader) + if err != nil { + t.Fatal(err) + } + if _, err := w.Write([]byte("hello")); err != nil { + t.Fatal(err) + } + runHeader := &zip.FileHeader{Name: "wrapped/bin/run.sh", Method: zip.Deflate} + runHeader.SetMode(0o755) + w, err = zw.CreateHeader(runHeader) + if err != nil { + t.Fatal(err) + } + if _, err := w.Write([]byte("#!/bin/sh\n")); err != nil { + t.Fatal(err) + } + if err := zw.Close(); err != nil { + t.Fatal(err) + } + + root := t.TempDir() + init := New(fakeAPI{data: buf.Bytes()}, Options{Workdir: root}) + session := &selfhosted.Session{ID: "sess_1", Skills: []selfhosted.SkillRef{{Name: "demo", ID: "sk_1"}}} + if err := init.Setup(context.Background(), session); err != nil { + t.Fatal(err) + } + if _, err := os.Stat(filepath.Join(root, "skills", "demo", "SKILL.md")); err != nil { + t.Fatal(err) + } + if _, err := os.Stat(filepath.Join(root, "skills", "demo", "wrapped", "SKILL.md")); !os.IsNotExist(err) { + t.Fatalf("single root directory was not flattened, err=%v", err) + } + stat, err := os.Stat(filepath.Join(root, "skills", "demo", "bin", "run.sh")) + if err != nil { + t.Fatal(err) + } + if stat.Mode().Perm()&0o111 == 0 { + t.Fatalf("run.sh mode = %s, want executable bit", stat.Mode().Perm()) + } +} + +func TestSetupSkipsZipTraversal(t *testing.T) { + var buf bytes.Buffer + zw := zip.NewWriter(&buf) + if _, err := zw.Create("../escape"); err != nil { + t.Fatal(err) + } + if err := zw.Close(); err != nil { + t.Fatal(err) + } + root := t.TempDir() + init := New(fakeAPI{data: buf.Bytes()}, Options{Workdir: root}) + session := &selfhosted.Session{ID: "sess_1", Skills: []selfhosted.SkillRef{{Name: "demo", ID: "sk_1"}}} + if err := init.Setup(context.Background(), session); err != nil { + t.Fatal(err) + } + if _, err := os.Stat(filepath.Join(root, "escape")); !os.IsNotExist(err) { + t.Fatalf("archive escaped workdir, err=%v", err) + } + if _, err := os.Stat(filepath.Join(root, "skills", "demo")); !os.IsNotExist(err) { + t.Fatalf("skill dir should not be installed, err=%v", err) + } +} + +func TestSafeSkillNameRejectsUnsafeComponent(t *testing.T) { + tests := []string{ + "../escape", + `nested\escape`, + ".", + "..", + "demo\nname", + "demo name", + "中文", + } + for _, tt := range tests { + if _, err := safeSkillName(tt); err == nil { + t.Fatalf("safeSkillName(%q) succeeded", tt) + } + } +} + +func TestExtractZipEnforcesActualExtractedBytes(t *testing.T) { + target := filepath.Join(t.TempDir(), "large.txt") + written, err := writeFile(target, bytes.NewBufferString("0123456789"), 0o600, 4) + if err == nil || !strings.Contains(err.Error(), "skill extracted content too large") { + t.Fatalf("writeFile() bytes=%d error=%v", written, err) + } + if _, err := os.Stat(target); !os.IsNotExist(err) { + t.Fatalf("partial output was not removed: %v", err) + } +} + +func TestExtractZipEnforcesEntryLimit(t *testing.T) { + var buffer bytes.Buffer + writer := zip.NewWriter(&buffer) + for _, name := range []string{"one.txt", "two.txt"} { + entry, err := writer.Create(name) + if err != nil { + t.Fatal(err) + } + if _, err := entry.Write([]byte(name)); err != nil { + t.Fatal(err) + } + } + if err := writer.Close(); err != nil { + t.Fatal(err) + } + path := filepath.Join(t.TempDir(), "skill.zip") + if err := os.WriteFile(path, buffer.Bytes(), 0o600); err != nil { + t.Fatal(err) + } + + initializer := New(fakeAPI{}, Options{ + Workdir: t.TempDir(), + MaxArchiveEntries: 1, + }) + err := initializer.extractZip(path, t.TempDir()) + if err == nil || !strings.Contains(err.Error(), "too many entries") { + t.Fatalf("extractZip() error = %v", err) + } +} diff --git a/arkruntime/selfhosted/errors.go b/arkruntime/selfhosted/errors.go new file mode 100644 index 0000000..072886a --- /dev/null +++ b/arkruntime/selfhosted/errors.go @@ -0,0 +1,58 @@ +// Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +// SPDX-License-Identifier: Apache-2.0 +package selfhosted + +import ( + "context" + "errors" + "net/http" +) + +// ClassifyWorkerError 把底层请求错误转换成稳定的 WorkerError。 +func ClassifyWorkerError(err error) *WorkerError { + if err == nil { + return nil + } + out := &WorkerError{ + Kind: WorkerErrorKindNetwork, + Message: err.Error(), + Err: err, + } + if errors.Is(err, context.DeadlineExceeded) { + out.Kind = WorkerErrorKindTimeout + out.Retryable = true + return out + } + if errors.Is(err, context.Canceled) { + out.Kind = WorkerErrorKindTimeout + return out + } + var apiErr *APIError + if !errors.As(err, &apiErr) { + out.Retryable = true + return out + } + out.RequestID = apiErr.RequestID + switch apiErr.StatusCode { + case http.StatusUnauthorized: + out.Kind = WorkerErrorKindAuth + case http.StatusForbidden: + out.Kind = WorkerErrorKindPermission + case http.StatusConflict, http.StatusPreconditionFailed: + out.Kind = WorkerErrorKindLeaseConflict + case http.StatusRequestTimeout: + out.Kind = WorkerErrorKindTimeout + out.Retryable = true + case http.StatusTooManyRequests: + out.Kind = WorkerErrorKindRateLimit + out.Retryable = true + default: + if apiErr.StatusCode >= 500 { + out.Kind = WorkerErrorKindNetwork + out.Retryable = true + } else { + out.Kind = WorkerErrorKindInvalidResponse + } + } + return out +} diff --git a/arkruntime/selfhosted/session_event.go b/arkruntime/selfhosted/session_event.go new file mode 100644 index 0000000..c4cbc10 --- /dev/null +++ b/arkruntime/selfhosted/session_event.go @@ -0,0 +1,43 @@ +// Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +// SPDX-License-Identifier: Apache-2.0 +package selfhosted + +import "context" + +// EventStreamer 是支持 session SSE 事件流的可选 API 能力。 +type EventStreamer interface { + StreamEvents(ctx context.Context, req StreamEventsRequest) (*EventStream, error) +} + +// EventStream 是 session event stream 的读取句柄。 +type EventStream struct { + events <-chan Event + close func() error + err func() error +} + +// Events 返回事件通道。 +func (s *EventStream) Events() <-chan Event { + if s == nil { + ch := make(chan Event) + close(ch) + return ch + } + return s.events +} + +// Close 关闭事件流。 +func (s *EventStream) Close() error { + if s == nil || s.close == nil { + return nil + } + return s.close() +} + +// Err 返回事件流结束后的错误。 +func (s *EventStream) Err() error { + if s == nil || s.err == nil { + return nil + } + return s.err() +} diff --git a/arkruntime/selfhosted/session_tool_runner.go b/arkruntime/selfhosted/session_tool_runner.go new file mode 100644 index 0000000..2ad6501 --- /dev/null +++ b/arkruntime/selfhosted/session_tool_runner.go @@ -0,0 +1,1110 @@ +// Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +// SPDX-License-Identifier: Apache-2.0 +package selfhosted + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "log" + "math/rand" + "net/http" + "runtime/debug" + "sync" + "time" + + "github.com/volcengine/ark-runtime-go/arkruntime/internal/selfhostedlog" + "github.com/volcengine/ark-runtime-go/arkruntime/toolset" +) + +const ( + sessionRunnerStreamBackoffStart = 500 * time.Millisecond + sessionRunnerStreamBackoffCap = 10 * time.Second + sessionRunnerStreamHealthyAfter = 30 * time.Second + sessionRunnerSendRetries = 3 + sessionRunnerResultsBuffer = 32 + + defaultSessionRunnerSendTimeout = 15 * time.Second + defaultSessionRunnerDrainTimeout = 30 * time.Second +) + +var ( + // ErrSessionTerminated 表示 session 已由控制面终止。 + ErrSessionTerminated = errors.New("session terminated") + // ErrIdleTimeout 表示 session 在 idle 后超过 MaxIdle。 + ErrIdleTimeout = errors.New("session idle after end_turn") +) + +// ToolCallStoreDecision 是持久化 store 对某个 tool_use 的裁决。 +type ToolCallStoreDecision struct { + Sent bool + Result Event +} + +// ToolResultStore 持久化 tool_use 执行状态,避免重启后重复执行有副作用的工具。 +type ToolResultStore interface { + Recover() (map[string]Event, map[string]bool, error) + Begin(callID string, event Event) (ToolCallStoreDecision, error) + SaveResult(callID string, result Event) error + MarkSent(callID string) error +} + +// SessionToolRunnerOptions 配置单 session 的 tool event loop。 +type SessionToolRunnerOptions struct { + WorkID string + // Deprecated: MA send events 接口不再接收 lease_id。 + LeaseID string + Tools *toolset.Set + CustomTools map[string]toolset.Tool + ResultStore ToolResultStore + EventPage string + EventPollInterval time.Duration + EventLimit int + MaxIdle *time.Duration + ToolTimeout time.Duration + SendTimeout time.Duration + DrainTimeout time.Duration + Logger *log.Logger + OnToolError func(Event, error) +} + +// ToolCallResult 是一次 tool call 执行并回写后的结果。 +type ToolCallResult struct { + ToolUseID string + Name string + Custom bool + Confirmation string + Posted bool + Event Event + Result Event +} + +// SessionToolRunner 执行一个 session 的 tool event loop。 +type SessionToolRunner struct { + ctx context.Context + cancel context.CancelFunc + api API + opts SessionToolRunnerOptions + session string + events chan ToolCallResult + logger *selfhostedlog.Logger + + started bool + current ToolCallResult + err error + done chan struct{} + + inFlight sync.WaitGroup +} + +// NewSessionToolRunner 创建单 session tool runner;第一次 Next 时开始消费事件。 +func NewSessionToolRunner(ctx context.Context, api API, sessionID string, opts SessionToolRunnerOptions) *SessionToolRunner { + if ctx == nil { + ctx = context.Background() + } + runCtx, cancel := context.WithCancel(ctx) + logger := selfhostedlog.New(opts.Logger) + r := &SessionToolRunner{ + ctx: runCtx, + cancel: cancel, + api: api, + opts: opts, + session: sessionID, + logger: logger.With( + "component", "session-tool-runner", + "work_id", opts.WorkID, + "session_id", sessionID, + ), + } + return r +} + +// Next 阻塞直到下一次 tool call 结果可用或 runner 结束。 +func (r *SessionToolRunner) Next() bool { + r.start() + result, ok := <-r.events + if !ok { + return false + } + r.current = result + return true +} + +// Current 返回最近一次 Next 产出的 tool call 结果。 +func (r *SessionToolRunner) Current() ToolCallResult { + return r.current +} + +// Err 返回 runner 结束原因。 +func (r *SessionToolRunner) Err() error { + return r.err +} + +// Close 取消 runner,并等待后台循环完成。 +func (r *SessionToolRunner) Close() error { + r.cancel() + if !r.started || r.events == nil { + return nil + } + for range r.events { + } + return nil +} + +func (r *SessionToolRunner) start() { + if r.started { + return + } + r.started = true + r.events = make(chan ToolCallResult, sessionRunnerResultsBuffer) + r.done = make(chan struct{}) + go r.run() +} + +func (r *SessionToolRunner) run() { + defer close(r.done) + defer close(r.events) + defer r.drainInFlight() + r.err = normalizeRunnerErr(r.runLoop()) +} + +func (r *SessionToolRunner) runLoop() error { + if r.api == nil { + return errors.New("session tool runner api must not be nil") + } + if r.session == "" { + return errors.New("session id must not be empty") + } + if r.opts.Tools == nil { + return errors.New("session tool runner tools must not be nil") + } + state := &toolRunnerState{ + runner: r, + page: r.opts.EventPage, + processed: map[string]bool{}, + seen: map[string]bool{}, + answered: map[string]bool{}, + pendingResults: map[string]Event{}, + pendingAsk: map[string]Event{}, + confirmations: map[string]Event{}, + externalTools: map[string]Event{}, + } + if r.opts.ResultStore != nil { + pending, processed, err := r.opts.ResultStore.Recover() + if err != nil { + return fmt.Errorf("recover tool result store: %w", err) + } + state.pendingResults = pending + state.processed = processed + for id := range processed { + state.answered[id] = true + } + } + if streamer, ok := r.api.(EventStreamer); ok { + if err := state.consumeStreamLoop(r.ctx, streamer); err != nil { + switch { + case errors.Is(err, ErrEventStreamUnsupported): + case errors.Is(err, ErrIdleTimeout), + errors.Is(err, ErrSessionTerminated), + errors.Is(err, context.Canceled), + errors.Is(err, context.DeadlineExceeded): + return err + default: + return err + } + } + } + return state.consumeList(r.ctx) +} + +type toolRunnerState struct { + runner *SessionToolRunner + page string + processed map[string]bool + seen map[string]bool + answered map[string]bool + pendingResults map[string]Event + pendingAsk map[string]Event + confirmations map[string]Event + externalTools map[string]Event + idleArmedAt time.Time + idleArmPending bool +} + +type pendingToolEvent struct { + event Event + custom bool +} + +func (s *toolRunnerState) consumeStreamLoop(ctx context.Context, streamer EventStreamer) error { + backoff := sessionRunnerStreamBackoffStart + for { + if ctx.Err() != nil { + return ctx.Err() + } + stream, err := streamer.StreamEvents(ctx, StreamEventsRequest{ + SessionID: s.runner.session, + }) + if err != nil { + if errors.Is(err, ErrEventStreamUnsupported) { + return err + } + if isFatal4xxStatus(err) { + return fmt.Errorf("stream events: %w", err) + } + s.runner.logger.Warn("open event stream failed, retrying", "err", err, "backoff", backoff) + if err := s.sleepOrIdle(ctx, jitterDuration(backoff)); err != nil { + return err + } + backoff = minDuration(backoff*2, sessionRunnerStreamBackoffCap) + continue + } + if err := s.reconcile(ctx); err != nil { + _ = stream.Close() + return err + } + if ctx.Err() != nil { + _ = stream.Close() + return ctx.Err() + } + connectedAt := time.Now() + err = s.consumeStream(ctx, stream, connectedAt, &backoff) + _ = stream.Close() + if err != nil { + if errors.Is(err, ErrIdleTimeout) || + errors.Is(err, ErrSessionTerminated) || + errors.Is(err, context.Canceled) || + errors.Is(err, context.DeadlineExceeded) { + return err + } + if isFatal4xxStatus(err) { + return fmt.Errorf("stream events: %w", err) + } + s.runner.logger.Warn("event stream disconnected, retrying", "err", err, "backoff", backoff) + } + if ctx.Err() != nil { + return ctx.Err() + } + if err := s.sleepOrIdle(ctx, jitterDuration(backoff)); err != nil { + return err + } + backoff = minDuration(backoff*2, sessionRunnerStreamBackoffCap) + } +} + +func (s *toolRunnerState) consumeStream(ctx context.Context, stream *EventStream, connectedAt time.Time, backoff *time.Duration) error { + if stream == nil { + return nil + } + idle := s.maxIdle() + var timer *time.Timer + var timerC <-chan time.Time + if idle > 0 { + timer = time.NewTimer(time.Hour) + if !timer.Stop() { + <-timer.C + } + defer timer.Stop() + } + resetIdle := func() { + if timer == nil { + return + } + if !timer.Stop() { + select { + case <-timer.C: + default: + } + } + timerC = nil + if !s.idleArmedAt.IsZero() { + remaining := idle - time.Since(s.idleArmedAt) + if remaining < time.Millisecond { + remaining = time.Millisecond + } + timer.Reset(remaining) + timerC = timer.C + } + } + for { + if err := s.flushResults(ctx); err != nil { + s.runner.logger.Warn("send pending tool result failed", "err", err) + } + resetIdle() + select { + case <-ctx.Done(): + return ctx.Err() + case <-timerC: + if s.idleExpired() { + return ErrIdleTimeout + } + resetIdle() + case event, ok := <-stream.Events(): + if !ok { + return stream.Err() + } + if time.Since(connectedAt) > sessionRunnerStreamHealthyAfter { + *backoff = sessionRunnerStreamBackoffStart + } + if err := s.handleStreamEvent(ctx, event); err != nil { + return err + } + } + } +} + +func (s *toolRunnerState) reconcile(ctx context.Context) error { + backoff := sessionRunnerStreamBackoffStart + for { + err := s.reconcileOnce(ctx) + if err == nil { + return nil + } + if ctx.Err() != nil { + return ctx.Err() + } + if isFatal4xxStatus(err) { + return fmt.Errorf("reconcile list events: %w", err) + } + s.runner.logger.Warn("reconcile list events failed, retrying", "err", err, "backoff", backoff) + if err := s.sleepOrIdle(ctx, jitterDuration(backoff)); err != nil { + return err + } + backoff = minDuration(backoff*2, sessionRunnerStreamBackoffCap) + } +} + +func (s *toolRunnerState) reconcileOnce(ctx context.Context) error { + page := "" + limit := s.runner.opts.EventLimit + if limit <= 0 { + limit = 1000 + } + var events []Event + for { + resp, err := s.runner.api.ListEvents(ctx, ListEventsRequest{ + SessionID: s.runner.session, + Page: page, + Limit: limit, + Order: EventListOrderAsc, + }) + if err != nil { + return err + } + if resp != nil { + events = append(events, resp.Events...) + if resp.NextPage != "" { + page = resp.NextPage + continue + } + } + return s.processListedEvents(ctx, events, true) + } +} + +func (s *toolRunnerState) sleepOrIdle(ctx context.Context, d time.Duration) error { + deadline := time.Now().Add(d) + for { + if s.idleExpired() { + return ErrIdleTimeout + } + remainingSleep := time.Until(deadline) + if remainingSleep <= 0 { + return nil + } + wait := remainingSleep + if !s.idleArmedAt.IsZero() { + if remainingIdle := s.maxIdle() - time.Since(s.idleArmedAt); remainingIdle > 0 && remainingIdle < wait { + wait = remainingIdle + } + } + timer := time.NewTimer(wait) + select { + case <-ctx.Done(): + timer.Stop() + return ctx.Err() + case <-timer.C: + } + } +} + +func (s *toolRunnerState) consumeList(ctx context.Context) error { + interval := s.runner.opts.EventPollInterval + if interval <= 0 { + interval = 500 * time.Millisecond + } + limit := s.runner.opts.EventLimit + if limit <= 0 { + limit = 100 + } + ticker := time.NewTicker(interval) + defer ticker.Stop() + var timer *time.Timer + var timerC <-chan time.Time + if s.maxIdle() > 0 { + timer = time.NewTimer(time.Hour) + if !timer.Stop() { + <-timer.C + } + defer timer.Stop() + } + resetIdle := func() { + if timer == nil { + return + } + if !timer.Stop() { + select { + case <-timer.C: + default: + } + } + timerC = nil + if !s.idleArmedAt.IsZero() { + remaining := s.maxIdle() - time.Since(s.idleArmedAt) + if remaining < time.Millisecond { + remaining = time.Millisecond + } + timer.Reset(remaining) + timerC = timer.C + } + } + for { + if err := s.flushResults(ctx); err != nil { + s.runner.logger.Warn("send pending tool result failed", "err", err) + } + var events []Event + page := s.page + listed := false + listFailed := false + for { + resp, err := s.runner.api.ListEvents(ctx, ListEventsRequest{ + SessionID: s.runner.session, + Page: page, + Limit: limit, + Order: EventListOrderAsc, + }) + if err != nil { + s.runner.logger.Warn("list events failed", "err", err) + listFailed = true + break + } + listed = true + if resp == nil { + break + } + events = append(events, resp.Events...) + if resp.NextPage == "" { + page = "" + break + } + page = resp.NextPage + } + if len(events) > 0 { + if err := s.processListedEvents(ctx, events, false); err != nil { + return err + } + } + if listed && !listFailed && page == "" { + // MA list events 没有可从响应事件安全推进的稳定 cursor。 + // fallback 模式每轮重读完整历史,并由 seen/answered 去重, + // 避免使用 worker 本地时钟造成新事件永久漏读。 + s.page = "" + } else if listed && !listFailed { + s.page = page + } + if s.idleExpired() { + return ErrIdleTimeout + } + resetIdle() + select { + case <-ctx.Done(): + return ctx.Err() + case <-timerC: + if s.idleExpired() { + return ErrIdleTimeout + } + case <-ticker.C: + } + } +} + +func (s *toolRunnerState) processListedEvents(ctx context.Context, events []Event, reconcile bool) error { + var pending []pendingToolEvent + pendingIDs := map[string]bool{} + touchedIdle := false + lastWasEndTurn := false + for _, event := range events { + seenNow := s.markEventSeen(event) + if !reconcile && !seenNow { + continue + } + if event.Type != EventTypeUserToolConfirmation { + touchedIdle = true + lastWasEndTurn = event.Type == EventTypeSessionStatusIdle && + event.StopReasonType() == SessionStopReasonEndTurn + } + switch event.Type { + case EventTypeUserToolConfirmation: + s.recordConfirmation(event) + case EventTypeUserToolResult, EventTypeUserCustomToolResult: + s.markAnswered(toolResultCallID(event)) + case EventTypeAgentToolUse: + callID := toolUseCallID(event) + if !pendingIDs[callID] { + pending = append(pending, pendingToolEvent{event: event}) + pendingIDs[callID] = true + } + case EventTypeAgentCustomToolUse: + callID := toolUseCallID(event) + if !pendingIDs[callID] { + pending = append(pending, pendingToolEvent{event: event, custom: true}) + pendingIDs[callID] = true + } + case EventTypeSessionStatusTerminated, EventTypeSessionDeleted: + return ErrSessionTerminated + } + } + if touchedIdle { + s.disarmIdle() + } + for _, toolEvent := range pending { + if s.isAnswered(toolUseCallID(toolEvent.event)) { + continue + } + if err := s.handleToolUse(ctx, toolEvent.event, toolEvent.custom); err != nil { + return err + } + } + if err := s.releaseConfirmedToolUses(ctx); err != nil { + return err + } + if touchedIdle && lastWasEndTurn { + if s.hasUnblockedOutstandingTool(pending) { + s.disarmIdle() + } else { + s.armIdle() + } + } + return nil +} + +func (s *toolRunnerState) noteIdleEvent(event Event) { + if event.Type == EventTypeUserToolConfirmation { + return + } + if event.Type == EventTypeSessionStatusIdle && event.StopReasonType() == SessionStopReasonEndTurn { + s.armIdle() + return + } + s.disarmIdle() +} + +func (s *toolRunnerState) handleStreamEvent(ctx context.Context, event Event) error { + if !s.markEventSeen(event) { + return nil + } + s.noteIdleEvent(event) + return s.handleEvent(ctx, event) +} + +func (s *toolRunnerState) armIdle() { + if s.maxIdle() <= 0 { + return + } + if s.hasIdleBlockers() { + s.idleArmPending = true + s.idleArmedAt = time.Time{} + return + } + s.idleArmPending = false + s.idleArmedAt = time.Now() +} + +func (s *toolRunnerState) disarmIdle() { + s.idleArmPending = false + s.idleArmedAt = time.Time{} +} + +func (s *toolRunnerState) maybeArmPendingIdle() { + if !s.idleArmPending || s.hasIdleBlockers() { + return + } + s.idleArmPending = false + s.idleArmedAt = time.Now() +} + +func (s *toolRunnerState) hasIdleBlockers() bool { + return len(s.pendingAsk) > 0 || len(s.pendingResults) > 0 || len(s.externalTools) > 0 +} + +func (s *toolRunnerState) idleExpired() bool { + return s.maxIdle() > 0 && !s.idleArmedAt.IsZero() && time.Since(s.idleArmedAt) >= s.maxIdle() +} + +func (s *toolRunnerState) handleEvent(ctx context.Context, event Event) error { + switch event.Type { + case EventTypeUserToolConfirmation: + s.recordConfirmation(event) + return s.releaseConfirmedToolUses(ctx) + case EventTypeUserToolResult, EventTypeUserCustomToolResult: + s.markAnswered(toolResultCallID(event)) + case EventTypeAgentToolUse: + return s.handleToolUse(ctx, event, false) + case EventTypeAgentCustomToolUse: + return s.handleToolUse(ctx, event, true) + case EventTypeSessionStatusTerminated, EventTypeSessionDeleted: + return ErrSessionTerminated + } + return nil +} + +func (s *toolRunnerState) markEventSeen(event Event) bool { + key := event.ID + if key == "" { + key = toolUseCallID(event) + } + if key == "" { + return true + } + if s.seen[key] { + return false + } + s.seen[key] = true + return true +} + +func (s *toolRunnerState) markAnswered(callID string) { + if callID == "" { + return + } + s.answered[callID] = true + s.processed[callID] = true + delete(s.pendingResults, callID) + delete(s.pendingAsk, callID) + delete(s.externalTools, callID) + s.maybeArmPendingIdle() +} + +func (s *toolRunnerState) isAnswered(callID string) bool { + return callID != "" && s.answered[callID] +} + +func (s *toolRunnerState) recordConfirmation(event Event) { + callID := toolConfirmationCallID(event) + if callID == "" || s.isAnswered(callID) { + return + } + s.confirmations[callID] = event +} + +func (s *toolRunnerState) releaseConfirmedToolUses(ctx context.Context) error { + var ready []Event + for callID, event := range s.pendingAsk { + if _, ok := s.confirmations[callID]; ok { + ready = append(ready, event) + } + } + for _, event := range ready { + if err := s.handleToolUse(ctx, event, event.Type == EventTypeAgentCustomToolUse); err != nil { + return err + } + } + return nil +} + +func (s *toolRunnerState) hasUnblockedOutstandingTool(pending []pendingToolEvent) bool { + for _, toolEvent := range pending { + callID := toolUseCallID(toolEvent.event) + if callID == "" || s.isAnswered(callID) { + continue + } + if _, ok := s.pendingAsk[callID]; ok { + continue + } + if _, ok := s.pendingResults[callID]; ok { + continue + } + return true + } + return false +} + +func (s *toolRunnerState) handleToolUse(ctx context.Context, event Event, custom bool) error { + callID := toolUseCallID(event) + if callID == "" || s.isAnswered(callID) { + return nil + } + if pending := s.pendingResults[callID]; pending.ID != "" { + return s.sendResult(ctx, callID, event, custom, "", pending) + } + if !s.ownsTool(event, custom) { + s.runner.logger.Info("tool not owned by this runner", "tool_use_id", callID, "tool", event.Name, "custom", custom) + return s.skipExternalToolUse(ctx, event, custom, callID) + } + confirmation, allowed, err := s.permissionAllows(ctx, event, custom, callID) + if err != nil || !allowed { + return err + } + if s.runner.opts.ResultStore != nil { + decision, err := s.runner.opts.ResultStore.Begin(callID, event) + if err != nil { + return fmt.Errorf("begin tool result store for %s: %w", callID, err) + } + if decision.Sent { + s.markAnswered(callID) + return nil + } + if decision.Result.ID != "" { + s.pendingResults[callID] = decision.Result + return s.sendResult(ctx, callID, event, custom, "", decision.Result) + } + } + var result toolset.Result + input := json.RawMessage(event.Input) + if custom { + tool := s.runner.opts.CustomTools[event.Name] + result = s.executeWithTimeout(ctx, event, func(toolCtx context.Context) toolset.Result { + return tool.Execute(toolCtx, input) + }) + } else { + result = s.executeWithTimeout(ctx, event, func(toolCtx context.Context) toolset.Result { + return s.runner.opts.Tools.Execute(toolCtx, event.Name, input) + }) + } + return s.postResult(ctx, event, custom, callID, result, confirmation) +} + +func (s *toolRunnerState) ownsTool(event Event, custom bool) bool { + if custom { + return s.runner.opts.CustomTools[event.Name] != nil + } + return s.runner.opts.Tools.Has(event.Name) +} + +func (s *toolRunnerState) permissionAllows(ctx context.Context, event Event, custom bool, callID string) (string, bool, error) { + if custom { + return "", true, nil + } + switch event.EvaluatedPermission { + case "": + return "", true, nil + case PermissionAllow: + return "", true, nil + case PermissionAsk: + confirmation, ok := s.confirmations[callID] + if !ok { + s.pendingAsk[callID] = event + return "", false, nil + } + switch confirmation.Result { + case ConfirmationAllow: + return "allow", true, nil + case ConfirmationDeny: + return ConfirmationDeny, false, s.resolveToolUseWithoutPost(ctx, event, false, callID, ConfirmationDeny) + default: + return ConfirmationDeny, false, s.resolveToolUseWithoutPost(ctx, event, false, callID, ConfirmationDeny) + } + case PermissionDeny: + return ConfirmationDeny, false, s.resolveToolUseWithoutPost(ctx, event, false, callID, ConfirmationDeny) + default: + confirmation, ok := s.confirmations[callID] + if !ok { + s.pendingAsk[callID] = event + return "", false, nil + } + if confirmation.Result == ConfirmationAllow { + return "allow", true, nil + } + return ConfirmationDeny, false, s.resolveToolUseWithoutPost(ctx, event, false, callID, ConfirmationDeny) + } +} + +func (s *toolRunnerState) executeWithTimeout(ctx context.Context, event Event, fn func(context.Context) toolset.Result) toolset.Result { + timeout := s.runner.opts.ToolTimeout + if timeout <= 0 { + timeout = DefaultToolTimeout + } + toolCtx, cancel := context.WithTimeout(ctx, timeout) + defer cancel() + done := make(chan toolset.Result, 1) + s.runner.inFlight.Add(1) + go func() { + defer s.runner.inFlight.Done() + defer func() { + if recovered := recover(); recovered != nil { + s.runner.logger.Error( + "tool execution panicked", + "tool_use_id", toolUseCallID(event), + "tool", event.Name, + "panic", recovered, + "stack", string(debug.Stack()), + ) + done <- toolset.ErrorResult(fmt.Sprintf("tool execution panicked: %v", recovered)) + } + }() + done <- fn(toolCtx) + }() + select { + case result := <-done: + return result + case <-toolCtx.Done(): + if errors.Is(toolCtx.Err(), context.DeadlineExceeded) { + return toolset.ErrorResult(fmt.Sprintf("tool execution timed out after %s", timeout)) + } + return toolset.ErrorResult(toolCtx.Err().Error()) + } +} + +func (s *toolRunnerState) postResult(ctx context.Context, event Event, custom bool, callID string, result toolset.Result, confirmation string) error { + var out Event + if custom { + out = NewUserCustomToolResultEvent(callID, runnerContentBlocks(result.Content), result.IsError, event.SessionThreadID) + } else { + out = NewUserToolResultEvent(callID, runnerContentBlocks(result.Content), result.IsError, event.SessionThreadID) + } + if s.runner.opts.ResultStore != nil { + if err := s.runner.opts.ResultStore.SaveResult(callID, out); err != nil { + s.runner.logger.Warn("persist tool result failed", "tool_use_id", callID, "err", err) + } + s.pendingResults[callID] = out + } + return s.sendResult(ctx, callID, event, custom, confirmation, out) +} + +func (s *toolRunnerState) sendResult(ctx context.Context, callID string, event Event, custom bool, confirmation string, out Event) error { + req := SendEventRequest{ + SessionID: s.runner.session, + Event: out, + } + posted, err := s.retrySendEvent(ctx, req, callID) + if err != nil && s.runner.opts.OnToolError != nil { + s.runner.opts.OnToolError(event, err) + } + if posted { + s.markAnswered(callID) + if s.runner.opts.ResultStore != nil { + if err := s.runner.opts.ResultStore.MarkSent(callID); err != nil { + s.runner.logger.Warn("mark tool result sent failed", "tool_use_id", callID, "event_id", out.ID, "err", err) + } + } + } else if s.runner.opts.ResultStore != nil { + s.pendingResults[callID] = out + } + select { + case <-ctx.Done(): + return ctx.Err() + case s.runner.events <- ToolCallResult{ + ToolUseID: callID, + Name: event.Name, + Custom: custom, + Confirmation: confirmation, + Posted: posted, + Event: event, + Result: out, + }: + return nil + } +} + +func (s *toolRunnerState) resolveToolUseWithoutPost(ctx context.Context, event Event, custom bool, callID string, confirmation string) error { + s.markAnswered(callID) + s.maybeArmPendingIdle() + select { + case <-ctx.Done(): + return ctx.Err() + case s.runner.events <- ToolCallResult{ + ToolUseID: callID, + Name: event.Name, + Custom: custom, + Confirmation: confirmation, + Posted: false, + Event: event, + }: + return nil + } +} + +func (s *toolRunnerState) skipExternalToolUse(ctx context.Context, event Event, custom bool, callID string) error { + s.externalTools[callID] = event + s.maybeArmPendingIdle() + select { + case <-ctx.Done(): + return ctx.Err() + case s.runner.events <- ToolCallResult{ + ToolUseID: callID, + Name: event.Name, + Custom: custom, + Posted: false, + Event: event, + }: + return nil + } +} + +func (s *toolRunnerState) flushResults(ctx context.Context) error { + var first error + for callID, event := range s.pendingResults { + req := SendEventRequest{ + SessionID: s.runner.session, + Event: event, + } + posted, err := s.retrySendEvent(ctx, req, callID) + if !posted { + if first == nil { + first = err + } + continue + } + s.markAnswered(callID) + if s.runner.opts.ResultStore != nil { + if err := s.runner.opts.ResultStore.MarkSent(callID); err != nil { + s.runner.logger.Warn("mark pending tool result sent failed", "tool_use_id", callID, "event_id", event.ID, "err", err) + } + } + } + s.maybeArmPendingIdle() + return first +} + +func (s *toolRunnerState) retrySendEvent(ctx context.Context, req SendEventRequest, callID string) (bool, error) { + var err error + for attempt := 0; attempt < sessionRunnerSendRetries; attempt++ { + // 工具已经执行后,即使 runner 开始退出,也为结果保留一次有界的最终投递机会。 + sendCtx, cancel := context.WithTimeout(context.Background(), s.sendTimeout()) + err = s.runner.api.SendEvent(sendCtx, req) + cancel() + if err == nil { + return true, nil + } + if isFatal4xxStatus(err) { + s.runner.logger.Error("tool result send hit permanent 4xx", "tool_use_id", callID, "err", err) + return false, err + } + if ctx.Err() != nil { + s.runner.logger.Debug("tool result send abandoned after context cancellation", "tool_use_id", callID, "err", err) + return false, err + } + s.runner.logger.Warn("tool result send failed, retrying", "tool_use_id", callID, "attempt", attempt+1, "err", err) + if attempt < sessionRunnerSendRetries-1 { + sleepCtx(ctx, time.Duration(attempt+1)*time.Second) + } + } + return false, err +} + +func (s *toolRunnerState) maxIdle() time.Duration { + if s.runner.opts.MaxIdle == nil { + return DefaultMaxIdle + } + return *s.runner.opts.MaxIdle +} + +func (s *toolRunnerState) sendTimeout() time.Duration { + if s.runner.opts.SendTimeout > 0 { + return s.runner.opts.SendTimeout + } + return defaultSessionRunnerSendTimeout +} + +func (r *SessionToolRunner) drainInFlight() { + done := make(chan struct{}) + go func() { + r.inFlight.Wait() + close(done) + }() + timer := time.NewTimer(r.drainTimeout()) + defer timer.Stop() + select { + case <-done: + case <-timer.C: + r.logger.Warn("drain timeout exceeded; in-flight tools may still be running", "drain_timeout", r.drainTimeout()) + } +} + +func (r *SessionToolRunner) drainTimeout() time.Duration { + if r.opts.DrainTimeout > 0 { + return r.opts.DrainTimeout + } + return defaultSessionRunnerDrainTimeout +} + +func normalizeRunnerErr(err error) error { + if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) { + return nil + } + return err +} + +func sleepCtx(ctx context.Context, d time.Duration) { + timer := time.NewTimer(d) + defer timer.Stop() + select { + case <-ctx.Done(): + case <-timer.C: + } +} + +func jitterDuration(d time.Duration) time.Duration { + if d <= 1 { + return d + } + half := d / 2 + return half + time.Duration(rand.Int63n(int64(d-half))) +} + +func minDuration(a, b time.Duration) time.Duration { + if a < b { + return a + } + return b +} + +func isFatal4xxStatus(err error) bool { + var apiErr *APIError + if !errors.As(err, &apiErr) { + return false + } + return apiErr.StatusCode >= 400 && + apiErr.StatusCode < 500 && + apiErr.StatusCode != http.StatusRequestTimeout && + apiErr.StatusCode != http.StatusConflict && + apiErr.StatusCode != http.StatusTooManyRequests +} + +func runnerContentBlocks(blocks []toolset.ContentBlock) []ContentBlock { + out := make([]ContentBlock, 0, len(blocks)) + for _, block := range blocks { + out = append(out, ContentBlock{ + Type: block.Type, + Text: block.Text, + MediaType: block.MediaType, + Data: block.Data, + }) + } + return out +} + +func toolConfirmationCallID(event Event) string { + if event.ToolUseID != "" { + return event.ToolUseID + } + return event.CustomToolUseID +} + +func toolUseCallID(event Event) string { + if event.ToolUseID != "" { + return event.ToolUseID + } + if event.CustomToolUseID != "" { + return event.CustomToolUseID + } + return event.ID +} + +func toolResultCallID(event Event) string { + if event.ToolUseID != "" { + return event.ToolUseID + } + return event.CustomToolUseID +} diff --git a/arkruntime/selfhosted/session_tool_runner_test.go b/arkruntime/selfhosted/session_tool_runner_test.go new file mode 100644 index 0000000..84d6d7c --- /dev/null +++ b/arkruntime/selfhosted/session_tool_runner_test.go @@ -0,0 +1,259 @@ +// Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +// SPDX-License-Identifier: Apache-2.0 +package selfhosted + +import ( + "context" + "encoding/json" + "errors" + "io" + "log" + "net/http" + "testing" + "time" + + "github.com/volcengine/ark-runtime-go/arkruntime/toolset" +) + +type runnerTestAPI struct { + listCalls int + listEvent Event + sent []Event +} + +func (a *runnerTestAPI) PollWork(context.Context, PollWorkRequest) (*WorkItem, error) { + return nil, nil +} + +func (a *runnerTestAPI) AckWork(context.Context, AckWorkRequest) error { return nil } + +func (a *runnerTestAPI) HeartbeatWork(context.Context, HeartbeatWorkRequest) (*HeartbeatResponse, error) { + return nil, nil +} + +func (a *runnerTestAPI) StopWork(context.Context, StopWorkRequest) error { return nil } + +func (a *runnerTestAPI) GetSession(context.Context, GetSessionRequest) (*Session, error) { + return nil, nil +} + +func (a *runnerTestAPI) ListEvents(context.Context, ListEventsRequest) (*ListEventsResponse, error) { + a.listCalls++ + if a.listCalls == 1 { + return nil, errors.New("temporary list failure") + } + return &ListEventsResponse{Events: []Event{a.listEvent}}, nil +} + +func (a *runnerTestAPI) SendEvent(_ context.Context, req SendEventRequest) error { + a.sent = append(a.sent, req.Event) + return nil +} + +func (a *runnerTestAPI) OpenSkill(context.Context, OpenSkillRequest) (*SkillContent, error) { + return nil, nil +} + +type runnerTestTool struct { + calls int +} + +func (t *runnerTestTool) Name() string { return "custom" } + +func (t *runnerTestTool) Execute(context.Context, json.RawMessage) toolset.Result { + t.calls++ + return toolset.TextResult("ok") +} + +type runnerFailingMarkSentStore struct { + markCalls int +} + +func (s *runnerFailingMarkSentStore) Recover() (map[string]Event, map[string]bool, error) { + return nil, nil, nil +} + +func (s *runnerFailingMarkSentStore) Begin(string, Event) (ToolCallStoreDecision, error) { + return ToolCallStoreDecision{}, nil +} + +func (s *runnerFailingMarkSentStore) SaveResult(string, Event) error { return nil } + +func (s *runnerFailingMarkSentStore) MarkSent(string) error { + s.markCalls++ + return errors.New("mark sent failed") +} + +func TestSessionToolRunnerReconcileRetriesWithoutLosingToolUse(t *testing.T) { + tool := &runnerTestTool{} + api := &runnerTestAPI{listEvent: Event{ + ID: "event-id", + Type: EventTypeAgentCustomToolUse, + Name: tool.Name(), + ToolUseID: "call-id", + SessionThreadID: "thread-id", + Input: RawJSON(`{}`), + }} + runner := NewSessionToolRunner(context.Background(), api, "session-id", SessionToolRunnerOptions{ + CustomTools: map[string]toolset.Tool{tool.Name(): tool}, + Logger: log.New(io.Discard, "", 0), + ToolTimeout: time.Second, + }) + runner.events = make(chan ToolCallResult, 1) + state := &toolRunnerState{ + runner: runner, + processed: map[string]bool{}, + seen: map[string]bool{}, + answered: map[string]bool{}, + pendingResults: map[string]Event{}, + pendingAsk: map[string]Event{}, + confirmations: map[string]Event{}, + externalTools: map[string]Event{}, + } + + if err := state.reconcile(context.Background()); err != nil { + t.Fatal(err) + } + if api.listCalls != 2 || tool.calls != 1 || len(api.sent) != 1 { + t.Fatalf("list_calls=%d tool_calls=%d sent=%d", api.listCalls, tool.calls, len(api.sent)) + } + if api.sent[0].CustomToolUseID != "call-id" { + t.Fatalf("sent event=%+v", api.sent[0]) + } +} + +func TestSessionToolRunnerConvertsToolPanicToErrorResult(t *testing.T) { + runner := NewSessionToolRunner(context.Background(), &runnerTestAPI{}, "session-id", SessionToolRunnerOptions{ + Logger: log.New(io.Discard, "", 0), + ToolTimeout: time.Second, + }) + state := &toolRunnerState{runner: runner} + result := state.executeWithTimeout(context.Background(), Event{ + ID: "event-id", + ToolUseID: "call-id", + Name: "panicking-tool", + }, func(context.Context) toolset.Result { + panic("boom") + }) + if !result.IsError || len(result.Content) != 1 || result.Content[0].Text != "tool execution panicked: boom" { + t.Fatalf("result=%+v", result) + } +} + +func TestSessionToolRunnerReplayDoesNotResetIdleDeadline(t *testing.T) { + maxIdle := time.Second + runner := NewSessionToolRunner(context.Background(), &runnerTestAPI{}, "session-id", SessionToolRunnerOptions{ + MaxIdle: &maxIdle, + Logger: log.New(io.Discard, "", 0), + }) + state := &toolRunnerState{ + runner: runner, + seen: map[string]bool{"replayed-event": true}, + } + state.armIdle() + armedAt := state.idleArmedAt + if err := state.handleStreamEvent(context.Background(), Event{ID: "replayed-event", Type: "user.message"}); err != nil { + t.Fatal(err) + } + if state.idleArmedAt != armedAt { + t.Fatalf("idle deadline changed from %s to %s", armedAt, state.idleArmedAt) + } +} + +func TestSessionToolRunnerMarkSentFailureDoesNotRetainDeliveredResult(t *testing.T) { + api := &runnerTestAPI{} + store := &runnerFailingMarkSentStore{} + runner := NewSessionToolRunner(context.Background(), api, "session-id", SessionToolRunnerOptions{ + ResultStore: store, + Logger: log.New(io.Discard, "", 0), + }) + runner.events = make(chan ToolCallResult, 1) + out := NewUserToolResultEvent("call-id", []ContentBlock{{Type: "text", Text: "ok"}}, false, "thread-id") + state := &toolRunnerState{ + runner: runner, + processed: map[string]bool{}, + answered: map[string]bool{}, + pendingResults: map[string]Event{"call-id": out}, + } + if err := state.sendResult(context.Background(), "call-id", Event{Name: "bash"}, false, "", out); err != nil { + t.Fatal(err) + } + if len(api.sent) != 1 || store.markCalls != 1 { + t.Fatalf("sent=%d mark_calls=%d", len(api.sent), store.markCalls) + } + if !state.isAnswered("call-id") || len(state.pendingResults) != 0 { + t.Fatalf("answered=%v pending=%v", state.answered, state.pendingResults) + } +} + +func TestSessionToolRunnerFlushMarkSentFailureDoesNotRetainDeliveredResult(t *testing.T) { + api := &runnerTestAPI{} + store := &runnerFailingMarkSentStore{} + runner := NewSessionToolRunner(context.Background(), api, "session-id", SessionToolRunnerOptions{ + ResultStore: store, + Logger: log.New(io.Discard, "", 0), + }) + out := NewUserToolResultEvent("call-id", []ContentBlock{{Type: "text", Text: "ok"}}, false, "thread-id") + state := &toolRunnerState{ + runner: runner, + processed: map[string]bool{}, + answered: map[string]bool{}, + pendingResults: map[string]Event{"call-id": out}, + } + if err := state.flushResults(context.Background()); err != nil { + t.Fatal(err) + } + if len(api.sent) != 1 || store.markCalls != 1 { + t.Fatalf("sent=%d mark_calls=%d", len(api.sent), store.markCalls) + } + if !state.isAnswered("call-id") || len(state.pendingResults) != 0 { + t.Fatalf("answered=%v pending=%v", state.answered, state.pendingResults) + } +} + +func TestSessionToolRunnerPermissionDenyDoesNotPostToolResult(t *testing.T) { + api := &runnerTestAPI{} + tools, err := toolset.NewDefault(toolset.Options{Workdir: t.TempDir()}) + if err != nil { + t.Fatal(err) + } + defer tools.Close() + runner := NewSessionToolRunner(context.Background(), api, "session-id", SessionToolRunnerOptions{ + Tools: tools, + Logger: log.New(io.Discard, "", 0), + }) + runner.events = make(chan ToolCallResult, 1) + state := &toolRunnerState{ + runner: runner, + processed: map[string]bool{}, + answered: map[string]bool{}, + pendingResults: map[string]Event{}, + pendingAsk: map[string]Event{}, + confirmations: map[string]Event{}, + externalTools: map[string]Event{}, + } + event := Event{ + ID: "event-id", + Type: EventTypeAgentToolUse, + Name: "read", + ToolUseID: "call-id", + EvaluatedPermission: PermissionDeny, + } + if err := state.handleToolUse(context.Background(), event, false); err != nil { + t.Fatal(err) + } + if len(api.sent) != 0 || !state.isAnswered("call-id") { + t.Fatalf("sent=%d answered=%v", len(api.sent), state.answered) + } + result := <-runner.events + if result.Posted || result.Confirmation != ConfirmationDeny { + t.Fatalf("result=%+v", result) + } +} + +func TestSessionToolRunnerConflictIsRetryable(t *testing.T) { + err := &APIError{StatusCode: http.StatusConflict, Message: "temporary conflict"} + if isFatal4xxStatus(err) { + t.Fatal("409 conflict must remain retryable") + } +} diff --git a/arkruntime/selfhosted/skill_hub.go b/arkruntime/selfhosted/skill_hub.go new file mode 100644 index 0000000..30d2108 --- /dev/null +++ b/arkruntime/selfhosted/skill_hub.go @@ -0,0 +1,130 @@ +// Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +// SPDX-License-Identifier: Apache-2.0 + +package selfhosted + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "io" + "net/http" + "net/url" + "strings" +) + +const ( + skillTypeSkillHub = "skill_hub" + defaultSkillHubBaseURL = "https://skills.volces.com/v1/skills" + maxSkillHubMetadataResponse = 1 << 20 +) + +type skillHubListResponse struct { + Skills []skillHubSkill `json:"Skills"` +} + +type skillHubSkill struct { + ID string `json:"Id"` + Slug string `json:"Slug"` +} + +func (a *ClientAPI) openSkillHub(ctx context.Context, skill SkillRef) (*SkillContent, error) { + skillID := strings.TrimSpace(skill.IDValue()) + if skillID == "" { + return nil, errors.New("skill id is required") + } + version := strings.TrimSpace(skill.Version) + if version == "" { + return nil, errors.New("skill version is required") + } + slug, err := lookupSkillHubSlug(ctx, a.client.HTTPClient(), skillID) + if err != nil { + return nil, err + } + downloadURL, err := skillHubDownloadURL(slug, version) + if err != nil { + return nil, err + } + return openSignedSkillURL(ctx, a.client.HTTPClient(), downloadURL) +} + +func lookupSkillHubSlug(ctx context.Context, client *http.Client, skillID string) (string, error) { + metadataURL, err := url.Parse(defaultSkillHubBaseURL) + if err != nil { + return "", fmt.Errorf("parse skill hub base url: %w", err) + } + query := metadataURL.Query() + query.Set("skillIds", skillID) + metadataURL.RawQuery = query.Encode() + req, err := http.NewRequestWithContext(ctx, http.MethodGet, metadataURL.String(), nil) + if err != nil { + return "", fmt.Errorf("build skill hub metadata request: %w", err) + } + if client == nil { + client = http.DefaultClient + } + resp, err := client.Do(req) + if err != nil { + return "", fmt.Errorf("lookup skill hub metadata: %w", err) + } + defer resp.Body.Close() //nolint:errcheck // response body close errors are non-actionable + if resp.StatusCode < http.StatusOK || resp.StatusCode >= http.StatusMultipleChoices { + return "", skillHubHTTPError(resp) + } + body, err := io.ReadAll(io.LimitReader(resp.Body, maxSkillHubMetadataResponse+1)) + if err != nil { + return "", fmt.Errorf("read skill hub metadata: %w", err) + } + if len(body) > maxSkillHubMetadataResponse { + return "", errors.New("skill hub metadata response is too large") + } + var metadata skillHubListResponse + if err := json.Unmarshal(body, &metadata); err != nil { + return "", fmt.Errorf("decode skill hub metadata: %w", err) + } + for _, candidate := range metadata.Skills { + if strings.TrimSpace(candidate.ID) != skillID { + continue + } + slug := strings.Trim(strings.TrimSpace(candidate.Slug), "/") + if slug == "" { + return "", fmt.Errorf("skill hub slug is empty: %s", skillID) + } + return slug, nil + } + return "", &APIError{StatusCode: http.StatusNotFound, Message: "skill hub skill not found: " + skillID} +} + +func skillHubDownloadURL(slug, version string) (string, error) { + baseURL, err := url.Parse(defaultSkillHubBaseURL) + if err != nil { + return "", fmt.Errorf("parse skill hub base url: %w", err) + } + segments := strings.Split(strings.Trim(slug, "/"), "/") + cleanSegments := make([]string, 0, len(segments)) + for _, segment := range segments { + segment = strings.TrimSpace(segment) + if segment == "" || segment == "." || segment == ".." { + return "", fmt.Errorf("invalid skill hub slug: %q", slug) + } + cleanSegments = append(cleanSegments, segment) + } + baseURL.Path = strings.TrimRight(baseURL.Path, "/") + "/download/" + strings.Join(cleanSegments, "/") + query := baseURL.Query() + query.Set("version", version) + baseURL.RawQuery = query.Encode() + return baseURL.String(), nil +} + +func skillHubHTTPError(resp *http.Response) error { + message, _ := io.ReadAll(io.LimitReader(resp.Body, 4096)) + if len(message) == 0 { + message = []byte(http.StatusText(resp.StatusCode)) + } + return &APIError{ + StatusCode: resp.StatusCode, + Message: strings.TrimSpace(string(message)), + RequestID: firstNonEmpty(resp.Header.Get("X-Skill-Request-Id"), resp.Header.Get("X-Request-Id")), + } +} diff --git a/arkruntime/selfhosted/tool_result_store.go b/arkruntime/selfhosted/tool_result_store.go new file mode 100644 index 0000000..df2b1d1 --- /dev/null +++ b/arkruntime/selfhosted/tool_result_store.go @@ -0,0 +1,223 @@ +// Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +// SPDX-License-Identifier: Apache-2.0 +package selfhosted + +import ( + "crypto/sha256" + "encoding/hex" + "encoding/json" + "errors" + "fmt" + "os" + "path/filepath" + "strings" + "time" +) + +const ( + fileToolResultStateStarted = "started" + fileToolResultStateResult = "result" + fileToolResultStateSent = "sent" +) + +// FileToolResultStore 用本地文件持久化 tool_use 执行状态。 +type FileToolResultStore struct { + dir string +} + +type fileToolResultRecord struct { + CallID string `json:"call_id"` + State string `json:"state"` + Event Event `json:"event"` + Result Event `json:"result,omitempty"` + UpdatedAt string `json:"updated_at"` +} + +// NewFileToolResultStore 在 workdir 下创建 tool result 持久化 store。 +func NewFileToolResultStore(workdir string) (*FileToolResultStore, error) { + if workdir == "" { + return nil, errors.New("workdir must not be empty") + } + dir := filepath.Join(workdir, ".ma_self_host_worker", "tool_ledger") + if err := os.MkdirAll(dir, 0o700); err != nil { + return nil, fmt.Errorf("create tool result store: %w", err) + } + return &FileToolResultStore{dir: dir}, nil +} + +// Recover 恢复未回写的 tool result 和已经完成回写的 tool_use。 +func (s *FileToolResultStore) Recover() (map[string]Event, map[string]bool, error) { + pending := map[string]Event{} + processed := map[string]bool{} + entries, err := os.ReadDir(s.dir) + if err != nil { + return nil, nil, err + } + for _, entry := range entries { + if entry.IsDir() { + continue + } + if strings.HasPrefix(entry.Name(), ".tool-result-") { + _ = os.Remove(filepath.Join(s.dir, entry.Name())) + continue + } + if filepath.Ext(entry.Name()) != ".json" { + continue + } + record, err := s.readPath(filepath.Join(s.dir, entry.Name())) + if err != nil { + return nil, nil, err + } + switch record.State { + case fileToolResultStateSent: + processed[record.CallID] = true + case fileToolResultStateResult: + pending[record.CallID] = record.Result + case fileToolResultStateStarted: + record.Result = unknownToolExecutionResult(record.CallID, record.Event) + record.State = fileToolResultStateResult + if err := s.write(record); err != nil { + return nil, nil, err + } + pending[record.CallID] = record.Result + default: + return nil, nil, fmt.Errorf("unknown tool result state %q for call %s", record.State, record.CallID) + } + } + return pending, processed, nil +} + +// Begin 记录一个 tool_use 即将开始执行,并返回是否已有可复用结果。 +func (s *FileToolResultStore) Begin(callID string, event Event) (ToolCallStoreDecision, error) { + if callID == "" { + return ToolCallStoreDecision{}, errors.New("call id must not be empty") + } + record, err := s.read(callID) + if err == nil { + switch record.State { + case fileToolResultStateSent: + return ToolCallStoreDecision{Sent: true}, nil + case fileToolResultStateResult: + return ToolCallStoreDecision{Result: record.Result}, nil + case fileToolResultStateStarted: + record.Result = unknownToolExecutionResult(record.CallID, record.Event) + record.State = fileToolResultStateResult + if err := s.write(record); err != nil { + return ToolCallStoreDecision{}, err + } + return ToolCallStoreDecision{Result: record.Result}, nil + default: + return ToolCallStoreDecision{}, fmt.Errorf("unknown tool result state %q for call %s", record.State, callID) + } + } + if !errors.Is(err, os.ErrNotExist) { + return ToolCallStoreDecision{}, err + } + return ToolCallStoreDecision{}, s.write(fileToolResultRecord{ + CallID: callID, + State: fileToolResultStateStarted, + Event: event, + UpdatedAt: time.Now().UTC().Format(time.RFC3339Nano), + }) +} + +// SaveResult 持久化 tool_use 结果。 +func (s *FileToolResultStore) SaveResult(callID string, result Event) error { + record, err := s.read(callID) + if err != nil { + return err + } + record.State = fileToolResultStateResult + record.Result = result + record.UpdatedAt = time.Now().UTC().Format(time.RFC3339Nano) + return s.write(record) +} + +// MarkSent 标记 tool result 已经成功回写到控制面。 +func (s *FileToolResultStore) MarkSent(callID string) error { + record, err := s.read(callID) + if err != nil { + return err + } + record.State = fileToolResultStateSent + record.UpdatedAt = time.Now().UTC().Format(time.RFC3339Nano) + return s.write(record) +} + +func (s *FileToolResultStore) read(callID string) (fileToolResultRecord, error) { + return s.readPath(s.path(callID)) +} + +func (s *FileToolResultStore) readPath(path string) (fileToolResultRecord, error) { + var record fileToolResultRecord + data, err := os.ReadFile(path) + if err != nil { + return record, err + } + if err := json.Unmarshal(data, &record); err != nil { + return record, fmt.Errorf("decode tool result store %s: %w", path, err) + } + if record.CallID == "" { + return record, fmt.Errorf("tool result store %s missing call_id", path) + } + return record, nil +} + +func (s *FileToolResultStore) write(record fileToolResultRecord) error { + if record.CallID == "" { + return errors.New("call id must not be empty") + } + record.UpdatedAt = time.Now().UTC().Format(time.RFC3339Nano) + data, err := json.MarshalIndent(record, "", " ") + if err != nil { + return err + } + target := s.path(record.CallID) + tmp, err := os.CreateTemp(s.dir, ".tool-result-*") + if err != nil { + return err + } + tmpName := tmp.Name() + if _, err := tmp.Write(data); err != nil { + _ = tmp.Close() + _ = os.Remove(tmpName) + return err + } + if err := tmp.Sync(); err != nil { + _ = tmp.Close() + _ = os.Remove(tmpName) + return err + } + if err := tmp.Close(); err != nil { + _ = os.Remove(tmpName) + return err + } + if err := os.Rename(tmpName, target); err != nil { + _ = os.Remove(tmpName) + return err + } + dir, err := os.Open(s.dir) + if err != nil { + return err + } + defer func() { _ = dir.Close() }() + return dir.Sync() +} + +func (s *FileToolResultStore) path(callID string) string { + sum := sha256.Sum256([]byte(callID)) + return filepath.Join(s.dir, hex.EncodeToString(sum[:])+".json") +} + +func unknownToolExecutionResult(callID string, event Event) Event { + content := []ContentBlock{{ + Type: "text", + Text: "tool execution state is unknown after worker restart; refusing to re-execute this tool_use to avoid duplicate side effects", + }} + if event.Type == EventTypeAgentCustomToolUse { + return NewUserCustomToolResultEvent(callID, content, true, event.SessionThreadID) + } + return NewUserToolResultEvent(callID, content, true, event.SessionThreadID) +} + +var _ ToolResultStore = (*FileToolResultStore)(nil) diff --git a/arkruntime/selfhosted/tool_result_store_test.go b/arkruntime/selfhosted/tool_result_store_test.go new file mode 100644 index 0000000..10b5e2a --- /dev/null +++ b/arkruntime/selfhosted/tool_result_store_test.go @@ -0,0 +1,59 @@ +// Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +// SPDX-License-Identifier: Apache-2.0 +package selfhosted + +import ( + "os" + "path/filepath" + "testing" +) + +func TestFileToolResultStoreRecoverIgnoresInterruptedTempFile(t *testing.T) { + store, err := NewFileToolResultStore(t.TempDir()) + if err != nil { + t.Fatal(err) + } + tempPath := filepath.Join(store.dir, ".tool-result-interrupted") + if err := os.WriteFile(tempPath, []byte("{"), 0o600); err != nil { + t.Fatal(err) + } + + pending, processed, err := store.Recover() + if err != nil { + t.Fatal(err) + } + if len(pending) != 0 || len(processed) != 0 { + t.Fatalf("pending=%v processed=%v", pending, processed) + } + if _, err := os.Stat(tempPath); !os.IsNotExist(err) { + t.Fatalf("temporary record was not removed: %v", err) + } +} + +func TestFileToolResultStoreRecoverUsesPersistedCallID(t *testing.T) { + store, err := NewFileToolResultStore(t.TempDir()) + if err != nil { + t.Fatal(err) + } + event := Event{ + ID: "event-id", + Type: EventTypeAgentToolUse, + ToolUseID: "call-id", + SessionThreadID: "thread-id", + } + if _, err := store.Begin("call-id", event); err != nil { + t.Fatal(err) + } + + pending, _, err := store.Recover() + if err != nil { + t.Fatal(err) + } + result, ok := pending["call-id"] + if !ok { + t.Fatalf("pending=%v", pending) + } + if result.ToolUseID != "call-id" { + t.Fatalf("tool_use_id=%q", result.ToolUseID) + } +} diff --git a/arkruntime/selfhosted/types.go b/arkruntime/selfhosted/types.go new file mode 100644 index 0000000..89692b2 --- /dev/null +++ b/arkruntime/selfhosted/types.go @@ -0,0 +1,614 @@ +// Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +// SPDX-License-Identifier: Apache-2.0 +package selfhosted + +import ( + "bytes" + "context" + "crypto/rand" + "encoding/hex" + "encoding/json" + "errors" + "fmt" + "io" + "time" +) + +const ( + // EventTypeAgentToolUse 表示 agent 请求执行内置工具。 + EventTypeAgentToolUse = "agent.tool_use" + // EventTypeAgentCustomToolUse 表示 agent 请求执行客户自定义工具。 + EventTypeAgentCustomToolUse = "agent.custom_tool_use" + // EventTypeUserToolConfirmation 表示用户确认工具执行。 + EventTypeUserToolConfirmation = "user.tool_confirmation" + // EventTypeUserToolResult 表示 self-host worker 回写内置工具结果。 + EventTypeUserToolResult = "user.tool_result" + // EventTypeUserCustomToolResult 表示 self-host worker 回写自定义工具结果。 + EventTypeUserCustomToolResult = "user.custom_tool_result" + // EventTypeSessionStatusIdle 表示 session 已经进入空闲状态。 + EventTypeSessionStatusIdle = "session.status_idle" + // EventTypeSessionStatusTerminated 表示 session 已经终止。 + EventTypeSessionStatusTerminated = "session.status_terminated" + // EventTypeSessionDeleted 表示 session 已经删除。 + EventTypeSessionDeleted = "session.deleted" +) + +const ( + // PermissionAllow 表示工具调用已经被授权执行。 + PermissionAllow = "allow" + // PermissionAsk 表示工具调用需要等待确认。 + PermissionAsk = "ask" + // PermissionDeny 表示工具调用被拒绝。 + PermissionDeny = "deny" +) + +const ( + // ConfirmationAllow 表示用户批准执行。 + ConfirmationAllow = "allow" + // ConfirmationDeny 表示用户拒绝执行。 + ConfirmationDeny = "deny" +) + +const ( + // EventListOrderAsc 表示按事件创建时间正序读取。 + EventListOrderAsc = "asc" + // EventListOrderDesc 表示按事件创建时间倒序读取。 + EventListOrderDesc = "desc" +) + +const ( + // SessionStopReasonEndTurn 表示 agent 自然完成本轮。 + SessionStopReasonEndTurn = "end_turn" + // SessionStopReasonRequiresAction 表示 agent 正等待用户输入或外部结果。 + SessionStopReasonRequiresAction = "requires_action" + // SessionStopReasonRetriesExhausted 表示重试预算用尽导致本轮终止。 + SessionStopReasonRetriesExhausted = "retries_exhausted" +) + +const ( + // WorkStateQueued 表示 work 仍在队列中。 + WorkStateQueued = "queued" + // WorkStateStarting 表示 work 已 ack 并正在启动执行环境。 + WorkStateStarting = "starting" + // WorkStateActive 表示 work lease 仍归当前 worker 所有。 + WorkStateActive = "active" + // WorkStateRunning 兼容 MA 早期返回的 running 状态。 + // + // Deprecated: MA work state 不再包含 running,使用 active。 + WorkStateRunning = "running" + // WorkStateStopping 表示控制面要求 worker 停止当前 work。 + WorkStateStopping = "stopping" + // WorkStateStopped 表示当前 work 已经停止。 + WorkStateStopped = "stopped" +) + +const ( + // ExpectedLastHeartbeatNoHeartbeat 表示 work 尚未产生成功 heartbeat。 + ExpectedLastHeartbeatNoHeartbeat = "NO_HEARTBEAT" +) + +const ( + // WorkerErrorKindAuth 表示 worker credential 无效。 + WorkerErrorKindAuth = "auth" + // WorkerErrorKindPermission 表示当前 worker 无权访问资源。 + WorkerErrorKindPermission = "permission" + // WorkerErrorKindLeaseConflict 表示 work lease 已丢失或版本冲突。 + WorkerErrorKindLeaseConflict = "lease_conflict" + // WorkerErrorKindRateLimit 表示控制面限流。 + WorkerErrorKindRateLimit = "rate_limit" + // WorkerErrorKindTimeout 表示请求或工具执行超时。 + WorkerErrorKindTimeout = "timeout" + // WorkerErrorKindNetwork 表示网络层错误。 + WorkerErrorKindNetwork = "network" + // WorkerErrorKindToolError 表示本地工具执行错误。 + WorkerErrorKindToolError = "tool_error" + // WorkerErrorKindInvalidResponse 表示控制面响应无法解析。 + WorkerErrorKindInvalidResponse = "invalid_response" +) + +const ( + // DefaultWorkerClientType 保留用于源码兼容。 + // + // Deprecated: MA work API 不再接收 worker client header。 + DefaultWorkerClientType = "ark-self-hosted-worker-go" + // DefaultWorkerClientVersion 保留用于源码兼容。 + // + // Deprecated: MA work API 不再接收 worker client header。 + DefaultWorkerClientVersion = "0.1.0" +) + +// ErrEventStreamUnsupported 表示当前 API 实现不支持 SSE 事件流。 +var ErrEventStreamUnsupported = errors.New("ark event stream is not configured") + +// API 是 worker 依赖的最小控制面接口。 +type API interface { + PollWork(ctx context.Context, req PollWorkRequest) (*WorkItem, error) + AckWork(ctx context.Context, req AckWorkRequest) error + HeartbeatWork(ctx context.Context, req HeartbeatWorkRequest) (*HeartbeatResponse, error) + StopWork(ctx context.Context, req StopWorkRequest) error + GetSession(ctx context.Context, req GetSessionRequest) (*Session, error) + ListEvents(ctx context.Context, req ListEventsRequest) (*ListEventsResponse, error) + SendEvent(ctx context.Context, req SendEventRequest) error + OpenSkill(ctx context.Context, req OpenSkillRequest) (*SkillContent, error) +} + +// SkillResolver optionally enriches a session skill reference with authoritative metadata. +type SkillResolver interface { + ResolveSkill(ctx context.Context, skill SkillRef) (SkillRef, error) +} + +// PollWorkRequest 是 worker 轮询 work queue 的请求。 +type PollWorkRequest struct { + EnvironmentID string `json:"environment_id"` + WorkerID string `json:"worker_id,omitempty"` + // Deprecated: MA poll 每次最多返回一个 work。 + MaxItems int `json:"max_items,omitempty"` + BlockMS int `json:"block_ms,omitempty"` + ReclaimOlderThanMS int `json:"reclaim_older_than_ms,omitempty"` + // Deprecated: MA work API 不再接收 worker client header。 + WorkerClientType string `json:"-"` + // Deprecated: MA work API 不再接收 worker client header。 + WorkerClientVersion string `json:"-"` +} + +// AckWorkRequest 是 worker 确认已接收 work 的请求。 +type AckWorkRequest struct { + EnvironmentID string `json:"environment_id"` + WorkID string `json:"work_id"` + WorkerID string `json:"worker_id,omitempty"` + // Deprecated: MA ack 接口不再接收 lease_id。 + LeaseID string `json:"lease_id,omitempty"` +} + +// HeartbeatWorkRequest 是 worker 刷新 work lease 的请求。 +type HeartbeatWorkRequest struct { + EnvironmentID string `json:"environment_id"` + WorkID string `json:"work_id"` + ExpectedLastHeartbeat string `json:"expected_last_heartbeat,omitempty"` + DesiredTTLSeconds int `json:"desired_ttl_seconds,omitempty"` + // Deprecated: MA heartbeat 接口不再接收 lease_id。 + LeaseID string `json:"lease_id,omitempty"` + // Deprecated: MA heartbeat 接口不再接收 worker_id。 + WorkerID string `json:"worker_id,omitempty"` + // Deprecated: 使用 ExpectedLastHeartbeat 表达 CAS 期望。 + LastHeartbeat string `json:"last_heartbeat,omitempty"` + // Deprecated: MA heartbeat 接口不再接收 active_tool_call_id。 + ActiveToolCallID string `json:"active_tool_call_id,omitempty"` +} + +// HeartbeatResponse 是控制面返回的 heartbeat 结果。 +type HeartbeatResponse struct { + LastHeartbeat string `json:"last_heartbeat,omitempty"` + LeaseExtended *bool `json:"lease_extended,omitempty"` + State string `json:"state,omitempty"` + TTLSeconds int `json:"ttl_seconds,omitempty"` + Type string `json:"type,omitempty"` + // Deprecated: MA heartbeat 响应不再返回 lease_expires_at。 + LeaseExpiresAt string `json:"lease_expires_at,omitempty"` + // Deprecated: 使用 TTLSeconds。 + LeaseSeconds int `json:"lease_seconds,omitempty"` + // Deprecated: 使用 State。 + Status string `json:"status,omitempty"` + // Deprecated: 使用 State 判断 stopping/stopped。 + StopRequested bool `json:"stop_requested,omitempty"` + // Deprecated: MA heartbeat 响应体不再返回 request_id。 + RequestID string `json:"request_id,omitempty"` +} + +// StopWorkRequest 是 worker 请求控制面结束或停止 work 的请求。 +type StopWorkRequest struct { + EnvironmentID string `json:"environment_id"` + WorkID string `json:"work_id"` + Force bool `json:"force,omitempty"` + // Deprecated: MA stop 接口不再接收 lease_id。 + LeaseID string `json:"lease_id,omitempty"` + // Deprecated: MA stop 接口不再接收 worker_id。 + WorkerID string `json:"worker_id,omitempty"` + // Deprecated: MA stop 接口只接收 force。 + Reason string `json:"reason,omitempty"` + // Deprecated: MA stop 接口只接收 force。 + Message string `json:"message,omitempty"` +} + +// GetSessionRequest 是读取 session 配置的请求。 +type GetSessionRequest struct { + SessionID string `json:"session_id"` +} + +// ListEventsRequest 是读取 session 增量事件的请求。 +type ListEventsRequest struct { + SessionID string `json:"session_id"` + Page string `json:"page,omitempty"` + CreatedAtGt string `json:"created_at_gt,omitempty"` + Limit int `json:"limit,omitempty"` + Order string `json:"order,omitempty"` + Types []string `json:"types,omitempty"` + // Deprecated: MA list events 接口不接收 work_id。 + WorkID string `json:"work_id,omitempty"` + // Deprecated: MA list events 接口不接收 lease_id。 + LeaseID string `json:"lease_id,omitempty"` +} + +// StreamEventsRequest 是读取 session SSE 事件流的请求。 +type StreamEventsRequest struct { + SessionID string `json:"session_id"` + EventDeltas []string `json:"event_deltas,omitempty"` +} + +// SendEventRequest 是回写用户侧事件的请求。 +type SendEventRequest struct { + SessionID string `json:"session_id"` + Event Event `json:"event"` + // Deprecated: MA send events 接口不接收 Idempotency-Key header。 + IdempotencyKey string `json:"idempotency_key,omitempty"` + // Deprecated: MA send events 接口不接收 work_id。 + WorkID string `json:"work_id,omitempty"` + // Deprecated: MA send events 接口不接收 lease_id。 + LeaseID string `json:"lease_id,omitempty"` +} + +// OpenSkillRequest 是下载单个 skill 归档的请求。 +type OpenSkillRequest struct { + SessionID string `json:"session_id,omitempty"` + Skill SkillRef `json:"skill"` +} + +// WorkItem 表示控制面分配给 worker 的一份工作。 +type WorkItem struct { + ID string `json:"id"` + AcknowledgedAt string `json:"acknowledged_at,omitempty"` + CreatedAt string `json:"created_at,omitempty"` + Data WorkData `json:"data,omitempty"` + EnvironmentID string `json:"environment_id,omitempty"` + LatestHeartbeatAt string `json:"latest_heartbeat_at,omitempty"` + Tags []WorkTag `json:"tags,omitempty"` + Secret string `json:"secret,omitempty"` + StartedAt string `json:"started_at,omitempty"` + State string `json:"state,omitempty"` + StopRequestedAt string `json:"stop_requested_at,omitempty"` + StoppedAt string `json:"stopped_at,omitempty"` + Type string `json:"type,omitempty"` + // Deprecated: 使用 Data.ID。 + SessionID string `json:"session_id,omitempty"` + // Deprecated: 使用 State。 + Status string `json:"status,omitempty"` + // Deprecated: MA work item 不再返回 attempt。 + Attempt int `json:"attempt,omitempty"` + // Deprecated: MA work item 不再返回 lease_id。 + LeaseID string `json:"lease_id,omitempty"` + // Deprecated: MA work item 不再返回 lease_expires_at。 + LeaseExpiresAt string `json:"lease_expires_at,omitempty"` + // Deprecated: MA work item 不再返回 request_id。 + RequestID string `json:"request_id,omitempty"` + // Deprecated: MA work item 不再返回 ttl_seconds。 + TTLSeconds int `json:"ttl_seconds,omitempty"` + // Deprecated: MA work item 不再返回 lease_seconds。 + LeaseSeconds int `json:"lease_seconds,omitempty"` + // Deprecated: 使用 LatestHeartbeatAt。 + LastHeartbeat string `json:"last_heartbeat,omitempty"` +} + +// SessionIDValue 返回 work item 对应的 session id。 +func (w WorkItem) SessionIDValue() string { + if w.SessionID != "" { + return w.SessionID + } + if w.Data.SessionID != "" { + return w.Data.SessionID + } + if w.Data.ID != "" && (w.Data.Type == "" || w.Data.Type == "session") { + return w.Data.ID + } + return "" +} + +// LeaseTTLSeconds 返回 work item 的 lease TTL,兼容旧字段 lease_seconds。 +func (w WorkItem) LeaseTTLSeconds() int { + if w.TTLSeconds > 0 { + return w.TTLSeconds + } + return w.LeaseSeconds +} + +// LatestHeartbeatValue 返回 work item 的 heartbeat CAS 值,兼容旧字段 last_heartbeat。 +func (w WorkItem) LatestHeartbeatValue() string { + if w.LatestHeartbeatAt != "" { + return w.LatestHeartbeatAt + } + return w.LastHeartbeat +} + +// WorkTag 是 work item 关联的火山资源标签。 +type WorkTag struct { + Key string `json:"key"` + Value string `json:"value,omitempty"` +} + +// WorkData 是 work item 的业务载荷。 +type WorkData struct { + Type string `json:"type,omitempty"` + ID string `json:"id,omitempty"` + // Deprecated: 使用 ID。 + SessionID string `json:"session_id,omitempty"` +} + +// Session 是 worker 启动一个 session 所需的最小配置快照。 +type Session struct { + ID string `json:"id"` + Agent AgentConfig `json:"agent,omitempty"` + Skills []SkillRef `json:"skills,omitempty"` +} + +// SkillRefs 返回 session 绑定的 skills,兼容 skills 在顶层或 agent 下的两种形态。 +func (s Session) SkillRefs() []SkillRef { + if len(s.Skills) > 0 { + return s.Skills + } + return s.Agent.Skills +} + +// AgentConfig 是 session 内 agent 的最小配置。 +type AgentConfig struct { + Skills []SkillRef `json:"skills,omitempty"` +} + +// SkillRef 描述需要安装到 workdir 的 skill。 +type SkillRef struct { + Name string `json:"name"` + DisplayName string `json:"display_name,omitempty"` + ID string `json:"id,omitempty"` + SkillID string `json:"skill_id,omitempty"` + Type string `json:"type,omitempty"` + Version string `json:"version,omitempty"` + DownloadURL string `json:"download_url,omitempty"` +} + +// IDValue 返回 skill 的稳定 id。 +func (s SkillRef) IDValue() string { + if s.SkillID != "" { + return s.SkillID + } + return s.ID +} + +// NameValue 返回 skill 的展示名,缺失时回退到 skill id。 +func (s SkillRef) NameValue() string { + if s.Name != "" { + return s.Name + } + if s.DisplayName != "" { + return s.DisplayName + } + return s.IDValue() +} + +// SkillContent 是一次 skill 下载响应。 +type SkillContent struct { + Body io.ReadCloser + ContentLength int64 + FileName string + ContentType string +} + +// ListEventsResponse 是 session 事件列表响应。 +type ListEventsResponse struct { + Events []Event `json:"events"` + NextPage string `json:"next_page,omitempty"` +} + +// Event 是 worker 关心的扁平事件视图。 +type Event struct { + ID string `json:"id"` + Type string `json:"type"` + Name string `json:"name,omitempty"` + Input RawJSON `json:"input,omitempty"` + ProcessedAt string `json:"processed_at,omitempty"` + EvaluatedPermission string `json:"evaluated_permission,omitempty"` + SessionThreadID string `json:"session_thread_id,omitempty"` + ToolUseID string `json:"tool_use_id,omitempty"` + CustomToolUseID string `json:"custom_tool_use_id,omitempty"` + Result string `json:"result,omitempty"` + DenyMessage string `json:"deny_message,omitempty"` + StopReason *SessionStopReason `json:"stop_reason,omitempty"` + Content []ContentBlock `json:"content,omitempty"` + IsError *bool `json:"is_error,omitempty"` + Extra map[string]RawJSON `json:"-"` +} + +// UnmarshalJSON 兼容 input 是 JSON 对象或 JSON 字符串的事件形态。 +func (e *Event) UnmarshalJSON(data []byte) error { + type alias Event + var a alias + if err := json.Unmarshal(data, &a); err != nil { + return err + } + var raw map[string]RawJSON + _ = json.Unmarshal(data, &raw) + for _, k := range []string{ + "id", "type", "name", "input", "processed_at", "evaluated_permission", + "session_thread_id", "tool_use_id", "custom_tool_use_id", "result", + "deny_message", "stop_reason", "content", "is_error", + } { + delete(raw, k) + } + a.Extra = raw + *e = Event(a) + return nil +} + +// StopReasonType 返回 session idle stop_reason 的 union type。 +func (e Event) StopReasonType() string { + if e.StopReason == nil { + return "" + } + return e.StopReason.Type +} + +// SessionStopReason 是 MA session.status_idle.stop_reason 的扁平视图。 +type SessionStopReason struct { + Type string `json:"type,omitempty"` + EventIDs []string `json:"event_ids,omitempty"` + Raw RawJSON `json:"-"` +} + +// UnmarshalJSON 兼容 MA union object,并容忍早期字符串形态。 +func (r *SessionStopReason) UnmarshalJSON(data []byte) error { + if r == nil { + return nil + } + trimmed := bytes.TrimSpace(data) + if string(trimmed) == "null" || len(trimmed) == 0 { + *r = SessionStopReason{} + return nil + } + if trimmed[0] == '"' { + var typ string + if err := json.Unmarshal(trimmed, &typ); err != nil { + return err + } + *r = SessionStopReason{Type: typ, Raw: append((*r).Raw[:0], trimmed...)} + return nil + } + var raw struct { + Type string `json:"type"` + EventIDs []string `json:"event_ids"` + } + if err := json.Unmarshal(trimmed, &raw); err != nil { + return err + } + *r = SessionStopReason{ + Type: raw.Type, + EventIDs: append([]string(nil), raw.EventIDs...), + Raw: append((*r).Raw[:0], trimmed...), + } + return nil +} + +// ContentBlock 是 tool result 的内容块。 +type ContentBlock struct { + Type string `json:"type"` + Text string `json:"text,omitempty"` + MediaType string `json:"media_type,omitempty"` + Data []byte `json:"data,omitempty"` +} + +// RawJSON 保存未解释的 JSON 对象,兼容 wire 上以字符串承载 raw JSON。 +type RawJSON []byte + +// UnmarshalJSON 解析 raw JSON。 +func (r *RawJSON) UnmarshalJSON(data []byte) error { + if string(data) == "null" || len(data) == 0 { + *r = []byte("{}") + return nil + } + if data[0] == '"' { + var s string + if err := json.Unmarshal(data, &s); err != nil { + return err + } + if s == "" { + s = "{}" + } + *r = []byte(s) + return nil + } + *r = append((*r)[:0], data...) + return nil +} + +// MarshalJSON 输出原始 JSON。 +func (r RawJSON) MarshalJSON() ([]byte, error) { + if len(r) == 0 { + return []byte("{}"), nil + } + return r, nil +} + +// BoolPtr 返回 bool 指针。 +func BoolPtr(v bool) *bool { + return &v +} + +// NewUserToolResultEvent 构造 user.tool_result 事件。 +func NewUserToolResultEvent(toolUseID string, content []ContentBlock, isError bool, threadID string) Event { + return Event{ + ID: NewEventID("evt"), + Type: EventTypeUserToolResult, + ToolUseID: toolUseID, + Content: content, + IsError: BoolPtr(isError), + ProcessedAt: time.Now().UTC().Format(time.RFC3339Nano), + SessionThreadID: threadID, + } +} + +// NewUserCustomToolResultEvent 构造 user.custom_tool_result 事件。 +func NewUserCustomToolResultEvent(customToolUseID string, content []ContentBlock, isError bool, threadID string) Event { + return Event{ + ID: NewEventID("evt"), + Type: EventTypeUserCustomToolResult, + CustomToolUseID: customToolUseID, + Content: content, + IsError: BoolPtr(isError), + ProcessedAt: time.Now().UTC().Format(time.RFC3339Nano), + SessionThreadID: threadID, + } +} + +// NewEventID 生成事件 id。 +func NewEventID(prefix string) string { + var b [8]byte + if _, err := rand.Read(b[:]); err != nil { + return fmt.Sprintf("%s-%d", prefix, time.Now().UnixNano()) + } + return prefix + "-" + hex.EncodeToString(b[:]) +} + +// APIError 表示控制面返回的非 2xx 错误。 +type APIError struct { + StatusCode int + Message string + RequestID string +} + +// Error 返回错误字符串。 +func (e *APIError) Error() string { + return fmt.Sprintf("worker api status %d: %s", e.StatusCode, e.Message) +} + +// IsStatus 判断错误是否是指定 HTTP 状态码。 +func IsStatus(err error, status int) bool { + var apiErr *APIError + return errors.As(err, &apiErr) && apiErr.StatusCode == status +} + +// WorkerError 是 SDK 暴露给调用方的可分类错误。 +type WorkerError struct { + Kind string + RequestID string + Retryable bool + Message string + Err error +} + +// Error 返回错误字符串。 +func (e *WorkerError) Error() string { + if e == nil { + return "" + } + if e.Message != "" { + return e.Message + } + if e.Err != nil { + return e.Err.Error() + } + return e.Kind +} + +// Unwrap 返回底层错误。 +func (e *WorkerError) Unwrap() error { + if e == nil { + return nil + } + return e.Err +} diff --git a/arkruntime/session_events_raw.go b/arkruntime/session_events_raw.go new file mode 100644 index 0000000..3cdbd69 --- /dev/null +++ b/arkruntime/session_events_raw.go @@ -0,0 +1,49 @@ +// Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +// SPDX-License-Identifier: Apache-2.0 + +package arkruntime + +import ( + "context" + "errors" + "fmt" + "net/http" + + "github.com/volcengine/ark-runtime-go/arkruntime/model/session" +) + +// SendSessionEventsRaw sends raw session event payloads. +// +// This is primarily used by self-hosted workers, where the event shape is +// already produced by the local runner and must be passed through without +// losing forward-compatible fields. +func (c *Client) SendSessionEventsRaw( + ctx context.Context, + sessionID string, + events []any, + setters ...requestOption, +) error { + if sessionID == "" { + return errors.New("missing required session_id") + } + if len(events) == 0 { + return errors.New("missing required events") + } + body := map[string]any{"events": events} + u := c.fullURL(fmt.Sprintf("%s/%s/events", sessionsPrefix, session.PathEscape(sessionID))) + wrap := &session.SendSessionEventsResponseWrapper{} + return c.doControlPlaneRequest(ctx, http.MethodPost, u, wrap, append(setters, withBody(body))...) +} + +// SendSessionEventRaw sends one raw session event payload. +func (c *Client) SendSessionEventRaw( + ctx context.Context, + sessionID string, + event any, + setters ...requestOption, +) error { + if event == nil { + return errors.New("missing required event") + } + return c.SendSessionEventsRaw(ctx, sessionID, []any{event}, setters...) +} diff --git a/arkruntime/sessions.go b/arkruntime/sessions.go index a096c06..0daedb4 100644 --- a/arkruntime/sessions.go +++ b/arkruntime/sessions.go @@ -196,11 +196,29 @@ func (c *Client) StreamSessionEvents( ctx context.Context, sessionID string, setters ...requestOption, +) (*session.StreamDecoder, error) { + return c.StreamSessionEventsWithParams(ctx, sessionID, nil, setters...) +} + +// StreamSessionEventsWithParams opens a text/event-stream connection with +// query parameters such as event_deltas. +func (c *Client) StreamSessionEventsWithParams( + ctx context.Context, + sessionID string, + params *session.SessionEventsStreamEventsParams, + setters ...requestOption, ) (*session.StreamDecoder, error) { if sessionID == "" { return nil, errors.New("missing required session_id") } + q, qerr := session.URLQuerySessionEventsStream(params) + if qerr != nil { + return nil, qerr + } u := c.fullURL(fmt.Sprintf("%s/%s/events/stream", sessionsPrefix, session.PathEscape(sessionID))) + if encoded := q.Encode(); encoded != "" { + u = u + "?" + encoded + } opts := append(setters, withBody(nil), WithCustomHeader("Accept", "text/event-stream")) req, reqErr := c.newRequest(ctx, http.MethodGet, u, "", "", opts...) @@ -354,6 +372,17 @@ func (c *Client) StreamSessionThreadEvents( ctx context.Context, sessionID, threadID string, setters ...requestOption, +) (*session.StreamDecoder, error) { + return c.StreamSessionThreadEventsWithParams(ctx, sessionID, threadID, nil, setters...) +} + +// StreamSessionThreadEventsWithParams opens a thread-scoped SSE stream with +// query parameters such as event_deltas. +func (c *Client) StreamSessionThreadEventsWithParams( + ctx context.Context, + sessionID, threadID string, + params *session.SessionThreadsStreamEventsParams, + setters ...requestOption, ) (*session.StreamDecoder, error) { if sessionID == "" { return nil, errors.New("missing required session_id") @@ -361,7 +390,14 @@ func (c *Client) StreamSessionThreadEvents( if threadID == "" { return nil, errors.New("missing required thread_id") } + q, qerr := session.URLQuerySessionThreadEventsStream(params) + if qerr != nil { + return nil, qerr + } u := c.fullURL(sessionThreadPath(sessionID, threadID) + "/stream") + if encoded := q.Encode(); encoded != "" { + u = u + "?" + encoded + } opts := append(setters, withBody(nil), WithCustomHeader("Accept", "text/event-stream")) req, reqErr := c.newRequest(ctx, http.MethodGet, u, "", "", opts...) if reqErr != nil { diff --git a/arkruntime/skill_content.go b/arkruntime/skill_content.go new file mode 100644 index 0000000..02519cf --- /dev/null +++ b/arkruntime/skill_content.go @@ -0,0 +1,75 @@ +// Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +// SPDX-License-Identifier: Apache-2.0 + +package arkruntime + +import ( + "context" + "errors" + "io" + "net/http" + "path" + "strings" + + "github.com/volcengine/ark-runtime-go/arkruntime/model/skill" +) + +// SkillContent is a streaming Skill archive download response. +type SkillContent struct { + Body io.ReadCloser + ContentLength int64 + FileName string + ContentType string +} + +// OpenSkillContent opens the archive content for a specific Skill version. +func (c *Client) OpenSkillContent( + ctx context.Context, + skillID string, + version string, + setters ...requestOption, +) (*SkillContent, error) { + if skillID == "" { + return nil, errors.New("missing required skill_id") + } + if version == "" { + return nil, errors.New("missing required version") + } + u := c.fullURL( + strings.Join([]string{ + skillsPrefix, + skill.PathEscape(skillID), + "versions", + skill.PathEscape(version), + "content", + }, "/"), + ) + req, reqErr := c.newRequest(ctx, http.MethodGet, u, "", "", append(setters, withBody(nil))...) + if reqErr != nil { + return nil, reqErr + } + resp, err := c.config.HTTPClient.Do(req) + if err != nil { + return nil, err + } + if resp.StatusCode >= 300 { + defer resp.Body.Close() //nolint:errcheck // response body close errors are non-actionable + return nil, c.handleErrorResp(resp) + } + return &SkillContent{ + Body: resp.Body, + ContentLength: resp.ContentLength, + FileName: path.Base(resp.Request.URL.Path), + ContentType: resp.Header.Get("Content-Type"), + }, nil +} + +// DownloadSkillVersionContent is an alias for OpenSkillContent. +func (c *Client) DownloadSkillVersionContent( + ctx context.Context, + skillID string, + version string, + setters ...requestOption, +) (*SkillContent, error) { + return c.OpenSkillContent(ctx, skillID, version, setters...) +} diff --git a/arkruntime/skills.go b/arkruntime/skills.go index 732a377..a2e02dc 100644 --- a/arkruntime/skills.go +++ b/arkruntime/skills.go @@ -16,6 +16,12 @@ import ( const skillsPrefix = "/skills" +// CreateSkillOptions controls optional multipart metadata for CreateSkill. +type CreateSkillOptions struct { + DisplayTitle string + ProtectionEnabled *bool +} + // CreateSkill uploads a zip package as multipart/form-data and creates a Skill. // `fileReader` supplies the zip bytes; `displayTitle` is optional. func (c *Client) CreateSkill( @@ -23,11 +29,29 @@ func (c *Client) CreateSkill( fileReader io.Reader, fileName, displayTitle string, setters ...requestOption, +) (*skill.Skill, error) { + return c.CreateSkillWithOptions(ctx, fileReader, fileName, CreateSkillOptions{ + DisplayTitle: displayTitle, + }, setters...) +} + +// CreateSkillWithOptions uploads a zip package with optional Skill metadata. +func (c *Client) CreateSkillWithOptions( + ctx context.Context, + fileReader io.Reader, + fileName string, + options CreateSkillOptions, + setters ...requestOption, ) (*skill.Skill, error) { if fileReader == nil { return nil, errors.New("missing required file reader") } - form := &skill.UploadForm{File: fileReader, FileName: fileName, DisplayTitle: displayTitle} + form := &skill.UploadForm{ + File: fileReader, + FileName: fileName, + DisplayTitle: options.DisplayTitle, + ProtectionEnabled: options.ProtectionEnabled, + } body, contentType, merr := form.MarshalMultipart() if merr != nil { return nil, merr diff --git a/arkruntime/tools/agenttoolset/agenttoolset.go b/arkruntime/tools/agenttoolset/agenttoolset.go new file mode 100644 index 0000000..f935d0b --- /dev/null +++ b/arkruntime/tools/agenttoolset/agenttoolset.go @@ -0,0 +1,91 @@ +// Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +// SPDX-License-Identifier: Apache-2.0 +// Package agenttoolset 提供 MA agent 默认内置工具集的公共入口。 +package agenttoolset + +import ( + "context" + "time" + + selfhosted "github.com/volcengine/ark-runtime-go/arkruntime/selfhosted" + "github.com/volcengine/ark-runtime-go/arkruntime/selfhosted/envinit" + "github.com/volcengine/ark-runtime-go/arkruntime/toolset" +) + +type ( + // Tool 是 worker 可执行的本地工具。 + Tool = toolset.Tool + // Set 是可执行工具集合。 + Set = toolset.Set + // Result 是一次工具执行结果。 + Result = toolset.Result + // ContentBlock 是工具结果内容块。 + ContentBlock = toolset.ContentBlock + // Limits 是工具执行的尺寸与数量上限。 + Limits = toolset.Limits +) + +// AgentToolContext 携带单个 session 绑定的本地工具执行上下文。 +type AgentToolContext struct { + Workdir string + UnrestrictedPaths bool + // Env 非 nil 时完全替换继承的进程环境;敏感凭证字段始终会被移除。 + Env map[string]string + Limits Limits + BashPath string + ToolTimeout time.Duration + MaxArchiveBytes int64 + MaxExtractedBytes int64 + MaxArchiveEntries int +} + +// AgentToolset20260401 返回 agent_toolset_20260401 对应的内置工具集合。 +func AgentToolset20260401(env *AgentToolContext) (*Set, error) { + if env == nil { + env = &AgentToolContext{} + } + return toolset.NewDefault(env.ToolOptions()) +} + +// CloseAll 释放工具集合持有的资源。 +func CloseAll(tools *Set) { + if tools != nil { + _ = tools.Close() + } +} + +// SetupSkills 下载并安装 session 绑定的 skills。 +func (e *AgentToolContext) SetupSkills(ctx context.Context, api selfhosted.API, session *selfhosted.Session) error { + if e == nil { + e = &AgentToolContext{} + } + return envinit.New(api, e.InitOptions()).Setup(ctx, session) +} + +// ToolOptions 转换为内部工具集合配置。 +func (e *AgentToolContext) ToolOptions() toolset.Options { + if e == nil { + return toolset.Options{} + } + return toolset.Options{ + Workdir: e.Workdir, + UnrestrictedPaths: e.UnrestrictedPaths, + Env: e.Env, + Limits: e.Limits, + BashPath: e.BashPath, + ToolTimeout: e.ToolTimeout, + } +} + +// InitOptions 转换为内部 session 初始化配置。 +func (e *AgentToolContext) InitOptions() envinit.Options { + if e == nil { + return envinit.Options{} + } + return envinit.Options{ + Workdir: e.Workdir, + MaxArchiveBytes: e.MaxArchiveBytes, + MaxExtractedBytes: e.MaxExtractedBytes, + MaxArchiveEntries: e.MaxArchiveEntries, + } +} diff --git a/arkruntime/toolset/bash.go b/arkruntime/toolset/bash.go new file mode 100644 index 0000000..d9c2fe5 --- /dev/null +++ b/arkruntime/toolset/bash.go @@ -0,0 +1,475 @@ +// Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +// SPDX-License-Identifier: Apache-2.0 +package toolset + +import ( + "bytes" + "context" + "crypto/rand" + "encoding/hex" + "encoding/json" + "errors" + "fmt" + "io" + "os" + "os/exec" + "path/filepath" + "strconv" + "strings" + "sync" + "time" +) + +// BashTool 实现持久 bash 工具。 +type BashTool struct { + session *BashSession + timeout time.Duration +} + +// NewBashTool 创建 bash 工具。 +func NewBashTool(opts Options) (*BashTool, error) { + if opts.Limits == (Limits{}) { + opts.Limits = DefaultLimits() + } + if opts.ToolTimeout <= 0 { + opts.ToolTimeout = 120 * time.Second + } + session, err := NewBashSession(opts) + if err != nil { + return nil, err + } + return &BashTool{session: session, timeout: opts.ToolTimeout}, nil +} + +// Name 返回工具名。 +func (t *BashTool) Name() string { return "bash" } + +// Execute 执行 bash 命令。 +func (t *BashTool) Execute(ctx context.Context, input json.RawMessage) Result { + var req struct { + Command string `json:"command"` + Cmd string `json:"cmd,omitempty"` + TimeoutMS int `json:"timeout_ms,omitempty"` + Restart bool `json:"restart,omitempty"` + } + if err := decodeInput(input, &req); err != nil { + return ErrorResult(err.Error()) + } + if req.Restart { + if err := t.session.Restart(); err != nil { + return ErrorResult(err.Error()) + } + return TextResult("bash session restarted") + } + command := req.Command + if command == "" { + command = req.Cmd + } + if command == "" { + return ErrorResult("command must not be empty") + } + timeout := t.timeout + if req.TimeoutMS > 0 { + timeout = time.Duration(req.TimeoutMS) * time.Millisecond + } + out, exit, timedOut, err := t.session.Run(ctx, command, timeout) + if err != nil { + return ErrorResult(err.Error()) + } + if timedOut { + return ErrorResult(fmt.Sprintf("command timed out after %s\n%s", timeout, out)) + } + return TextResult(fmt.Sprintf("exit_code: %d\n%s", exit, out)) +} + +// Close 关闭 bash 会话。 +func (t *BashTool) Close() error { + return t.session.Close() +} + +// BashSession 是 per-worker 的持久 bash 子进程。 +type BashSession struct { + mu sync.Mutex + cond *sync.Cond + cmd *exec.Cmd + stdin *os.File + outR *os.File + ctrlR *os.File + done chan struct{} + buf bytes.Buffer + ctrlBuf bytes.Buffer + dead bool + gen int64 + workdir string + privateDir string + env map[string]string + bashPath string + limits Limits +} + +// NewBashSession 创建持久 bash 会话。 +func NewBashSession(opts Options) (*BashSession, error) { + if opts.BashPath == "" { + opts.BashPath = "/bin/bash" + } + if opts.Limits == (Limits{}) { + opts.Limits = DefaultLimits() + } + if opts.Workdir == "" { + return nil, errors.New("workdir must not be empty") + } + if err := os.MkdirAll(opts.Workdir, 0o755); err != nil { + return nil, err + } + s := &BashSession{ + workdir: opts.Workdir, + env: opts.Env, + bashPath: opts.BashPath, + limits: opts.Limits, + } + s.cond = sync.NewCond(&s.mu) + if err := s.startLocked(); err != nil { + _ = os.RemoveAll(s.privateDir) + return nil, err + } + return s, nil +} + +// Restart 重建持久 bash。 +func (s *BashSession) Restart() error { + s.mu.Lock() + defer s.mu.Unlock() + s.closeLocked() + return s.startLocked() +} + +// Run 在持久 bash 中执行命令。 +func (s *BashSession) Run(ctx context.Context, command string, timeout time.Duration) (string, int, bool, error) { + s.mu.Lock() + defer s.mu.Unlock() + if s.dead { + if err := s.startLocked(); err != nil { + return "", -1, false, err + } + } + cmdID := nonce() + controlID := nonce() + info, err := os.Lstat(s.privateDir) + if err != nil { + return "", -1, false, err + } + if !info.IsDir() || info.Mode()&os.ModeSymlink != 0 { + return "", -1, false, errors.New("bash private directory is not a directory") + } + cmdPath := filepath.Join(s.privateDir, cmdID+".sh") + statePath := filepath.Join(s.privateDir, cmdID+".state") + if err := os.WriteFile(cmdPath, []byte(command), 0o600); err != nil { + return "", -1, false, err + } + defer func() { _ = os.Remove(cmdPath) }() + defer func() { _ = os.Remove(statePath) }() + + s.buf.Reset() + s.ctrlBuf.Reset() + frame := s.execFrame(cmdPath, statePath, controlID) + if _, err := io.WriteString(s.stdin, frame); err != nil { + s.dead = true + return "", -1, false, fmt.Errorf("write bash stdin: %w", err) + } + + deadline := time.Now().Add(timeout) + wakeDone := make(chan struct{}) + go func() { + timer := time.NewTimer(timeout) + defer timer.Stop() + select { + case <-ctx.Done(): + case <-timer.C: + case <-wakeDone: + return + } + s.mu.Lock() + s.cond.Broadcast() + s.mu.Unlock() + }() + defer close(wakeDone) + + for { + if _, exit, consumed, ok := parseControl(s.ctrlBuf.Bytes(), controlID); ok { + remaining := append([]byte(nil), s.ctrlBuf.Bytes()[consumed:]...) + s.ctrlBuf.Reset() + _, _ = s.ctrlBuf.Write(remaining) + return truncateOutput(s.buf.String(), s.limits.MaxOutputBytes), exit, false, nil + } + if s.dead { + return "", -1, false, errors.New("persistent bash terminated before command completed") + } + if err := ctx.Err(); err != nil { + s.closeLocked() + return "", -1, true, err + } + if timeout > 0 && time.Now().After(deadline) { + partial := truncateOutput(s.buf.String(), s.limits.MaxOutputBytes) + s.closeLocked() + return partial, -1, true, nil + } + s.cond.Wait() + } +} + +// Close 关闭持久 bash。 +func (s *BashSession) Close() error { + s.mu.Lock() + defer s.mu.Unlock() + s.closeLocked() + err := os.RemoveAll(s.privateDir) + s.privateDir = "" + return err +} + +func (s *BashSession) startLocked() error { + if s.privateDir == "" { + privateDir, err := os.MkdirTemp("", "ark-self-host-shell-*") + if err != nil { + return err + } + s.privateDir = privateDir + } + cmd := exec.Command(s.bashPath, "--noprofile", "--norc") + cmd.Dir = s.workdir + cmd.Env = scrubbedEnv(s.env) + setProcessGroup(cmd) + + inR, inW, err := os.Pipe() + if err != nil { + return err + } + outR, outW, err := os.Pipe() + if err != nil { + _ = inR.Close() + _ = inW.Close() + return err + } + ctrlR, ctrlW, err := os.Pipe() + if err != nil { + _ = inR.Close() + _ = inW.Close() + _ = outR.Close() + _ = outW.Close() + return err + } + cmd.Stdin = inR + cmd.Stdout = outW + cmd.Stderr = outW + cmd.ExtraFiles = []*os.File{ctrlW} + if err := cmd.Start(); err != nil { + _ = inR.Close() + _ = inW.Close() + _ = outR.Close() + _ = outW.Close() + _ = ctrlR.Close() + _ = ctrlW.Close() + return err + } + _ = inR.Close() + _ = outW.Close() + _ = ctrlW.Close() + s.cmd = cmd + s.stdin = inW + s.outR = outR + s.ctrlR = ctrlR + s.done = make(chan struct{}) + s.buf.Reset() + s.ctrlBuf.Reset() + s.dead = false + s.gen++ + gen := s.gen + go s.readOutputLoop(gen, outR) + go s.readControlLoop(gen, ctrlR) + go s.waitLoop(gen, cmd, s.done) + return nil +} + +func (s *BashSession) closeLocked() { + if s.stdin != nil { + _ = s.stdin.Close() + s.stdin = nil + } + if s.outR != nil { + _ = s.outR.Close() + s.outR = nil + } + if s.ctrlR != nil { + _ = s.ctrlR.Close() + s.ctrlR = nil + } + killCommandGroup(s.cmd) + s.cmd = nil + s.dead = true + s.buf.Reset() + s.ctrlBuf.Reset() + s.cond.Broadcast() +} + +func (s *BashSession) readOutputLoop(gen int64, r *os.File) { + buf := make([]byte, 4096) + for { + n, err := r.Read(buf) + s.mu.Lock() + current := gen == s.gen + if current && n > 0 { + _, _ = s.buf.Write(buf[:n]) + if s.limits.MaxOutputBytes > 0 && int64(s.buf.Len()) > s.limits.MaxOutputBytes*4 { + data := s.buf.Bytes() + keep := int(s.limits.MaxOutputBytes * 2) + if keep < len(data) { + s.buf.Reset() + _, _ = s.buf.Write(data[len(data)-keep:]) + } + } + } + if current && err != nil { + s.dead = true + } + if current { + s.cond.Broadcast() + } + s.mu.Unlock() + if err != nil { + return + } + } +} + +func (s *BashSession) readControlLoop(gen int64, r *os.File) { + buf := make([]byte, 256) + for { + n, err := r.Read(buf) + s.mu.Lock() + current := gen == s.gen + if current && n > 0 { + _, _ = s.ctrlBuf.Write(buf[:n]) + } + if current && err != nil { + s.dead = true + } + if current { + s.cond.Broadcast() + } + s.mu.Unlock() + if err != nil { + return + } + } +} + +func (s *BashSession) waitLoop(gen int64, cmd *exec.Cmd, done chan struct{}) { + _ = cmd.Wait() + close(done) + s.mu.Lock() + if gen == s.gen { + s.dead = true + s.cond.Broadcast() + } + s.mu.Unlock() +} + +func (s *BashSession) execFrame(cmdPath, statePath, nonce string) string { + return fmt.Sprintf( + "builtin set +e; builtin set +o pipefail\n"+ + "__ark_st=%s\n"+ + "( builtin trap '__ark_c=$?; { builtin printf \"builtin cd %%q\\n\" \"$PWD\"; builtin set +o | command grep -vE \" (errexit|pipefail|xtrace)$\"; ( builtin unset __ark_st __ark_c __ark_ec; builtin declare -p ); builtin declare -f; } > \"$__ark_st\" 2>/dev/null; builtin exit $__ark_c' EXIT\n"+ + " . %s ) 3>&-\n"+ + "__ark_ec=$?\n"+ + "{ builtin set +e; . \"$__ark_st\"; } >/dev/null 2>&1 3>&-\n"+ + "command rm -f %s \"$__ark_st\"\n"+ + "builtin printf '__MA_WORKER_DONE_%s:%%d\\n' \"$__ark_ec\" >&3\n", + shellQuote(statePath), shellQuote(cmdPath), shellQuote(cmdPath), nonce, + ) +} + +func parseControl(data []byte, nonce string) ([]byte, int, int, bool) { + prefix := []byte("__MA_WORKER_DONE_" + nonce + ":") + idx := bytes.Index(data, prefix) + if idx < 0 { + return nil, 0, 0, false + } + start := idx + len(prefix) + endRel := bytes.IndexByte(data[start:], '\n') + if endRel < 0 { + return nil, 0, 0, false + } + end := start + endRel + exit, err := strconv.Atoi(string(data[start:end])) + if err != nil { + return nil, 0, 0, false + } + output := append([]byte(nil), data[:idx]...) + return output, exit, end + 1, true +} + +func scrubbedEnv(extra map[string]string) []string { + if extra != nil { + env := make([]string, 0, len(extra)) + for key, value := range extra { + if !isSensitiveEnvKey(key) { + env = append(env, key+"="+value) + } + } + return env + } + env := make([]string, 0, len(os.Environ())) + for _, item := range os.Environ() { + key, _, _ := strings.Cut(item, "=") + if !isSensitiveEnvKey(key) { + env = append(env, item) + } + } + return env +} + +func isSensitiveEnvKey(key string) bool { + upper := strings.ToUpper(strings.TrimSpace(key)) + for _, prefix := range []string{"ARK_", "MA_", "ANTHROPIC_", "OPENAI_", "AWS_", "AZURE_", "GOOGLE_"} { + if strings.HasPrefix(upper, prefix) { + return true + } + } + for _, exact := range []string{ + "VOLC_ACCESSKEY", "VOLC_SECRETKEY", "BYTEPLUS_ACCESSKEY", "BYTEPLUS_SECRETKEY", + "TOKEN", "SECRET", "PASSWORD", "PASSWD", "PRIVATE_KEY", "API_KEY", "ACCESS_KEY", "SECRET_KEY", + } { + if upper == exact { + return true + } + } + for _, suffix := range []string{"_TOKEN", "_SECRET", "_PASSWORD", "_PASSWD", "_PRIVATE_KEY", "_API_KEY", "_ACCESS_KEY", "_SECRET_KEY"} { + if strings.HasSuffix(upper, suffix) { + return true + } + } + return false +} + +func shellQuote(s string) string { + return "'" + strings.ReplaceAll(s, "'", "'\\''") + "'" +} + +func nonce() string { + var b [8]byte + if _, err := rand.Read(b[:]); err != nil { + return fmt.Sprintf("%d", time.Now().UnixNano()) + } + return hex.EncodeToString(b[:]) +} + +func truncateOutput(output string, limit int64) string { + if limit <= 0 || int64(len(output)) <= limit { + return output + } + const marker = "\n[output truncated]\n" + if limit <= int64(len(marker)) { + return marker[:limit] + } + return output[:limit-int64(len(marker))] + marker +} diff --git a/arkruntime/toolset/file.go b/arkruntime/toolset/file.go new file mode 100644 index 0000000..a0b5836 --- /dev/null +++ b/arkruntime/toolset/file.go @@ -0,0 +1,268 @@ +// Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +// SPDX-License-Identifier: Apache-2.0 +package toolset + +import ( + "context" + "encoding/json" + "fmt" + "os" + "path/filepath" + "strings" + "unicode/utf8" +) + +// ReadTool 实现 read 工具。 +type ReadTool struct { + resolver *Resolver + limits Limits +} + +// NewReadTool 创建 read 工具。 +func NewReadTool(resolver *Resolver, limits Limits) *ReadTool { + return &ReadTool{resolver: resolver, limits: limits} +} + +// Name 返回工具名。 +func (t *ReadTool) Name() string { return "read" } + +// Execute 执行 read。 +func (t *ReadTool) Execute(_ context.Context, input json.RawMessage) Result { + var req struct { + FilePath string `json:"file_path"` + ViewRange []int `json:"view_range,omitempty"` + Offset int `json:"offset,omitempty"` + Limit int `json:"limit,omitempty"` + } + if err := decodeInput(input, &req); err != nil { + return ErrorResult(err.Error()) + } + host, err := t.resolver.ResolveExisting(req.FilePath) + if err != nil { + return ErrorResult(err.Error()) + } + info, err := os.Stat(host) + if err != nil { + return ErrorResult(err.Error()) + } + if !info.Mode().IsRegular() { + return ErrorResult("path is not a regular file") + } + if t.limits.MaxInputFileBytes > 0 && info.Size() > t.limits.MaxInputFileBytes { + return ErrorResult(fmt.Sprintf("file too large: %d bytes", info.Size())) + } + data, err := os.ReadFile(host) + if err != nil { + return ErrorResult(err.Error()) + } + if !utf8.Valid(data) { + return ErrorResult("binary file cannot be read directly") + } + if len(req.ViewRange) > 0 { + if len(req.ViewRange) != 2 { + return ErrorResult("view_range must be [start_line, end_line]") + } + lines := strings.Split(string(data), "\n") + start := req.ViewRange[0] - 1 + if start < 0 { + start = 0 + } + if start >= len(lines) { + return TextResult("") + } + end := len(lines) + if req.ViewRange[1] > 0 && req.ViewRange[1] < end { + end = req.ViewRange[1] + } + if end < start { + return ErrorResult(fmt.Sprintf("view_range end line %d is before start line %d", req.ViewRange[1], req.ViewRange[0])) + } + return TextResult(strings.Join(lines[start:end], "\n")) + } + if req.Offset == 0 && req.Limit == 0 { + return TextResult(string(data)) + } + lines := strings.Split(string(data), "\n") + if len(lines) > 0 && lines[len(lines)-1] == "" { + lines = lines[:len(lines)-1] + } + offset := req.Offset + if offset < 0 { + offset = 0 + } + if offset > len(lines) { + offset = len(lines) + } + limit := req.Limit + if limit <= 0 { + limit = t.limits.ReadDefaultLines + } + if limit <= 0 { + limit = 2000 + } + end := offset + limit + if end > len(lines) { + end = len(lines) + } + var b strings.Builder + for i := offset; i < end; i++ { + line := lines[i] + if t.limits.ReadMaxLineChars > 0 && utf8.RuneCountInString(line) > t.limits.ReadMaxLineChars { + line = string([]rune(line)[:t.limits.ReadMaxLineChars]) + } + fmt.Fprintf(&b, "%6d\t%s\n", i+1, line) + } + if end < len(lines) { + fmt.Fprintf(&b, "\n[truncated: showing lines %d-%d of %d]\n", offset+1, end, len(lines)) + } + return TextResult(b.String()) +} + +// WriteTool 实现 write 工具。 +type WriteTool struct { + resolver *Resolver + limits Limits +} + +// NewWriteTool 创建 write 工具。 +func NewWriteTool(resolver *Resolver, limits Limits) *WriteTool { + return &WriteTool{resolver: resolver, limits: limits} +} + +// Name 返回工具名。 +func (t *WriteTool) Name() string { return "write" } + +// Execute 执行 write。 +func (t *WriteTool) Execute(_ context.Context, input json.RawMessage) Result { + var req struct { + FilePath string `json:"file_path"` + Content string `json:"content"` + } + if err := decodeInput(input, &req); err != nil { + return ErrorResult(err.Error()) + } + if t.limits.MaxInputFileBytes > 0 && int64(len(req.Content)) > t.limits.MaxInputFileBytes { + return ErrorResult(fmt.Sprintf("content too large: %d bytes", len(req.Content))) + } + host, err := t.resolver.ResolveForWrite(req.FilePath) + if err != nil { + return ErrorResult(err.Error()) + } + if err := os.MkdirAll(filepath.Dir(host), 0o755); err != nil { + return ErrorResult(err.Error()) + } + verified, err := t.resolver.ResolveForWrite(req.FilePath) + if err != nil { + return ErrorResult(err.Error()) + } + if verified != host { + return ErrorResult("path resolution changed while writing") + } + if err := writeFileAtomically(host, []byte(req.Content), 0o600); err != nil { + return ErrorResult(err.Error()) + } + return TextResult(fmt.Sprintf("wrote %d bytes to %s", len(req.Content), req.FilePath)) +} + +// EditTool 实现 edit 工具。 +type EditTool struct { + resolver *Resolver + limits Limits +} + +// NewEditTool 创建 edit 工具。 +func NewEditTool(resolver *Resolver, limits Limits) *EditTool { + return &EditTool{resolver: resolver, limits: limits} +} + +// Name 返回工具名。 +func (t *EditTool) Name() string { return "edit" } + +// Execute 执行 edit。 +func (t *EditTool) Execute(_ context.Context, input json.RawMessage) Result { + var req struct { + FilePath string `json:"file_path"` + OldString string `json:"old_string"` + NewString string `json:"new_string"` + ReplaceAll bool `json:"replace_all,omitempty"` + } + if err := decodeInput(input, &req); err != nil { + return ErrorResult(err.Error()) + } + if req.OldString == "" { + return ErrorResult("old_string must not be empty") + } + host, err := t.resolver.ResolveExisting(req.FilePath) + if err != nil { + return ErrorResult(err.Error()) + } + writable, err := t.resolver.ResolveForWrite(req.FilePath) + if err != nil { + return ErrorResult(err.Error()) + } + if writable != host { + return ErrorResult("path resolution changed while editing") + } + info, err := os.Stat(host) + if err != nil { + return ErrorResult(err.Error()) + } + if !info.Mode().IsRegular() { + return ErrorResult("path is not a regular file") + } + if t.limits.MaxInputFileBytes > 0 && info.Size() > t.limits.MaxInputFileBytes { + return ErrorResult(fmt.Sprintf("file too large: %d bytes", info.Size())) + } + data, err := os.ReadFile(host) + if err != nil { + return ErrorResult(err.Error()) + } + content := string(data) + count := strings.Count(content, req.OldString) + switch { + case count == 0: + return ErrorResult("old_string not found") + case count > 1 && !req.ReplaceAll: + return ErrorResult(fmt.Sprintf("old_string is not unique: %d matches", count)) + } + n := 1 + if req.ReplaceAll { + n = -1 + } + next := strings.Replace(content, req.OldString, req.NewString, n) + if int64(len(next)) > t.limits.MaxInputFileBytes && t.limits.MaxInputFileBytes > 0 { + return ErrorResult(fmt.Sprintf("edited file too large: %d bytes", len(next))) + } + verified, err := t.resolver.ResolveForWrite(req.FilePath) + if err != nil { + return ErrorResult(err.Error()) + } + if verified != host { + return ErrorResult("path resolution changed while editing") + } + if err := writeFileAtomically(host, []byte(next), info.Mode().Perm()); err != nil { + return ErrorResult(err.Error()) + } + return TextResult(fmt.Sprintf("replaced %d occurrence(s) in %s", count, req.FilePath)) +} + +func writeFileAtomically(target string, data []byte, mode os.FileMode) error { + tmp, err := os.CreateTemp(filepath.Dir(target), ".ark-write-*") + if err != nil { + return err + } + tmpName := tmp.Name() + defer func() { _ = os.Remove(tmpName) }() + if err := tmp.Chmod(mode); err != nil { + _ = tmp.Close() + return err + } + if _, err := tmp.Write(data); err != nil { + _ = tmp.Close() + return err + } + if err := tmp.Close(); err != nil { + return err + } + return os.Rename(tmpName, target) +} diff --git a/arkruntime/toolset/path.go b/arkruntime/toolset/path.go new file mode 100644 index 0000000..e838994 --- /dev/null +++ b/arkruntime/toolset/path.go @@ -0,0 +1,145 @@ +// Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +// SPDX-License-Identifier: Apache-2.0 +package toolset + +import ( + "errors" + "fmt" + "os" + "path/filepath" + "strings" +) + +var errPathEscape = errors.New("path escapes workdir") + +// Resolver 负责把工具输入路径解析到 workdir 内的真实路径。 +type Resolver struct { + root string + rootReal string + unrestricted bool +} + +// NewResolver 创建 workdir 路径解析器。 +func NewResolver(root string) (*Resolver, error) { + return NewResolverWithOptions(root, false) +} + +// NewResolverWithOptions 创建 workdir 路径解析器。 +func NewResolverWithOptions(root string, unrestricted bool) (*Resolver, error) { + if root == "" { + return nil, errors.New("workdir must not be empty") + } + abs, err := filepath.Abs(root) + if err != nil { + return nil, err + } + if err := os.MkdirAll(abs, 0o755); err != nil { + return nil, err + } + real, err := filepath.EvalSymlinks(abs) + if err != nil { + return nil, err + } + return &Resolver{root: abs, rootReal: real, unrestricted: unrestricted}, nil +} + +// Root 返回 workdir 绝对路径。 +func (r *Resolver) Root() string { + return r.root +} + +// ResolveExisting 解析一个必须已存在的路径。 +func (r *Resolver) ResolveExisting(p string) (string, error) { + clean, err := r.clean(p) + if err != nil { + return "", err + } + real, err := filepath.EvalSymlinks(clean) + if err != nil { + if os.IsNotExist(err) { + return "", fmt.Errorf("not found: %s", p) + } + return "", err + } + if !r.unrestricted && !withinRoot(r.rootReal, real) { + return "", errPathEscape + } + return real, nil +} + +// ResolveForWrite 解析写入目标,允许目标尚不存在。 +func (r *Resolver) ResolveForWrite(p string) (string, error) { + clean, err := r.clean(p) + if err != nil { + return "", err + } + dir, base := filepath.Split(clean) + realDir, err := r.resolveLongest(filepath.Clean(dir)) + if err != nil { + return "", err + } + if !r.unrestricted && !withinRoot(r.rootReal, realDir) { + return "", errPathEscape + } + target := filepath.Join(realDir, base) + if real, err := filepath.EvalSymlinks(target); err == nil { + if !r.unrestricted && !withinRoot(r.rootReal, real) { + return "", errPathEscape + } + return real, nil + } else if !os.IsNotExist(err) { + return "", err + } + if info, err := os.Lstat(target); err == nil && info.Mode()&os.ModeSymlink != 0 && !r.unrestricted { + return "", fmt.Errorf("%w: unresolved symbolic link", errPathEscape) + } + return target, nil +} + +func (r *Resolver) clean(p string) (string, error) { + if p == "" { + return "", errors.New("path must not be empty") + } + switch { + case p == "~": + p = r.root + case strings.HasPrefix(p, "~/"): + p = filepath.Join(r.root, p[2:]) + case !filepath.IsAbs(p): + p = filepath.Join(r.root, p) + } + return filepath.Clean(p), nil +} + +func (r *Resolver) resolveLongest(clean string) (string, error) { + cur := clean + var rest []string + for { + real, err := filepath.EvalSymlinks(cur) + if err == nil { + return filepath.Join(append([]string{real}, rest...)...), nil + } + if !os.IsNotExist(err) { + return "", err + } + if info, lstatErr := os.Lstat(cur); lstatErr == nil && info.Mode()&os.ModeSymlink != 0 && !r.unrestricted { + return "", fmt.Errorf("%w: unresolved symbolic link", errPathEscape) + } else if lstatErr != nil && !os.IsNotExist(lstatErr) { + return "", lstatErr + } + parent := filepath.Dir(cur) + if parent == cur { + return clean, nil + } + rest = append([]string{filepath.Base(cur)}, rest...) + cur = parent + } +} + +func withinRoot(root, target string) bool { + rel, err := filepath.Rel(root, target) + if err != nil { + return false + } + return rel == "." || (rel != ".." && !strings.HasPrefix(rel, ".."+string(os.PathSeparator)) && !filepath.IsAbs(rel)) +} diff --git a/arkruntime/toolset/process_other.go b/arkruntime/toolset/process_other.go new file mode 100644 index 0000000..70b594e --- /dev/null +++ b/arkruntime/toolset/process_other.go @@ -0,0 +1,15 @@ +// Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +// SPDX-License-Identifier: Apache-2.0 +//go:build !unix + +package toolset + +import "os/exec" + +func setProcessGroup(_ *exec.Cmd) {} + +func killCommandGroup(cmd *exec.Cmd) { + if cmd != nil && cmd.Process != nil { + _ = cmd.Process.Kill() + } +} diff --git a/arkruntime/toolset/process_unix.go b/arkruntime/toolset/process_unix.go new file mode 100644 index 0000000..fa1bad3 --- /dev/null +++ b/arkruntime/toolset/process_unix.go @@ -0,0 +1,21 @@ +// Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +// SPDX-License-Identifier: Apache-2.0 +//go:build unix + +package toolset + +import ( + "os/exec" + "syscall" +) + +func setProcessGroup(cmd *exec.Cmd) { + cmd.SysProcAttr = &syscall.SysProcAttr{Setpgid: true} +} + +func killCommandGroup(cmd *exec.Cmd) { + if cmd == nil || cmd.Process == nil { + return + } + _ = syscall.Kill(-cmd.Process.Pid, syscall.SIGKILL) +} diff --git a/arkruntime/toolset/search.go b/arkruntime/toolset/search.go new file mode 100644 index 0000000..66dbed3 --- /dev/null +++ b/arkruntime/toolset/search.go @@ -0,0 +1,339 @@ +// Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +// SPDX-License-Identifier: Apache-2.0 +package toolset + +import ( + "bufio" + "context" + "encoding/json" + "fmt" + "os" + "path/filepath" + "regexp" + "sort" + "strings" +) + +const grepOutputContent = "content" + +// GlobTool 实现 glob 工具。 +type GlobTool struct { + resolver *Resolver + limits Limits +} + +// NewGlobTool 创建 glob 工具。 +func NewGlobTool(resolver *Resolver, limits Limits) *GlobTool { + return &GlobTool{resolver: resolver, limits: limits} +} + +// Name 返回工具名。 +func (t *GlobTool) Name() string { return "glob" } + +// Execute 执行 glob。 +func (t *GlobTool) Execute(ctx context.Context, input json.RawMessage) Result { + var req struct { + Pattern string `json:"pattern"` + Path string `json:"path,omitempty"` + } + if err := decodeInput(input, &req); err != nil { + return ErrorResult(err.Error()) + } + if req.Pattern == "" { + return ErrorResult("pattern must not be empty") + } + root := t.resolver.Root() + if req.Path != "" { + resolved, err := t.resolver.ResolveExisting(req.Path) + if err != nil { + return ErrorResult(err.Error()) + } + root = resolved + } + re, err := globRegexp(req.Pattern) + if err != nil { + return ErrorResult(err.Error()) + } + type match struct { + path string + mod int64 + } + var matches []match + err = filepath.WalkDir(root, func(p string, d os.DirEntry, err error) error { + if err != nil { + return nil + } + select { + case <-ctx.Done(): + return ctx.Err() + default: + } + if d.IsDir() { + return nil + } + rel, err := filepath.Rel(t.resolver.Root(), p) + if err != nil { + return nil + } + rel = filepath.ToSlash(rel) + name := filepath.ToSlash(d.Name()) + if !re.MatchString(rel) && !re.MatchString(name) { + return nil + } + info, _ := d.Info() + mod := int64(0) + if info != nil { + mod = info.ModTime().UnixNano() + } + matches = append(matches, match{path: rel, mod: mod}) + if t.limits.GlobMaxMatches > 0 && len(matches) >= t.limits.GlobMaxMatches { + return filepath.SkipAll + } + return nil + }) + if err != nil && err != filepath.SkipAll { + return ErrorResult(err.Error()) + } + sort.Slice(matches, func(i, j int) bool { + if matches[i].mod == matches[j].mod { + return matches[i].path < matches[j].path + } + return matches[i].mod > matches[j].mod + }) + var b strings.Builder + for _, m := range matches { + b.WriteString(m.path) + b.WriteByte('\n') + } + if t.limits.GlobMaxMatches > 0 && len(matches) >= t.limits.GlobMaxMatches { + fmt.Fprintf(&b, "\n[truncated: first %d matches]\n", t.limits.GlobMaxMatches) + } + return TextResult(b.String()) +} + +// GrepTool 实现 grep 工具。 +type GrepTool struct { + resolver *Resolver + limits Limits +} + +// NewGrepTool 创建 grep 工具。 +func NewGrepTool(resolver *Resolver, limits Limits) *GrepTool { + return &GrepTool{resolver: resolver, limits: limits} +} + +// Name 返回工具名。 +func (t *GrepTool) Name() string { return "grep" } + +// Execute 执行 grep。 +func (t *GrepTool) Execute(ctx context.Context, input json.RawMessage) Result { + var req struct { + Pattern string `json:"pattern"` + Path string `json:"path,omitempty"` + GlobFilter string `json:"glob_filter,omitempty"` + CaseInsensitive bool `json:"case_insensitive,omitempty"` + OutputMode string `json:"output_mode,omitempty"` + } + if err := decodeInput(input, &req); err != nil { + return ErrorResult(err.Error()) + } + if req.Pattern == "" { + return ErrorResult("pattern must not be empty") + } + pattern := req.Pattern + if req.CaseInsensitive { + pattern = "(?i)" + pattern + } + re, err := regexp.Compile(pattern) + if err != nil { + return ErrorResult(err.Error()) + } + var filter *regexp.Regexp + if req.GlobFilter != "" { + filter, err = globRegexp(req.GlobFilter) + if err != nil { + return ErrorResult(err.Error()) + } + } + root := t.resolver.Root() + if req.Path != "" { + root, err = t.resolver.ResolveExisting(req.Path) + if err != nil { + return ErrorResult(err.Error()) + } + } + mode := req.OutputMode + if mode == "" { + mode = grepOutputContent + } + var b strings.Builder + count := 0 + outputTruncated := false + filesWithMatches := map[string]bool{} + err = filepath.WalkDir(root, func(p string, d os.DirEntry, err error) error { + if err != nil { + return nil + } + select { + case <-ctx.Done(): + return ctx.Err() + default: + } + if d.IsDir() { + return nil + } + rel, err := filepath.Rel(t.resolver.Root(), p) + if err != nil { + return nil + } + rel = filepath.ToSlash(rel) + if filter != nil && !filter.MatchString(rel) && !filter.MatchString(filepath.ToSlash(d.Name())) { + return nil + } + host, err := t.resolver.ResolveExisting(rel) + if err != nil { + return nil + } + remaining := -1 + if t.limits.GrepMaxMatches > 0 { + remaining = t.limits.GrepMaxMatches - count + } + fileCount, truncated, err := grepFile(ctx, host, rel, re, mode, &b, remaining, t.limits.MaxOutputBytes) + if err != nil { + return err + } + if truncated { + outputTruncated = true + } + if fileCount > 0 { + filesWithMatches[rel] = true + count += fileCount + } + if outputTruncated || t.limits.GrepMaxMatches > 0 && count >= t.limits.GrepMaxMatches { + return filepath.SkipAll + } + return nil + }) + if err != nil && err != filepath.SkipAll { + return ErrorResult(err.Error()) + } + if mode == "files_with_matches" { + files := make([]string, 0, len(filesWithMatches)) + for f := range filesWithMatches { + files = append(files, f) + } + sort.Strings(files) + b.Reset() + for _, f := range files { + b.WriteString(f) + b.WriteByte('\n') + } + } + if mode == "count" { + b.Reset() + fmt.Fprintf(&b, "%d\n", count) + } + if t.limits.GrepMaxMatches > 0 && count >= t.limits.GrepMaxMatches && mode == grepOutputContent { + fmt.Fprintf(&b, "\n[truncated: first %d matches]\n", t.limits.GrepMaxMatches) + } + if outputTruncated { + return TextResult(outputWithTruncationMarker(b.String(), t.limits.MaxOutputBytes)) + } + return TextResult(truncateOutput(b.String(), t.limits.MaxOutputBytes)) +} + +func grepFile(ctx context.Context, path, rel string, re *regexp.Regexp, mode string, b *strings.Builder, remaining int, maxOutputBytes int64) (int, bool, error) { + f, err := os.Open(path) + if err != nil { + return 0, false, err + } + defer func() { _ = f.Close() }() + scanner := bufio.NewScanner(f) + scanner.Buffer(make([]byte, 4096), 1<<20) + lineNo := 0 + count := 0 + for scanner.Scan() { + select { + case <-ctx.Done(): + return count, false, ctx.Err() + default: + } + lineNo++ + line := scanner.Text() + if !re.MatchString(line) { + continue + } + count++ + if mode == grepOutputContent && (remaining < 0 || count <= remaining) { + entry := fmt.Sprintf("%s:%d:%s\n", rel, lineNo, line) + if !appendWithinLimit(b, entry, maxOutputBytes) { + return count, true, nil + } + } + if remaining >= 0 && count >= remaining { + break + } + } + if err := scanner.Err(); err != nil { + return count, false, fmt.Errorf("scan %s: %w", rel, err) + } + return count, false, nil +} + +func appendWithinLimit(builder *strings.Builder, text string, limit int64) bool { + if limit <= 0 { + builder.WriteString(text) + return true + } + remaining := limit - int64(builder.Len()) + if remaining <= 0 { + return false + } + if int64(len(text)) <= remaining { + builder.WriteString(text) + return true + } + builder.WriteString(text[:remaining]) + return false +} + +func outputWithTruncationMarker(output string, limit int64) string { + const marker = "\n[output truncated]\n" + if limit <= 0 { + return output + marker + } + if limit <= int64(len(marker)) { + return marker[:limit] + } + keep := limit - int64(len(marker)) + if int64(len(output)) > keep { + output = output[:keep] + } + return output + marker +} + +func globRegexp(pattern string) (*regexp.Regexp, error) { + var b strings.Builder + b.WriteString("^") + for i := 0; i < len(pattern); i++ { + ch := pattern[i] + switch ch { + case '*': + if i+1 < len(pattern) && pattern[i+1] == '*' { + b.WriteString(".*") + i++ + } else { + b.WriteString("[^/]*") + } + case '?': + b.WriteString("[^/]") + case '.', '+', '(', ')', '|', '^', '$', '{', '}', '[', ']', '\\': + b.WriteByte('\\') + b.WriteByte(ch) + default: + b.WriteByte(ch) + } + } + b.WriteString("$") + return regexp.Compile(b.String()) +} diff --git a/arkruntime/toolset/toolset_test.go b/arkruntime/toolset/toolset_test.go new file mode 100644 index 0000000..60e9bf8 --- /dev/null +++ b/arkruntime/toolset/toolset_test.go @@ -0,0 +1,410 @@ +// Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +// SPDX-License-Identifier: Apache-2.0 +package toolset + +import ( + "context" + "os" + "path/filepath" + "strings" + "testing" +) + +func TestFileToolsStayInsideWorkdir(t *testing.T) { + root := t.TempDir() + outside := filepath.Join(t.TempDir(), "outside.txt") + if err := os.WriteFile(outside, []byte("secret"), 0o600); err != nil { + t.Fatal(err) + } + if err := os.Symlink(outside, filepath.Join(root, "link")); err != nil { + t.Fatal(err) + } + set, err := NewDefault(Options{Workdir: root}) + if err != nil { + t.Fatal(err) + } + defer set.Close() + + res := set.Execute(context.Background(), "read", []byte(`{"file_path":"link"}`)) + if !res.IsError || !strings.Contains(res.Content[0].Text, "path escapes workdir") { + t.Fatalf("read result = %+v", res) + } +} + +func TestFileToolsAllowUnrestrictedPaths(t *testing.T) { + root := t.TempDir() + outside := filepath.Join(t.TempDir(), "outside.txt") + if err := os.WriteFile(outside, []byte("visible"), 0o600); err != nil { + t.Fatal(err) + } + if err := os.Symlink(outside, filepath.Join(root, "link")); err != nil { + t.Fatal(err) + } + set, err := NewDefault(Options{Workdir: root, UnrestrictedPaths: true}) + if err != nil { + t.Fatal(err) + } + defer set.Close() + + res := set.Execute(context.Background(), "read", []byte(`{"file_path":"link"}`)) + if res.IsError || !strings.Contains(res.Content[0].Text, "visible") { + t.Fatalf("read result = %+v", res) + } +} + +func TestWriteRejectsDanglingIntermediateSymlink(t *testing.T) { + root := t.TempDir() + outside := filepath.Join(t.TempDir(), "missing") + if err := os.Symlink(outside, filepath.Join(root, "link")); err != nil { + t.Fatal(err) + } + set, err := NewDefault(Options{Workdir: root}) + if err != nil { + t.Fatal(err) + } + defer set.Close() + + res := set.Execute(context.Background(), "write", []byte(`{"file_path":"link/new.txt","content":"escape"}`)) + if !res.IsError || !strings.Contains(res.Content[0].Text, "path escapes workdir") { + t.Fatalf("write result = %+v", res) + } + if _, err := os.Stat(filepath.Join(outside, "new.txt")); !os.IsNotExist(err) { + t.Fatalf("write escaped workdir: %v", err) + } +} + +func TestWriteRejectsMultiLayerDanglingSymlink(t *testing.T) { + root := t.TempDir() + outside := filepath.Join(t.TempDir(), "missing") + if err := os.Symlink(outside, filepath.Join(root, "second")); err != nil { + t.Fatal(err) + } + if err := os.Symlink("second", filepath.Join(root, "first")); err != nil { + t.Fatal(err) + } + set, err := NewDefault(Options{Workdir: root}) + if err != nil { + t.Fatal(err) + } + defer set.Close() + + res := set.Execute(context.Background(), "write", []byte(`{"file_path":"first/new.txt","content":"escape"}`)) + if !res.IsError || !strings.Contains(res.Content[0].Text, "path escapes workdir") { + t.Fatalf("write result = %+v", res) + } +} + +func TestWriteAllowsOrdinaryMissingDirectories(t *testing.T) { + root := t.TempDir() + set, err := NewDefault(Options{Workdir: root}) + if err != nil { + t.Fatal(err) + } + defer set.Close() + + res := set.Execute(context.Background(), "write", []byte(`{"file_path":"new/dir/file.txt","content":"ok"}`)) + if res.IsError { + t.Fatalf("write result = %+v", res) + } + data, err := os.ReadFile(filepath.Join(root, "new", "dir", "file.txt")) + if err != nil || string(data) != "ok" { + t.Fatalf("written data=%q err=%v", data, err) + } +} + +func TestWriteAllowsResolvedSymlinkInsideWorkdir(t *testing.T) { + root := t.TempDir() + actual := filepath.Join(root, "actual") + if err := os.MkdirAll(actual, 0o755); err != nil { + t.Fatal(err) + } + if err := os.Symlink("actual", filepath.Join(root, "link")); err != nil { + t.Fatal(err) + } + set, err := NewDefault(Options{Workdir: root}) + if err != nil { + t.Fatal(err) + } + defer set.Close() + + res := set.Execute(context.Background(), "write", []byte(`{"file_path":"link/file.txt","content":"ok"}`)) + if res.IsError { + t.Fatalf("write result = %+v", res) + } + data, err := os.ReadFile(filepath.Join(actual, "file.txt")) + if err != nil || string(data) != "ok" { + t.Fatalf("written data=%q err=%v", data, err) + } +} + +func TestEditAtomicallyPreservesFileMode(t *testing.T) { + root := t.TempDir() + path := filepath.Join(root, "script.sh") + if err := os.WriteFile(path, []byte("echo old\n"), 0o750); err != nil { + t.Fatal(err) + } + before, err := os.Stat(path) + if err != nil { + t.Fatal(err) + } + set, err := NewDefault(Options{Workdir: root}) + if err != nil { + t.Fatal(err) + } + defer set.Close() + + res := set.Execute(context.Background(), "edit", []byte(`{"file_path":"script.sh","old_string":"old","new_string":"new"}`)) + if res.IsError { + t.Fatalf("edit result = %+v", res) + } + data, err := os.ReadFile(path) + if err != nil || string(data) != "echo new\n" { + t.Fatalf("edited data=%q err=%v", data, err) + } + info, err := os.Stat(path) + if err != nil { + t.Fatal(err) + } + if info.Mode().Perm() != before.Mode().Perm() { + t.Fatalf("edited mode=%o want=%o", info.Mode().Perm(), before.Mode().Perm()) + } +} + +func TestReadSupportsViewRange(t *testing.T) { + root := t.TempDir() + if err := os.WriteFile(filepath.Join(root, "demo.txt"), []byte("a\nb\nc\nd\n"), 0o600); err != nil { + t.Fatal(err) + } + set, err := NewDefault(Options{Workdir: root}) + if err != nil { + t.Fatal(err) + } + defer set.Close() + + res := set.Execute(context.Background(), "read", []byte(`{"file_path":"demo.txt","view_range":[2,3]}`)) + if res.IsError || res.Content[0].Text != "b\nc" { + t.Fatalf("read result = %+v", res) + } +} + +func TestReadDefaultsToRawContent(t *testing.T) { + root := t.TempDir() + if err := os.WriteFile(filepath.Join(root, "demo.txt"), []byte("a\nb\n"), 0o600); err != nil { + t.Fatal(err) + } + set, err := NewDefault(Options{Workdir: root}) + if err != nil { + t.Fatal(err) + } + defer set.Close() + + res := set.Execute(context.Background(), "read", []byte(`{"file_path":"demo.txt"}`)) + if res.IsError || res.Content[0].Text != "a\nb\n" { + t.Fatalf("read result = %+v", res) + } +} + +func TestGrepDoesNotReadSymlinkOutsideWorkdir(t *testing.T) { + root := t.TempDir() + outside := filepath.Join(t.TempDir(), "secret.txt") + if err := os.WriteFile(outside, []byte("outside-secret"), 0o600); err != nil { + t.Fatal(err) + } + if err := os.Symlink(outside, filepath.Join(root, "link.txt")); err != nil { + t.Fatal(err) + } + set, err := NewDefault(Options{Workdir: root}) + if err != nil { + t.Fatal(err) + } + defer set.Close() + + res := set.Execute(context.Background(), "grep", []byte(`{"pattern":"outside-secret"}`)) + if res.IsError { + t.Fatalf("grep result = %+v", res) + } + if strings.Contains(res.Content[0].Text, "outside-secret") || strings.Contains(res.Content[0].Text, "link.txt") { + t.Fatalf("grep leaked symlink target: %+v", res) + } +} + +func TestGrepReportsScannerError(t *testing.T) { + root := t.TempDir() + if err := os.WriteFile(filepath.Join(root, "long.txt"), []byte(strings.Repeat("x", 1<<20+1)), 0o600); err != nil { + t.Fatal(err) + } + set, err := NewDefault(Options{Workdir: root}) + if err != nil { + t.Fatal(err) + } + defer set.Close() + + res := set.Execute(context.Background(), "grep", []byte(`{"pattern":"x"}`)) + if !res.IsError || !strings.Contains(res.Content[0].Text, "scan long.txt") { + t.Fatalf("grep result = %+v", res) + } +} + +func TestBashPersistsStateAndSurvivesExit(t *testing.T) { + root := t.TempDir() + set, err := NewDefault(Options{Workdir: root}) + if err != nil { + t.Fatal(err) + } + defer set.Close() + + res := set.Execute(context.Background(), "bash", []byte(`{"command":"export SELF_HOST_TEST=ok; mkdir -p sub; cd sub"}`)) + if res.IsError { + t.Fatalf("bash setup = %+v", res) + } + res = set.Execute(context.Background(), "bash", []byte(`{"command":"pwd; echo $SELF_HOST_TEST"}`)) + if res.IsError || !strings.Contains(res.Content[0].Text, "/sub") || !strings.Contains(res.Content[0].Text, "ok") { + t.Fatalf("bash persisted = %+v", res) + } + res = set.Execute(context.Background(), "bash", []byte(`{"command":"exit 7"}`)) + if res.IsError || !strings.Contains(res.Content[0].Text, "exit_code: 7") { + t.Fatalf("bash exit = %+v", res) + } + res = set.Execute(context.Background(), "bash", []byte(`{"command":"echo alive"}`)) + if res.IsError || !strings.Contains(res.Content[0].Text, "alive") { + t.Fatalf("bash alive = %+v", res) + } +} + +func TestBashUsesPrivateTemporaryDirectory(t *testing.T) { + root := t.TempDir() + session, err := NewBashSession(Options{Workdir: root}) + if err != nil { + t.Fatal(err) + } + privateDir := session.privateDir + if withinRoot(root, privateDir) { + t.Fatalf("private directory is inside workdir: %s", privateDir) + } + info, err := os.Stat(privateDir) + if err != nil { + t.Fatal(err) + } + if !info.IsDir() || info.Mode().Perm() != 0o700 { + t.Fatalf("private directory mode=%s", info.Mode()) + } + if err := session.Close(); err != nil { + t.Fatal(err) + } + if _, err := os.Stat(privateDir); !os.IsNotExist(err) { + t.Fatalf("private directory was not removed: %v", err) + } +} + +func TestBashIgnoresInBandSentinelSpoof(t *testing.T) { + root := t.TempDir() + set, err := NewDefault(Options{Workdir: root}) + if err != nil { + t.Fatal(err) + } + defer set.Close() + + res := set.Execute(context.Background(), "bash", []byte(`{"command":"n=${BASH_SOURCE[0]##*/}; n=${n%.sh}; printf '\\n__MA_WORKER_EXIT_%s:0\\n' \"$n\"; printf '__MA_WORKER_DONE_%s:0\\n' \"$n\"; sleep 0.1; echo after-spoof"}`)) + if res.IsError { + t.Fatalf("bash spoof = %+v", res) + } + text := res.Content[0].Text + if !strings.Contains(text, "after-spoof") { + t.Fatalf("bash returned before real completion: %q", text) + } +} + +func TestBashScrubsSensitiveExtraEnv(t *testing.T) { + t.Setenv("MA_ENVIRONMENT_KEY", "host-leak") + t.Setenv("MA_WORKER_KEY", "host-leak") + t.Setenv("ARK_API_KEY", "host-leak") + t.Setenv("VOLC_SECRETKEY", "host-leak") + t.Setenv("GITHUB_TOKEN", "host-leak") + t.Setenv("INHERITED_SAFE_VAR", "host-value") + root := t.TempDir() + set, err := NewDefault(Options{ + Workdir: root, + Env: map[string]string{ + "MA_ENVIRONMENT_KEY": "extra-leak", + "MA_WORKER_KEY": "extra-leak", + "ARK_API_KEY": "extra-leak", + "VOLC_SECRETKEY": "extra-leak", + "GITHUB_TOKEN": "extra-leak", + "SAFE_VAR": "ok", + }, + }) + if err != nil { + t.Fatal(err) + } + defer set.Close() + + res := set.Execute(context.Background(), "bash", []byte(`{"command":"printf '%s/%s/%s/%s/%s/%s/%s' \"${MA_ENVIRONMENT_KEY-unset}\" \"${MA_WORKER_KEY-unset}\" \"${ARK_API_KEY-unset}\" \"${VOLC_SECRETKEY-unset}\" \"${GITHUB_TOKEN-unset}\" \"${INHERITED_SAFE_VAR-unset}\" \"$SAFE_VAR\""}`)) + if res.IsError { + t.Fatalf("bash env = %+v", res) + } + if !strings.Contains(res.Content[0].Text, "unset/unset/unset/unset/unset/unset/ok") { + t.Fatalf("bash env leaked sensitive values: %+v", res) + } +} + +func TestBashScrubsInheritedCloudCredentials(t *testing.T) { + t.Setenv("VOLC_ACCESSKEY", "host-leak") + t.Setenv("BYTEPLUS_SECRETKEY", "host-leak") + t.Setenv("AWS_SECRET_ACCESS_KEY", "host-leak") + t.Setenv("SAFE_INHERITED_VAR", "ok") + set, err := NewDefault(Options{Workdir: t.TempDir()}) + if err != nil { + t.Fatal(err) + } + defer set.Close() + + res := set.Execute(context.Background(), "bash", []byte(`{"command":"printf '%s/%s/%s/%s' \"${VOLC_ACCESSKEY-unset}\" \"${BYTEPLUS_SECRETKEY-unset}\" \"${AWS_SECRET_ACCESS_KEY-unset}\" \"$SAFE_INHERITED_VAR\""}`)) + if res.IsError || !strings.Contains(res.Content[0].Text, "unset/unset/unset/ok") { + t.Fatalf("bash inherited env = %+v", res) + } +} + +func TestGrepHonorsOutputLimit(t *testing.T) { + root := t.TempDir() + content := strings.Repeat("match-"+strings.Repeat("x", 256)+"\n", 20) + if err := os.WriteFile(filepath.Join(root, "large.txt"), []byte(content), 0o600); err != nil { + t.Fatal(err) + } + limits := DefaultLimits() + limits.MaxOutputBytes = 128 + set, err := NewDefault(Options{Workdir: root, Limits: limits}) + if err != nil { + t.Fatal(err) + } + defer set.Close() + + res := set.Execute(context.Background(), "grep", []byte(`{"pattern":"match"}`)) + if res.IsError { + t.Fatalf("grep result = %+v", res) + } + if got := len(res.Content[0].Text); got > int(limits.MaxOutputBytes) { + t.Fatalf("grep output bytes = %d, limit = %d", got, limits.MaxOutputBytes) + } + if !strings.Contains(res.Content[0].Text, "[output truncated]") { + t.Fatalf("grep output missing truncation marker: %q", res.Content[0].Text) + } +} + +func TestGrepHonorsCanceledContext(t *testing.T) { + root := t.TempDir() + if err := os.WriteFile(filepath.Join(root, "demo.txt"), []byte("match\n"), 0o600); err != nil { + t.Fatal(err) + } + set, err := NewDefault(Options{Workdir: root}) + if err != nil { + t.Fatal(err) + } + defer set.Close() + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + res := set.Execute(ctx, "grep", []byte(`{"pattern":"match"}`)) + if !res.IsError || !strings.Contains(res.Content[0].Text, context.Canceled.Error()) { + t.Fatalf("grep result = %+v", res) + } +} diff --git a/arkruntime/toolset/types.go b/arkruntime/toolset/types.go new file mode 100644 index 0000000..8d6326f --- /dev/null +++ b/arkruntime/toolset/types.go @@ -0,0 +1,167 @@ +// Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +// SPDX-License-Identifier: Apache-2.0 +// Package toolset 实现 self-hosted worker 内置工具集合。 +package toolset + +import ( + "context" + "encoding/json" + "fmt" + "sync" + "time" +) + +// ContentBlock 是工具结果内容块。 +type ContentBlock struct { + Type string `json:"type"` + Text string `json:"text,omitempty"` + MediaType string `json:"media_type,omitempty"` + Data []byte `json:"data,omitempty"` +} + +// Result 是一次工具执行结果。 +type Result struct { + Content []ContentBlock `json:"content"` + IsError bool `json:"is_error"` +} + +// Tool 是 worker 可执行的本地工具。 +type Tool interface { + Name() string + // Execute 必须监听 ctx.Done;worker 只能取消 context,不能安全终止任意 Go goroutine。 + Execute(ctx context.Context, input json.RawMessage) Result +} + +// ClosableTool 是需要释放资源的工具。 +type ClosableTool interface { + Tool + Close() error +} + +// Options 是默认工具集合配置。 +type Options struct { + Workdir string + UnrestrictedPaths bool + // Env 非 nil 时完全替换继承的进程环境;敏感凭证字段始终会被移除。 + Env map[string]string + Limits Limits + BashPath string + ToolTimeout time.Duration +} + +// Limits 是工具执行的尺寸与数量上限。 +type Limits struct { + MaxOutputBytes int64 + MaxInputFileBytes int64 + ReadDefaultLines int + ReadMaxLineChars int + GlobMaxMatches int + GrepMaxMatches int +} + +// DefaultLimits 返回默认工具限制。 +func DefaultLimits() Limits { + return Limits{ + MaxOutputBytes: 100 << 10, + MaxInputFileBytes: 256 << 10, + ReadDefaultLines: 2000, + ReadMaxLineChars: 2000, + GlobMaxMatches: 1000, + GrepMaxMatches: 200, + } +} + +// Set 是可执行工具集合。 +type Set struct { + mu sync.RWMutex + tools map[string]Tool +} + +// NewDefault 创建 bash/read/write/edit/glob/grep 默认工具集合。 +func NewDefault(opts Options) (*Set, error) { + if opts.Limits == (Limits{}) { + opts.Limits = DefaultLimits() + } + if opts.ToolTimeout <= 0 { + opts.ToolTimeout = 120 * time.Second + } + resolver, err := NewResolverWithOptions(opts.Workdir, opts.UnrestrictedPaths) + if err != nil { + return nil, err + } + bash, err := NewBashTool(opts) + if err != nil { + return nil, err + } + s := &Set{tools: map[string]Tool{}} + s.Register(bash) + s.Register(NewReadTool(resolver, opts.Limits)) + s.Register(NewWriteTool(resolver, opts.Limits)) + s.Register(NewEditTool(resolver, opts.Limits)) + s.Register(NewGlobTool(resolver, opts.Limits)) + s.Register(NewGrepTool(resolver, opts.Limits)) + return s, nil +} + +// Register 注册一个工具。 +func (s *Set) Register(tool Tool) { + s.mu.Lock() + defer s.mu.Unlock() + s.tools[tool.Name()] = tool +} + +// Execute 执行指定工具。 +func (s *Set) Execute(ctx context.Context, name string, input json.RawMessage) Result { + s.mu.RLock() + tool := s.tools[name] + s.mu.RUnlock() + if tool == nil { + return ErrorResult(fmt.Sprintf("unknown tool: %s", name)) + } + return tool.Execute(ctx, input) +} + +// Has 判断工具集合是否注册了指定工具。 +func (s *Set) Has(name string) bool { + s.mu.RLock() + defer s.mu.RUnlock() + return s.tools[name] != nil +} + +// Close 关闭工具集合持有的资源。 +func (s *Set) Close() error { + s.mu.RLock() + defer s.mu.RUnlock() + var first error + for _, tool := range s.tools { + if c, ok := tool.(ClosableTool); ok { + if err := c.Close(); err != nil && first == nil { + first = err + } + } + } + return first +} + +// TextResult 构造文本工具结果。 +func TextResult(text string) Result { + return Result{Content: []ContentBlock{{Type: "text", Text: text}}} +} + +// ErrorResult 构造错误工具结果。 +func ErrorResult(message string) Result { + return Result{ + Content: []ContentBlock{{Type: "text", Text: message}}, + IsError: true, + } +} + +func decodeInput(input json.RawMessage, out any) error { + if len(input) == 0 { + input = []byte("{}") + } + if err := json.Unmarshal(input, out); err != nil { + return fmt.Errorf("invalid tool input: %w", err) + } + return nil +} diff --git a/examples/README.md b/examples/README.md index f900810..b3a94ed 100644 --- a/examples/README.md +++ b/examples/README.md @@ -10,7 +10,7 @@ export ARK_MODEL=... go run examples/volc/responses/basic/main.go ``` -All service-calling examples are grouped by cloud: +Cloud-specific examples are grouped by cloud: - [`volc/`](./volc) uses `NewVolcClientWithApiKey` and Volcengine China model IDs. - [`byteplus/`](./byteplus) uses `NewByteplusClientWithApiKey` and BytePlus model IDs. @@ -18,3 +18,5 @@ All service-calling examples are grouped by cloud: The paired multimodal and sparse embedding examples default to `doubao-embedding-vision-251215` / `skylark-embedding-vision-251215`. The paired image examples default to `doubao-seedream-5-0-pro-260628` / `dola-seedream-5-0-pro-260628`. The paired video-generation examples default to `doubao-seedance-2-0-fast-260128` / `dreamina-seedance-2-0-fast-260128`. MCP is available in both clouds and its examples explicitly send `ark-beta-mcp: true`. Other built-in tools are CN-only: Web Search sends `ark-beta-web-search: true`, and Doubao App sends `ark-beta-doubao-app: true`. + +The [`self_hosted_worker/`](./self_hosted_worker) example runs a local Managed Agents worker for an existing self-hosted environment. It requires `MA_ENVIRONMENT_ID`; the client defaults to `https://ark.cn-beijing.volces.com/api/v3`. diff --git a/examples/byteplus/sessions_loop/main.go b/examples/byteplus/sessions_loop/main.go index 9bc10ce..3f0228b 100644 --- a/examples/byteplus/sessions_loop/main.go +++ b/examples/byteplus/sessions_loop/main.go @@ -79,7 +79,7 @@ func main() { // 3. Session — binds the agent to the environment. sess, err := client.CreateSession(ctx, &session.CreateSessionRequest{ Agent: session.NewStringAgentIdentifier(ag.ID), - EnvironmentID: env.ID, + EnvironmentID: session.NewOptString(env.ID), Title: session.NewOptString("ark-runtime-go example loop"), }) if err != nil { diff --git a/examples/self_hosted_worker/main.go b/examples/self_hosted_worker/main.go new file mode 100644 index 0000000..b7cf04d --- /dev/null +++ b/examples/self_hosted_worker/main.go @@ -0,0 +1,52 @@ +// Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +// SPDX-License-Identifier: Apache-2.0 + +// Self-hosted worker runner. +// +// Required: +// +// export ARK_API_KEY=... +// export MA_ENVIRONMENT_ID=env_xxx +// +// Run from examples/: +// +// go run ./self_hosted_worker +package main + +import ( + "context" + "log" + "os" + "os/signal" + "syscall" + + "github.com/volcengine/ark-runtime-go/arkruntime" + "github.com/volcengine/ark-runtime-go/arkruntime/lib/environments" +) + +func main() { + apiKey := os.Getenv("ARK_API_KEY") + baseURL := os.Getenv("ARK_BASE_URL") + envID := os.Getenv("MA_ENVIRONMENT_ID") + + if apiKey == "" || envID == "" { + log.Fatal("ARK_API_KEY and MA_ENVIRONMENT_ID are required") + } + + clientOptions := make([]arkruntime.ConfigOption, 0, 1) + if baseURL != "" { + clientOptions = append(clientOptions, arkruntime.WithBaseUrl(baseURL)) + } + client := arkruntime.NewClientWithApiKey(apiKey, clientOptions...) + + worker := environments.NewEnvironmentWorkerForClient(client, environments.EnvironmentWorkerOptions{ + EnvironmentID: envID, + Workdir: ".", + }) + + ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM) + defer stop() + if err := worker.Run(ctx); err != nil { + log.Fatal(err) + } +} diff --git a/examples/volc/sessions_loop/main.go b/examples/volc/sessions_loop/main.go index 46d6573..bfcab99 100644 --- a/examples/volc/sessions_loop/main.go +++ b/examples/volc/sessions_loop/main.go @@ -79,7 +79,7 @@ func main() { // 3. Session — binds the agent to the environment. sess, err := client.CreateSession(ctx, &session.CreateSessionRequest{ Agent: session.NewStringAgentIdentifier(ag.ID), - EnvironmentID: env.ID, + EnvironmentID: session.NewOptString(env.ID), Title: session.NewOptString("ark-runtime-go example loop"), }) if err != nil {