feat: forward agent version and extension context
This commit is contained in:
@@ -0,0 +1,8 @@
|
||||
---
|
||||
category: Added
|
||||
---
|
||||
|
||||
- **Agent version and extended context passthrough** (Aone 85384225) — adds
|
||||
validated `DWS_AGENT_VER` and sensitive JSON `DWS_AGENT_EXT` metadata to
|
||||
ordinary non-plugin MCP requests without forwarding it to A2A, OAuth,
|
||||
Discovery, or third-party plugins.
|
||||
@@ -19,3 +19,12 @@
|
||||
|
||||
# Cache directory (optional, defaults to ~/.dws/cache)
|
||||
# DWS_CACHE_DIR=
|
||||
|
||||
# Agent integration metadata (optional; ordinary non-plugin MCP requests only)
|
||||
# DWS_AGENT_PRODUCT=example-agent
|
||||
# DWS_AGENT_HOST=cloud
|
||||
# DWS_AGENT_VER=0.1.5
|
||||
# DWS_AGENT_EXT='{"umt":"example-redacted","miniwua":"example-redacted","ua":"ExampleAgent/0.1.5"}'
|
||||
# The outer single quotes above are shell syntax and are not part of the value.
|
||||
# DWS_AGENT_EXT is sensitive caller-declared JSON (max 8 KiB); never put real
|
||||
# tokens in committed files or use this metadata alone for authentication.
|
||||
|
||||
@@ -7,6 +7,8 @@
|
||||
| `DWS_CONFIG_DIR` | Override default config directory / 覆盖默认配置目录 |
|
||||
| `DWS_AGENT_PRODUCT` | Optional, caller-declared Agent product sent as `x-dws-agent-product` (for example `qwenwork`) for downstream logs/BI and used as the IM `clawType` display label when `--ai-tag` is enabled. `--ai-tag` defaults to `true`, so a configured Product changes the displayed label by default. With `--ai-tag=false`, native `chat message send` / `reply` calls send an empty `clawType`, while shortcut calls omit the argument. Surrounding ASCII spaces/tabs are trimmed; the remaining value must be at most 64 bytes and match `^[A-Za-z0-9][A-Za-z0-9_-]*$`. Unset or empty values omit the Header and use the edition's IM display default. This client never uses Product to change the separate HTTP `claw-type` PAT/routing label. / 可选、由调用方声明的 Agent 产品标识,经校验后作为 `x-dws-agent-product` 发送,并用于 IM 小尾巴;`--ai-tag` 默认为 `true`,因此配置 Product 后默认会改变展示标签。使用 `--ai-tag=false` 时,原生 `chat message send` / `reply` 发送空的 `clawType`,shortcut 调用则省略该参数。未设置时省略请求头且 IM 使用发行版默认值;本客户端不会用 Product 修改独立的 HTTP `claw-type` |
|
||||
| `DWS_AGENT_HOST` | Optional, caller-declared Agent runtime form sent as `x-dws-agent-host` (for example `cloud` or `desktop`) for downstream logs/BI. Surrounding ASCII spaces/tabs are trimmed; the remaining value must be at most 64 bytes and match `^[a-z0-9][a-z0-9_-]*$`; unset values are omitted. This client does not use Host for PAT, authentication, Discovery, or MCP endpoint selection. / 可选、由调用方声明的 Agent 运行形态,经校验后作为 `x-dws-agent-host` 发送给下游日志/BI;本客户端不使用该值进行 PAT、鉴权、Discovery 或 MCP 端点选择,未设置时省略 |
|
||||
| `DWS_AGENT_VER` | Optional caller-declared Agent version / 可选、由调用方声明的 Agent 版本。After trimming surrounding ASCII spaces/tabs, the value must be at most 64 bytes and match `^[A-Za-z0-9][A-Za-z0-9._+-]*$`; a non-empty valid value is sent as `x-dws-agent-ver`, while unset or empty values omit the Header. / 去除首尾 ASCII 空格和 Tab 后,值不得超过 64 字节且必须匹配上述格式;合法非空值通过 `x-dws-agent-ver` 发送,未设置或空值则省略请求头 |
|
||||
| `DWS_AGENT_EXT` | Optional caller-declared Agent extended context / 可选、由调用方声明的 Agent 扩展上下文。The value must be a UTF-8 JSON object no larger than 8 KiB, is compacted before being sent as the sensitive `x-dws-agent-ext` Header, and may use the recommended keys `umt`, `miniwua`, and `ua`; unknown keys remain supported. Unset or empty values omit the Header. / 值必须是 UTF-8 JSON 对象且不得超过 8 KiB,压缩后通过敏感请求头 `x-dws-agent-ext` 发送;推荐使用 `umt`、`miniwua`、`ua`,同时允许未知扩展键。未设置或空值则省略请求头 |
|
||||
| `DWS_<PRODUCT>_MCP_URL` | Override a product MCP endpoint for local development / 本地开发时覆盖指定产品 MCP endpoint |
|
||||
| `DWS_CLIENT_ID` | OAuth client ID (DingTalk AppKey) |
|
||||
| `DWS_CLIENT_SECRET` | OAuth client secret (DingTalk AppSecret) |
|
||||
@@ -14,6 +16,47 @@
|
||||
| `DWS_ALLOW_HTTP_ENDPOINTS` | Set `1` to allow HTTP for loopback during dev / 设为 `1` 允许回环地址 HTTP,仅用于开发调试 |
|
||||
| `DWS_DISABLE_KEYCHAIN` | macOS only. Set `1` to skip system Keychain for the encryption key and use file-based storage (same scheme as Linux). For sandboxed runtimes (e.g. Codex App) that block Keychain APIs. Weakens at-rest protection — DEK and ciphertext live in the same directory. / 仅 macOS。设为 `1` 时跳过系统 Keychain,密钥以文件形式存储(与 Linux 一致)。用于 Keychain API 被拦截的沙盒环境(如 Codex App)。代价是 DEK 与密文同目录,保护强度低于默认方案 |
|
||||
|
||||
### Agent Version and Extended Context / Agent 版本与扩展上下文
|
||||
|
||||
`DWS_AGENT_VER` and `DWS_AGENT_EXT` are sent only on the CLI's ordinary,
|
||||
non-plugin MCP requests. They do not change the standard HTTP `User-Agent` or
|
||||
the separate `X-Cli-Version` that identifies the DWS CLI version, and they are
|
||||
not forwarded to A2A, OAuth, Discovery, or third-party plugin requests.
|
||||
|
||||
`DWS_AGENT_EXT` is one JSON-object Header rather than a set of Headers. The
|
||||
recommended keys are `umt`, `miniwua`, and `ua`, but the open-source CLI keeps
|
||||
the object extensible and does not enforce a key allowlist. For example, using
|
||||
fictional, redacted values:
|
||||
|
||||
```bash
|
||||
DWS_AGENT_VER=0.1.5
|
||||
DWS_AGENT_EXT='{"umt":"example-redacted","miniwua":"example-redacted","ua":"ExampleAgent/0.1.5"}'
|
||||
```
|
||||
|
||||
The shell's outer single quotes group the JSON and are not part of the
|
||||
environment-variable value. The CLI trims surrounding ASCII spaces/tabs,
|
||||
omits either Header when its value is empty, and compacts EXT to a single-line
|
||||
JSON object. A representative current payload is about 657 bytes, well below
|
||||
the 8 KiB limit; integrations must still enforce the limit because values can
|
||||
grow. EXT may contain sensitive device or runtime signals: the CLI masks it in
|
||||
configuration and logs, and removes it on a cross-host redirect.
|
||||
|
||||
Both values are declared by the caller and are therefore forgeable. They can
|
||||
support compatibility checks, diagnostics, and observability, but they are not
|
||||
credentials or attestations and must never be sufficient on their own to
|
||||
authenticate a caller or authorize access.
|
||||
|
||||
`DWS_AGENT_VER` 与 `DWS_AGENT_EXT` 仅随 CLI 发起的普通非插件 MCP 请求发送,不会
|
||||
改变标准 HTTP `User-Agent`,也不会覆盖标识 DWS CLI 自身版本的 `X-Cli-Version`;
|
||||
二者不会进入 A2A、OAuth、Discovery 或第三方插件请求。EXT 使用单个 JSON 对象
|
||||
请求头,不拆成多个子请求头;推荐键为 `umt`、`miniwua`、`ua`,但开源 CLI 不限制
|
||||
扩展键。Shell 示例中的外层单引号只用于保护 JSON,不属于环境变量值。当前典型负载
|
||||
约为 657 字节,远低于 8 KiB 上限,但集成方仍须遵守大小限制。EXT 可能包含敏感的
|
||||
设备或运行时信号,配置展示和日志会对其脱敏,跨主机重定向时也会移除该请求头。
|
||||
|
||||
这两个值都由调用方自行声明,可以被伪造;它们可用于兼容性判断、诊断和可观测性,
|
||||
但不是凭据或可信证明,不能单独用于身份认证或访问授权。
|
||||
|
||||
### Agent Product, Host, and `claw-type` / Agent 产品、运行形态与 `claw-type`
|
||||
|
||||
`DWS_AGENT_PRODUCT` and `DWS_AGENT_HOST` are caller-declared observation
|
||||
|
||||
@@ -165,7 +165,7 @@ func TestResolveIdentityHeadersOmitsAbsentOrInvalidAgentHost(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestRootRejectsInvalidAgentHostBeforeEditionHook(t *testing.T) {
|
||||
func TestCrossPlatformCoverageRootRejectsInvalidAgentHostBeforeEditionHook(t *testing.T) {
|
||||
t.Setenv("DWS_CONFIG_DIR", t.TempDir())
|
||||
const invalidValue = "DO_NOT_ECHO"
|
||||
t.Setenv(envDWSAgentHost, invalidValue)
|
||||
|
||||
@@ -0,0 +1,248 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package app
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"os"
|
||||
"regexp"
|
||||
"strings"
|
||||
"unicode"
|
||||
"unicode/utf8"
|
||||
|
||||
apperrors "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/errors"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/executor"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/transport"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/configmeta"
|
||||
)
|
||||
|
||||
const (
|
||||
envDWSAgentVersion = "DWS_AGENT_VER"
|
||||
envDWSAgentExt = "DWS_AGENT_EXT"
|
||||
maxAgentVersionBytes = 64
|
||||
maxAgentExtensionBytes = 8 * 1024
|
||||
)
|
||||
|
||||
var agentVersionPattern = regexp.MustCompile(`^[A-Za-z0-9][A-Za-z0-9._+-]*$`)
|
||||
|
||||
type agentMetadataSnapshot struct {
|
||||
version string
|
||||
ext string
|
||||
versionErr error
|
||||
extErr error
|
||||
}
|
||||
|
||||
type agentMetadataSnapshotContextKey struct{}
|
||||
|
||||
func (snapshot agentMetadataSnapshot) validationError() error {
|
||||
if snapshot.versionErr != nil {
|
||||
return snapshot.versionErr
|
||||
}
|
||||
return snapshot.extErr
|
||||
}
|
||||
|
||||
func contextWithAgentMetadataSnapshot(ctx context.Context, snapshot agentMetadataSnapshot) context.Context {
|
||||
return context.WithValue(ctx, agentMetadataSnapshotContextKey{}, snapshot)
|
||||
}
|
||||
|
||||
func agentMetadataSnapshotFromContext(ctx context.Context) (agentMetadataSnapshot, bool) {
|
||||
if ctx == nil {
|
||||
return agentMetadataSnapshot{}, false
|
||||
}
|
||||
snapshot, ok := ctx.Value(agentMetadataSnapshotContextKey{}).(agentMetadataSnapshot)
|
||||
return snapshot, ok
|
||||
}
|
||||
|
||||
func init() {
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: envDWSAgentVersion,
|
||||
Category: configmeta.CategoryExternal,
|
||||
Description: "调用 DWS 的 Agent 版本;仅作为 x-dws-agent-ver 透传到非插件 MCP 请求",
|
||||
Example: "1.2.3-beta.1+build.7",
|
||||
})
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: envDWSAgentExt,
|
||||
Category: configmeta.CategoryExternal,
|
||||
Description: "调用 DWS 的 Agent 扩展上下文 JSON;仅作为 x-dws-agent-ext 透传到非插件 MCP 请求",
|
||||
Example: `{"umt":"<token>","miniwua":"<token>","ua":"agent/1.0"}`,
|
||||
Sensitive: true,
|
||||
})
|
||||
}
|
||||
|
||||
// parseAgentVersion normalizes and validates the caller-declared Agent
|
||||
// version. Only surrounding ASCII spaces and tabs are trimmed. An unset or
|
||||
// ASCII-whitespace-only value means "do not emit".
|
||||
func parseAgentVersion(raw string) (string, error) {
|
||||
value := strings.Trim(raw, " \t")
|
||||
if value == "" {
|
||||
return "", nil
|
||||
}
|
||||
if len(value) > maxAgentVersionBytes || !agentVersionPattern.MatchString(value) {
|
||||
return "", invalidAgentVersionError()
|
||||
}
|
||||
return value, nil
|
||||
}
|
||||
|
||||
// parseAgentExt validates one generic JSON object and returns its compact
|
||||
// one-line representation. Raw control characters other than horizontal tab
|
||||
// are rejected before JSON parsing; escaped JSON control characters remain
|
||||
// valid because they are safe on the HTTP header wire.
|
||||
func parseAgentExt(raw string) (string, error) {
|
||||
if len(raw) > maxAgentExtensionBytes || !utf8.ValidString(raw) {
|
||||
return "", invalidAgentExtError()
|
||||
}
|
||||
for _, r := range raw {
|
||||
if unicode.IsControl(r) && r != '\t' {
|
||||
return "", invalidAgentExtError()
|
||||
}
|
||||
}
|
||||
|
||||
value := strings.Trim(raw, " \t")
|
||||
if value == "" {
|
||||
return "", nil
|
||||
}
|
||||
|
||||
var compact bytes.Buffer
|
||||
if err := json.Compact(&compact, []byte(value)); err != nil {
|
||||
return "", invalidAgentExtError()
|
||||
}
|
||||
compactBytes := compact.Bytes()
|
||||
if len(compactBytes) > maxAgentExtensionBytes || len(compactBytes) < 2 || compactBytes[0] != '{' {
|
||||
return "", invalidAgentExtError()
|
||||
}
|
||||
return compact.String(), nil
|
||||
}
|
||||
|
||||
func invalidAgentVersionError() error {
|
||||
return apperrors.NewValidation(
|
||||
"DWS_AGENT_VER must be at most 64 bytes and match ^[A-Za-z0-9][A-Za-z0-9._+-]*$",
|
||||
apperrors.WithReason("invalid_agent_version"),
|
||||
)
|
||||
}
|
||||
|
||||
func invalidAgentExtError() error {
|
||||
return apperrors.NewValidation(
|
||||
"DWS_AGENT_EXT must be a UTF-8 JSON object of at most 8192 bytes without raw control characters",
|
||||
apperrors.WithReason("invalid_agent_ext"),
|
||||
)
|
||||
}
|
||||
|
||||
// readAgentMetadataSnapshot reads both environment variables from one
|
||||
// os.Environ snapshot, then parses them once. Normal CLI execution retains the
|
||||
// validated result through the invocation so hooks and transport observe the
|
||||
// same pair even in an embedding process that mutates its environment.
|
||||
func readAgentMetadataSnapshot() agentMetadataSnapshot {
|
||||
var rawVersion, rawExt string
|
||||
for _, entry := range os.Environ() {
|
||||
key, value, _ := strings.Cut(entry, "=")
|
||||
switch key {
|
||||
case envDWSAgentVersion:
|
||||
rawVersion = value
|
||||
case envDWSAgentExt:
|
||||
rawExt = value
|
||||
}
|
||||
}
|
||||
version, versionErr := parseAgentVersion(rawVersion)
|
||||
ext, extErr := parseAgentExt(rawExt)
|
||||
return agentMetadataSnapshot{
|
||||
version: version,
|
||||
ext: ext,
|
||||
versionErr: versionErr,
|
||||
extErr: extErr,
|
||||
}
|
||||
}
|
||||
|
||||
// removeAgentMetadataHeaders removes every case variant so edition or
|
||||
// credential hooks cannot smuggle MCP-only metadata into shared transports.
|
||||
func removeAgentMetadataHeaders(headers map[string]string) {
|
||||
for key := range headers {
|
||||
if strings.EqualFold(key, transport.HeaderAgentVersion) ||
|
||||
strings.EqualFold(key, transport.HeaderAgentExt) {
|
||||
delete(headers, key)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// applyAgentMetadataHeaders applies validated environment values as the final
|
||||
// authority for non-plugin MCP requests. Invalid values are omitted on
|
||||
// library paths that bypass root validation; normal CLI execution rejects
|
||||
// them before hooks or network access.
|
||||
func applyAgentMetadataHeaders(headers map[string]string) map[string]string {
|
||||
return applyAgentMetadataSnapshot(headers, readAgentMetadataSnapshot())
|
||||
}
|
||||
|
||||
func applyAgentMetadataSnapshot(headers map[string]string, snapshot agentMetadataSnapshot) map[string]string {
|
||||
removeAgentMetadataHeaders(headers)
|
||||
|
||||
if (snapshot.versionErr != nil || snapshot.version == "") && (snapshot.extErr != nil || snapshot.ext == "") {
|
||||
return headers
|
||||
}
|
||||
if headers == nil {
|
||||
headers = make(map[string]string)
|
||||
}
|
||||
if snapshot.versionErr == nil && snapshot.version != "" {
|
||||
headers[transport.HeaderAgentVersion] = snapshot.version
|
||||
}
|
||||
if snapshot.extErr == nil && snapshot.ext != "" {
|
||||
headers[transport.HeaderAgentExt] = snapshot.ext
|
||||
}
|
||||
return headers
|
||||
}
|
||||
|
||||
// resolveMCPRequestHeaders adds Agent version and extension metadata only to
|
||||
// the built-in DingTalk MCP request path. Shared identity consumers (notably
|
||||
// A2A) continue to use resolveIdentityHeaders and never receive these fields.
|
||||
func resolveMCPRequestHeaders() map[string]string {
|
||||
return resolveMCPRequestHeadersWithSnapshot(readAgentMetadataSnapshot())
|
||||
}
|
||||
|
||||
func resolveMCPRequestHeadersWithSnapshot(snapshot agentMetadataSnapshot) map[string]string {
|
||||
return applyAgentMetadataSnapshot(resolveIdentityHeaders(), snapshot)
|
||||
}
|
||||
|
||||
// resolveMCPRequestHeadersForInvocation resolves one immutable Header snapshot
|
||||
// for an invocation. The helper-only mcp-meta server performs endpoint
|
||||
// discovery rather than an ordinary MCP product call, so caller-declared
|
||||
// Agent metadata must not cross that boundary.
|
||||
func resolveMCPRequestHeadersForInvocation(invocation executor.Invocation, snapshots ...agentMetadataSnapshot) map[string]string {
|
||||
headers := resolveIdentityHeaders()
|
||||
if strings.EqualFold(strings.TrimSpace(invocation.CanonicalProduct), mcpMetaServerID) {
|
||||
return headers
|
||||
}
|
||||
snapshot := readAgentMetadataSnapshot()
|
||||
if len(snapshots) > 0 {
|
||||
snapshot = snapshots[0]
|
||||
}
|
||||
return applyAgentMetadataSnapshot(headers, snapshot)
|
||||
}
|
||||
|
||||
// pluginRequestHeaders returns a private, sanitized copy of plugin-owned
|
||||
// Headers. Third-party plugins never receive DWS-owned Agent metadata, even if
|
||||
// their manifest tries to declare the reserved Header names itself.
|
||||
func pluginRequestHeaders(pluginAuth *PluginAuth) map[string]string {
|
||||
if pluginAuth == nil || len(pluginAuth.ExtraHeaders) == 0 {
|
||||
return nil
|
||||
}
|
||||
headers := make(map[string]string, len(pluginAuth.ExtraHeaders))
|
||||
for key, value := range pluginAuth.ExtraHeaders {
|
||||
headers[key] = value
|
||||
}
|
||||
removeAgentMetadataHeaders(headers)
|
||||
if len(headers) == 0 {
|
||||
return nil
|
||||
}
|
||||
return headers
|
||||
}
|
||||
@@ -0,0 +1,743 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package app
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"io"
|
||||
"maps"
|
||||
"os"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/audit"
|
||||
apperrors "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/errors"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/executor"
|
||||
outputpkg "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/output"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/pipeline"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/testseam"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/transport"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/agentproduct"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/configmeta"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/mcptypes"
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
|
||||
func TestCrossPlatformCoverageParseAgentVersion(t *testing.T) {
|
||||
var nilContext context.Context
|
||||
if _, ok := agentMetadataSnapshotFromContext(nilContext); ok {
|
||||
t.Fatal("nil context unexpectedly contained Agent metadata")
|
||||
}
|
||||
wantSnapshot := agentMetadataSnapshot{version: "context-version", ext: "{}"}
|
||||
if got, ok := agentMetadataSnapshotFromContext(contextWithAgentMetadataSnapshot(context.Background(), wantSnapshot)); !ok || got != wantSnapshot {
|
||||
t.Fatalf("context Agent metadata = %#v, %v; want %#v", got, ok, wantSnapshot)
|
||||
}
|
||||
|
||||
valid := []struct {
|
||||
name string
|
||||
raw string
|
||||
want string
|
||||
}{
|
||||
{name: "unset", raw: "", want: ""},
|
||||
{name: "ASCII whitespace only", raw: " \t ", want: ""},
|
||||
{name: "semantic version", raw: "1.2.3", want: "1.2.3"},
|
||||
{name: "pre-release and build", raw: " v1.2.3-rc.1+build_7 ", want: "v1.2.3-rc.1+build_7"},
|
||||
{name: "maximum length", raw: strings.Repeat("a", maxAgentVersionBytes), want: strings.Repeat("a", maxAgentVersionBytes)},
|
||||
}
|
||||
for _, tc := range valid {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
got, err := parseAgentVersion(tc.raw)
|
||||
if err != nil {
|
||||
t.Fatalf("parseAgentVersion() error = %v", err)
|
||||
}
|
||||
if got != tc.want {
|
||||
t.Fatalf("parseAgentVersion() = %q, want %q", got, tc.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
invalid := []struct {
|
||||
name string
|
||||
raw string
|
||||
}{
|
||||
{name: "leading punctuation", raw: "-1.2.3"},
|
||||
{name: "internal space", raw: "1.2 3"},
|
||||
{name: "slash", raw: "1.2/3"},
|
||||
{name: "line feed", raw: "1.2.3\n"},
|
||||
{name: "carriage return", raw: "1.2.3\r"},
|
||||
{name: "NUL", raw: "1.2\x003"},
|
||||
{name: "Unicode", raw: "版本1"},
|
||||
{name: "too long", raw: strings.Repeat("a", maxAgentVersionBytes+1)},
|
||||
}
|
||||
for _, tc := range invalid {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
got, err := parseAgentVersion(tc.raw)
|
||||
if err == nil || got != "" {
|
||||
t.Fatalf("parseAgentVersion(%q) = %q, %v; want validation error", tc.raw, got, err)
|
||||
}
|
||||
assertAgentMetadataValidationError(t, err, "invalid_agent_version", tc.raw)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageParseAgentExt(t *testing.T) {
|
||||
boundary := `{"x":"` + strings.Repeat("a", maxAgentExtensionBytes-8) + `"}`
|
||||
if len(boundary) != maxAgentExtensionBytes {
|
||||
t.Fatalf("invalid boundary fixture size: %d", len(boundary))
|
||||
}
|
||||
|
||||
valid := []struct {
|
||||
name string
|
||||
raw string
|
||||
want string
|
||||
}{
|
||||
{name: "unset", raw: "", want: ""},
|
||||
{name: "ASCII whitespace only", raw: " \t ", want: ""},
|
||||
{name: "empty object", raw: "{}", want: "{}"},
|
||||
{name: "compact generic object", raw: " \t{ \"umt\": \"masked\",\t \"nested\": { \"ok\": true }, \"unknown\": [1, 2] }\t ", want: `{"umt":"masked","nested":{"ok":true},"unknown":[1,2]}`},
|
||||
{name: "Unicode value", raw: `{"ua":"千问办公/1.0"}`, want: `{"ua":"千问办公/1.0"}`},
|
||||
{name: "escaped control remains safe", raw: `{"ua":"line\nnext"}`, want: `{"ua":"line\nnext"}`},
|
||||
{name: "maximum length", raw: boundary, want: boundary},
|
||||
}
|
||||
for _, tc := range valid {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
got, err := parseAgentExt(tc.raw)
|
||||
if err != nil {
|
||||
t.Fatalf("parseAgentExt() error = %v", err)
|
||||
}
|
||||
if got != tc.want {
|
||||
t.Fatalf("parseAgentExt() = %q, want %q", got, tc.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
invalidUTF8 := string([]byte{'{', '"', 'x', '"', ':', '"', 0xff, '"', '}'})
|
||||
invalid := []struct {
|
||||
name string
|
||||
raw string
|
||||
}{
|
||||
{name: "too long raw input", raw: strings.Repeat(" ", maxAgentExtensionBytes+1)},
|
||||
{name: "invalid UTF-8", raw: invalidUTF8},
|
||||
{name: "array", raw: `[]`},
|
||||
{name: "string", raw: `"value"`},
|
||||
{name: "number", raw: `1`},
|
||||
{name: "boolean", raw: `true`},
|
||||
{name: "null", raw: `null`},
|
||||
{name: "malformed object", raw: `{"secret":"DO_NOT_ECHO"`},
|
||||
{name: "trailing value", raw: `{} {}`},
|
||||
{name: "line feed", raw: "{\n}"},
|
||||
{name: "carriage return", raw: "{\r}"},
|
||||
{name: "NUL", raw: "{\x00}"},
|
||||
{name: "vertical tab", raw: "{\v}"},
|
||||
{name: "form feed", raw: "{\f}"},
|
||||
{name: "DEL", raw: "{\x7f}"},
|
||||
{name: "C1 control", raw: "{\u0085}"},
|
||||
}
|
||||
for _, tc := range invalid {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
got, err := parseAgentExt(tc.raw)
|
||||
if err == nil || got != "" {
|
||||
t.Fatalf("parseAgentExt() = %q, %v; want validation error", got, err)
|
||||
}
|
||||
assertAgentMetadataValidationError(t, err, "invalid_agent_ext", tc.raw)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func assertAgentMetadataValidationError(t *testing.T, err error, reason, raw string) {
|
||||
t.Helper()
|
||||
var appErr *apperrors.Error
|
||||
if !errors.As(err, &appErr) {
|
||||
t.Fatalf("error type = %T, want *errors.Error", err)
|
||||
}
|
||||
if appErr.Category != apperrors.CategoryValidation || appErr.Reason != reason {
|
||||
t.Fatalf("error = category %q reason %q, want validation/%s", appErr.Category, appErr.Reason, reason)
|
||||
}
|
||||
if strings.Contains(raw, "DO_NOT_ECHO") && strings.Contains(err.Error(), "DO_NOT_ECHO") {
|
||||
t.Fatalf("error must not echo invalid value: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageAgentMetadataConfigRegistrationAndMasking(t *testing.T) {
|
||||
items := configmeta.All()
|
||||
var versionItem, extItem *configmeta.ConfigItem
|
||||
for i := range items {
|
||||
switch items[i].Name {
|
||||
case envDWSAgentVersion:
|
||||
versionItem = &items[i]
|
||||
case envDWSAgentExt:
|
||||
extItem = &items[i]
|
||||
}
|
||||
}
|
||||
if versionItem == nil || extItem == nil {
|
||||
t.Fatalf("Agent metadata config registration missing: version=%v ext=%v", versionItem != nil, extItem != nil)
|
||||
}
|
||||
if versionItem.Category != configmeta.CategoryExternal || versionItem.Sensitive {
|
||||
t.Fatalf("version config metadata = %#v", *versionItem)
|
||||
}
|
||||
if extItem.Category != configmeta.CategoryExternal || !extItem.Sensitive {
|
||||
t.Fatalf("extension config metadata = %#v", *extItem)
|
||||
}
|
||||
|
||||
const canary = `{"umt":"SENSITIVE_CANARY"}`
|
||||
t.Setenv(envDWSAgentExt, canary)
|
||||
got, ok := configmeta.Resolve(envDWSAgentExt)
|
||||
if !ok || got == "" || strings.Contains(got, "SENSITIVE_CANARY") || got == canary {
|
||||
t.Fatalf("sensitive extension was not masked: value=%q ok=%v", got, ok)
|
||||
}
|
||||
|
||||
t.Setenv(envDWSAgentVersion, "9.8.7")
|
||||
command := newConfigListCommand()
|
||||
var output strings.Builder
|
||||
command.SetOut(&output)
|
||||
command.SetArgs([]string{"--category", string(configmeta.CategoryExternal), "--show-values", "--json"})
|
||||
if err := command.Execute(); err != nil {
|
||||
t.Fatalf("config list failed: %v", err)
|
||||
}
|
||||
rawOutput := output.String()
|
||||
if !json.Valid([]byte(rawOutput)) {
|
||||
t.Fatalf("config list emitted invalid JSON: %q", rawOutput)
|
||||
}
|
||||
if !strings.Contains(rawOutput, envDWSAgentVersion) || !strings.Contains(rawOutput, envDWSAgentExt) {
|
||||
t.Fatalf("config list omitted Agent metadata variables: %s", rawOutput)
|
||||
}
|
||||
if strings.Contains(rawOutput, "SENSITIVE_CANARY") || strings.Contains(rawOutput, canary) {
|
||||
t.Fatalf("config list leaked Agent extension: %s", rawOutput)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageResolveMCPRequestHeadersScopesAndFinalizesAgentMetadata(t *testing.T) {
|
||||
t.Setenv("DWS_CONFIG_DIR", t.TempDir())
|
||||
t.Setenv(envDWSAgentHost, "")
|
||||
t.Setenv(agentproduct.EnvName, "")
|
||||
t.Setenv(envDWSAgentVersion, " 1.2.3-rc.1 ")
|
||||
t.Setenv(envDWSAgentExt, " { \"umt\": \"masked\", \"unknown\": true } ")
|
||||
|
||||
oldEdition := edition.Get()
|
||||
t.Cleanup(func() { edition.Override(oldEdition) })
|
||||
edition.Override(&edition.Hooks{
|
||||
MergeHeaders: func(headers map[string]string) map[string]string {
|
||||
headers["X-Dws-Agent-Ver"] = "merge-must-not-win"
|
||||
headers["X-Dws-Agent-Ext"] = `{"source":"merge"}`
|
||||
return headers
|
||||
},
|
||||
EnterpriseCredentialHeaders: func(headers map[string]string) map[string]string {
|
||||
headers[transport.HeaderAgentVersion] = "credential-must-not-win"
|
||||
headers[transport.HeaderAgentExt] = `{"source":"credential"}`
|
||||
return headers
|
||||
},
|
||||
})
|
||||
|
||||
for name, headers := range map[string]map[string]string{
|
||||
"shared identity": resolveIdentityHeaders(),
|
||||
"A2A export": MCPIdentityHeaders(),
|
||||
} {
|
||||
if hasHeaderFold(headers, transport.HeaderAgentVersion) || hasHeaderFold(headers, transport.HeaderAgentExt) {
|
||||
t.Fatalf("%s leaked MCP-only metadata: %#v", name, headers)
|
||||
}
|
||||
}
|
||||
|
||||
headers := resolveMCPRequestHeaders()
|
||||
if got := headers[transport.HeaderAgentVersion]; got != "1.2.3-rc.1" {
|
||||
t.Fatalf("%s = %q, want 1.2.3-rc.1", transport.HeaderAgentVersion, got)
|
||||
}
|
||||
if got := headers[transport.HeaderAgentExt]; got != `{"umt":"masked","unknown":true}` {
|
||||
t.Fatalf("%s = %q", transport.HeaderAgentExt, got)
|
||||
}
|
||||
if got := headers[transport.HeaderVersion]; got != version {
|
||||
t.Fatalf("%s = %q, want CLI version %q", transport.HeaderVersion, got, version)
|
||||
}
|
||||
if _, ok := headers["User-Agent"]; ok {
|
||||
t.Fatal("Agent extension must not create or replace the standard User-Agent header")
|
||||
}
|
||||
for _, key := range []string{"umt", "miniwua", "ua", "x-dws-agent-umt", "x-dws-agent-miniwua", "x-dws-agent-ua"} {
|
||||
if hasHeaderFold(headers, key) {
|
||||
t.Fatalf("Agent extension was split into an extra header %q: %#v", key, headers)
|
||||
}
|
||||
}
|
||||
|
||||
// Library paths are best-effort: one invalid value is omitted without
|
||||
// suppressing the other valid field or preserving hook-injected values.
|
||||
t.Setenv(envDWSAgentExt, `{"secret":"DO_NOT_ECHO"`)
|
||||
headers = resolveMCPRequestHeaders()
|
||||
if got := headers[transport.HeaderAgentVersion]; got != "1.2.3-rc.1" {
|
||||
t.Fatalf("valid version was suppressed: %q", got)
|
||||
}
|
||||
if hasHeaderFold(headers, transport.HeaderAgentExt) {
|
||||
t.Fatalf("invalid extension or hook value leaked: %#v", headers)
|
||||
}
|
||||
|
||||
// Exercise the nil-map and empty-input library paths. An absent environment
|
||||
// must not allocate a map, while an EXT-only value must allocate one and
|
||||
// remain a single compact Header.
|
||||
t.Setenv(envDWSAgentVersion, "")
|
||||
t.Setenv(envDWSAgentExt, "")
|
||||
if got := applyAgentMetadataHeaders(nil); got != nil {
|
||||
t.Fatalf("empty metadata allocated headers: %#v", got)
|
||||
}
|
||||
t.Setenv(envDWSAgentExt, " { } ")
|
||||
headers = applyAgentMetadataHeaders(nil)
|
||||
if got := headers[transport.HeaderAgentExt]; got != "{}" {
|
||||
t.Fatalf("EXT-only metadata = %q, want {}", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageRootRejectsInvalidAgentMetadataBeforeEditionHook(t *testing.T) {
|
||||
oldEdition := edition.Get()
|
||||
t.Cleanup(func() { edition.Override(oldEdition) })
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
env string
|
||||
value string
|
||||
reason string
|
||||
}{
|
||||
{name: "version", env: envDWSAgentVersion, value: "DO_NOT ECHO", reason: "invalid_agent_version"},
|
||||
{name: "extension", env: envDWSAgentExt, value: `{"secret":"DO_NOT_ECHO"`, reason: "invalid_agent_ext"},
|
||||
}
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Setenv("DWS_CONFIG_DIR", t.TempDir())
|
||||
t.Setenv(envDWSAgentHost, "")
|
||||
t.Setenv(agentproduct.EnvName, "")
|
||||
t.Setenv(envDWSAgentVersion, "")
|
||||
t.Setenv(envDWSAgentExt, "")
|
||||
t.Setenv(tc.env, tc.value)
|
||||
|
||||
headerHookCalled := false
|
||||
afterHookCalled := false
|
||||
edition.Override(&edition.Hooks{
|
||||
MergeHeaders: func(headers map[string]string) map[string]string {
|
||||
headerHookCalled = true
|
||||
return headers
|
||||
},
|
||||
EnterpriseCredentialHeaders: func(headers map[string]string) map[string]string {
|
||||
headerHookCalled = true
|
||||
return headers
|
||||
},
|
||||
AfterPersistentPreRun: func(_ *cobra.Command, _ []string) error {
|
||||
afterHookCalled = true
|
||||
return nil
|
||||
},
|
||||
})
|
||||
|
||||
root := NewRootCommand()
|
||||
root.SetOut(io.Discard)
|
||||
root.SetErr(io.Discard)
|
||||
root.SetArgs([]string{"version"})
|
||||
err := root.Execute()
|
||||
if err == nil {
|
||||
t.Fatalf("root command accepted invalid %s", tc.env)
|
||||
}
|
||||
if headerHookCalled || afterHookCalled {
|
||||
t.Fatalf("edition hook ran before %s validation", tc.env)
|
||||
}
|
||||
assertAgentMetadataValidationError(t, err, tc.reason, tc.value)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageAgentMetadataProcessEntryValidationPrecedesRootConstruction(t *testing.T) {
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
args []string
|
||||
want bool
|
||||
}{
|
||||
{name: "default JSON", args: []string{"version"}, want: true},
|
||||
{name: "long JSON", args: []string{"version", "--format", "JSON"}, want: true},
|
||||
{name: "long table", args: []string{"--format=table", "version"}, want: false},
|
||||
{name: "short attached JSON", args: []string{"version", "-fjson"}, want: true},
|
||||
{name: "short table", args: []string{"version", "-f", "table"}, want: false},
|
||||
{name: "last wins", args: []string{"--format", "table", "version", "-f=json"}, want: true},
|
||||
{name: "terminator", args: []string{"version", "--format", "table", "--", "--format", "json"}, want: false},
|
||||
{name: "missing value", args: []string{"version", "--format"}, want: false},
|
||||
} {
|
||||
t.Run("presentation/"+tc.name, func(t *testing.T) {
|
||||
if got := processArgsRequestJSON(tc.args); got != tc.want {
|
||||
t.Fatalf("processArgsRequestJSON(%q) = %v, want %v", tc.args, got, tc.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
t.Setenv("DWS_CONFIG_DIR", t.TempDir())
|
||||
t.Setenv(envDWSAgentHost, "")
|
||||
t.Setenv(agentproduct.EnvName, "")
|
||||
t.Setenv(envDWSAgentVersion, "")
|
||||
sensitiveRaw := "{\"umt\":\"must-not-leak\"}\n"
|
||||
t.Setenv(envDWSAgentExt, sensitiveRaw)
|
||||
oldEdition := edition.Get()
|
||||
t.Cleanup(func() { edition.Override(oldEdition) })
|
||||
extensionHookCalls := 0
|
||||
edition.Override(&edition.Hooks{
|
||||
Name: "presentation-test",
|
||||
RegisterExtraCommands: func(*cobra.Command, edition.ToolCaller) {
|
||||
extensionHookCalls++
|
||||
},
|
||||
VisibleProducts: func() []string {
|
||||
extensionHookCalls++
|
||||
return nil
|
||||
},
|
||||
StaticServers: func() []edition.ServerInfo {
|
||||
extensionHookCalls++
|
||||
return nil
|
||||
},
|
||||
})
|
||||
|
||||
oldArgs := os.Args
|
||||
os.Args = []string{"dws", "version"}
|
||||
t.Cleanup(func() { os.Args = oldArgs })
|
||||
|
||||
rootConstructed := false
|
||||
preParseCalled := false
|
||||
testseam.Swap(t, &rootNewRootCommandWithEngine, func(context.Context, *pipeline.Engine) *cobra.Command {
|
||||
rootConstructed = true
|
||||
return &cobra.Command{Use: "dws"}
|
||||
})
|
||||
testseam.Swap(t, &rootRunPreParse, func(*cobra.Command, *pipeline.Engine) error {
|
||||
preParseCalled = true
|
||||
return nil
|
||||
})
|
||||
|
||||
stderrFile, err := os.CreateTemp(t.TempDir(), "agent-metadata-stderr-*")
|
||||
if err != nil {
|
||||
t.Fatalf("create stderr capture: %v", err)
|
||||
}
|
||||
oldStderr := os.Stderr
|
||||
os.Stderr = stderrFile
|
||||
t.Cleanup(func() {
|
||||
os.Stderr = oldStderr
|
||||
_ = stderrFile.Close()
|
||||
})
|
||||
|
||||
if code := Execute(); code == 0 {
|
||||
t.Fatal("process entry accepted invalid Agent metadata")
|
||||
}
|
||||
if rootConstructed || preParseCalled {
|
||||
t.Fatalf("invalid Agent metadata reached root hooks: constructed=%v preParse=%v", rootConstructed, preParseCalled)
|
||||
}
|
||||
if extensionHookCalls != 0 {
|
||||
t.Fatalf("invalid Agent metadata executed %d extension hooks", extensionHookCalls)
|
||||
}
|
||||
if err := stderrFile.Sync(); err != nil {
|
||||
t.Fatalf("sync stderr capture: %v", err)
|
||||
}
|
||||
stderrOutput, err := os.ReadFile(stderrFile.Name())
|
||||
if err != nil {
|
||||
t.Fatalf("read stderr capture: %v", err)
|
||||
}
|
||||
if strings.Contains(string(stderrOutput), "must-not-leak") || strings.Contains(string(stderrOutput), sensitiveRaw) {
|
||||
t.Fatalf("process validation error leaked raw EXT: %q", stderrOutput)
|
||||
}
|
||||
if !json.Valid(stderrOutput) || !strings.Contains(string(stderrOutput), `"reason": "invalid_agent_ext"`) {
|
||||
t.Fatalf("default JSON error presentation = %q", stderrOutput)
|
||||
}
|
||||
|
||||
stdoutFile, err := os.CreateTemp(t.TempDir(), "agent-metadata-stdout-*")
|
||||
if err != nil {
|
||||
t.Fatalf("create stdout capture: %v", err)
|
||||
}
|
||||
oldStdout := os.Stdout
|
||||
os.Stdout = stdoutFile
|
||||
t.Cleanup(func() {
|
||||
os.Stdout = oldStdout
|
||||
_ = stdoutFile.Close()
|
||||
})
|
||||
emitEarlyAgentMetadataValidationError(invalidAgentExtError(), []string{"drive", "+list", "--format", "json"})
|
||||
if err := stdoutFile.Sync(); err != nil {
|
||||
t.Fatalf("sync stdout capture: %v", err)
|
||||
}
|
||||
unifiedOutput, err := os.ReadFile(stdoutFile.Name())
|
||||
if err != nil {
|
||||
t.Fatalf("read stdout capture: %v", err)
|
||||
}
|
||||
if !json.Valid(unifiedOutput) || !strings.Contains(string(unifiedOutput), `"outcome": "failure"`) ||
|
||||
!strings.Contains(string(unifiedOutput), `"subtype": "invalid_agent_ext"`) {
|
||||
t.Fatalf("unified JSON error presentation = %q", unifiedOutput)
|
||||
}
|
||||
if extensionHookCalls != 0 {
|
||||
t.Fatalf("presentation-only root executed %d extension hooks", extensionHookCalls)
|
||||
}
|
||||
|
||||
if err := stderrFile.Truncate(0); err != nil {
|
||||
t.Fatalf("truncate fallback stderr capture: %v", err)
|
||||
}
|
||||
if _, err := stderrFile.Seek(0, io.SeekStart); err != nil {
|
||||
t.Fatalf("rewind fallback stderr capture: %v", err)
|
||||
}
|
||||
testseam.Swap(t, &rootEmitResult, func(*cobra.Command, outputpkg.CommandResult) (int, error) {
|
||||
return 0, errors.New("injected result emission failure")
|
||||
})
|
||||
emitEarlyAgentMetadataValidationError(invalidAgentExtError(), []string{"drive", "+list", "--format", "json"})
|
||||
if err := stderrFile.Sync(); err != nil {
|
||||
t.Fatalf("sync fallback stderr capture: %v", err)
|
||||
}
|
||||
fallbackOutput, err := os.ReadFile(stderrFile.Name())
|
||||
if err != nil {
|
||||
t.Fatalf("read fallback stderr capture: %v", err)
|
||||
}
|
||||
if !json.Valid(fallbackOutput) || !strings.Contains(string(fallbackOutput), `"reason": "invalid_agent_ext"`) ||
|
||||
strings.Contains(string(fallbackOutput), "must-not-leak") {
|
||||
t.Fatalf("fallback validation error presentation = %q", fallbackOutput)
|
||||
}
|
||||
|
||||
if err := stderrFile.Truncate(0); err != nil {
|
||||
t.Fatalf("truncate stderr capture: %v", err)
|
||||
}
|
||||
if _, err := stderrFile.Seek(0, io.SeekStart); err != nil {
|
||||
t.Fatalf("rewind stderr capture: %v", err)
|
||||
}
|
||||
emitEarlyAgentMetadataValidationError(invalidAgentExtError(), []string{"version", "--format", "table"})
|
||||
if err := stderrFile.Sync(); err != nil {
|
||||
t.Fatalf("sync human stderr capture: %v", err)
|
||||
}
|
||||
humanOutput, err := os.ReadFile(stderrFile.Name())
|
||||
if err != nil {
|
||||
t.Fatalf("read human stderr capture: %v", err)
|
||||
}
|
||||
if json.Valid(humanOutput) || !strings.Contains(string(humanOutput), "DWS_AGENT_EXT") ||
|
||||
strings.Contains(string(humanOutput), "must-not-leak") {
|
||||
t.Fatalf("human validation error presentation = %q", humanOutput)
|
||||
}
|
||||
|
||||
var capturedRunner *runtimeRunner
|
||||
testseam.Swap(t, &rootNewCommandRunnerWithFlags, func(flags *GlobalFlags) executor.Runner {
|
||||
capturedRunner = newCommandRunnerWithFlags(flags).(*runtimeRunner)
|
||||
return capturedRunner
|
||||
})
|
||||
cachedSnapshot := agentMetadataSnapshot{version: "9.8.7", ext: `{"ua":"cached"}`}
|
||||
_ = newRootCommandWithMode(
|
||||
contextWithAgentMetadataSnapshot(context.Background(), cachedSnapshot),
|
||||
nil,
|
||||
false,
|
||||
true,
|
||||
true,
|
||||
)
|
||||
if capturedRunner == nil || capturedRunner.agentMetadata == nil || *capturedRunner.agentMetadata != cachedSnapshot {
|
||||
t.Fatalf("root runner Agent metadata = %#v, want %#v", capturedRunner, cachedSnapshot)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageAgentMetadataExcludedFromServiceDiscovery(t *testing.T) {
|
||||
t.Setenv("DWS_CONFIG_DIR", t.TempDir())
|
||||
t.Setenv(envDWSAgentHost, "")
|
||||
t.Setenv(agentproduct.EnvName, "")
|
||||
t.Setenv(envDWSAgentVersion, "3.0.0")
|
||||
t.Setenv(envDWSAgentExt, `{"umt":"test-value"}`)
|
||||
|
||||
headers := resolveMCPRequestHeadersForInvocation(executor.Invocation{
|
||||
CanonicalProduct: mcpMetaServerID,
|
||||
Tool: mcpMetaURLTool,
|
||||
})
|
||||
if hasHeaderFold(headers, transport.HeaderAgentVersion) || hasHeaderFold(headers, transport.HeaderAgentExt) {
|
||||
t.Fatalf("service-discovery request leaked Agent metadata: %#v", headers)
|
||||
}
|
||||
|
||||
headers = resolveMCPRequestHeadersForInvocation(executor.Invocation{CanonicalProduct: "doc", Tool: "read"})
|
||||
if headers[transport.HeaderAgentVersion] != "3.0.0" || headers[transport.HeaderAgentExt] == "" {
|
||||
t.Fatalf("ordinary MCP request omitted Agent metadata: %#v", headers)
|
||||
}
|
||||
cached := agentMetadataSnapshot{version: "3.1.0", ext: "{}"}
|
||||
headers = resolveMCPRequestHeadersForInvocation(executor.Invocation{CanonicalProduct: "doc", Tool: "read"}, cached)
|
||||
if headers[transport.HeaderAgentVersion] != "3.1.0" || headers[transport.HeaderAgentExt] != "{}" {
|
||||
t.Fatalf("ordinary MCP request ignored its validated snapshot: %#v", headers)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageAgentMetadataMCPAndPluginScoping(t *testing.T) {
|
||||
t.Setenv("DWS_CONFIG_DIR", t.TempDir())
|
||||
t.Setenv(envDWSAgentHost, "")
|
||||
t.Setenv(agentproduct.EnvName, "")
|
||||
t.Setenv(envDWSAgentVersion, "2.0.0")
|
||||
t.Setenv(envDWSAgentExt, `{"ua":"test-agent/2.0"}`)
|
||||
|
||||
oldEdition := edition.Get()
|
||||
t.Cleanup(func() { edition.Override(oldEdition) })
|
||||
edition.Override(&edition.Hooks{})
|
||||
|
||||
pluginAuthMu.Lock()
|
||||
oldPluginRegistry := pluginAuthRegistry
|
||||
pluginAuthRegistry = make(map[string]*PluginAuth)
|
||||
pluginAuthMu.Unlock()
|
||||
t.Cleanup(func() {
|
||||
pluginAuthMu.Lock()
|
||||
pluginAuthRegistry = oldPluginRegistry
|
||||
pluginAuthMu.Unlock()
|
||||
})
|
||||
dynamicMu.Lock()
|
||||
oldDynamicEndpoints := dynamicEndpoints
|
||||
oldDynamicProducts := dynamicProducts
|
||||
oldDynamicAliases := dynamicAliases
|
||||
oldDynamicToolEndpoints := dynamicToolEndpoints
|
||||
dynamicEndpoints = nil
|
||||
dynamicProducts = nil
|
||||
dynamicAliases = nil
|
||||
dynamicToolEndpoints = nil
|
||||
dynamicMu.Unlock()
|
||||
t.Cleanup(func() {
|
||||
dynamicMu.Lock()
|
||||
dynamicEndpoints = oldDynamicEndpoints
|
||||
dynamicProducts = oldDynamicProducts
|
||||
dynamicAliases = oldDynamicAliases
|
||||
dynamicToolEndpoints = oldDynamicToolEndpoints
|
||||
dynamicMu.Unlock()
|
||||
})
|
||||
|
||||
testseam.Swap(t, &runnerPreflightDocDownload, func(*runtimeRunner, context.Context, *transport.Client, string, executor.Invocation) error {
|
||||
return nil
|
||||
})
|
||||
type capturedRequest struct {
|
||||
headers map[string]string
|
||||
token string
|
||||
}
|
||||
var captured []capturedRequest
|
||||
testseam.Swap(t, &runnerCallTool, func(client *transport.Client, _ context.Context, _, _ string, _ map[string]any) (transport.ToolCallResult, error) {
|
||||
copyHeaders := make(map[string]string, len(client.ExtraHeaders))
|
||||
for key, value := range client.ExtraHeaders {
|
||||
copyHeaders[key] = value
|
||||
}
|
||||
captured = append(captured, capturedRequest{headers: copyHeaders, token: client.AuthToken})
|
||||
return transport.ToolCallResult{Content: map[string]any{"value": "ok"}}, nil
|
||||
})
|
||||
|
||||
created := newCommandRunnerWithFlags(&GlobalFlags{}).(*runtimeRunner)
|
||||
if hasHeaderFold(created.transport.ExtraHeaders, transport.HeaderAgentVersion) ||
|
||||
hasHeaderFold(created.transport.ExtraHeaders, transport.HeaderAgentExt) {
|
||||
t.Fatalf("new runner resolved Agent metadata before invocation validation: %#v", created.transport.ExtraHeaders)
|
||||
}
|
||||
|
||||
// runSingle must not cache ambient MCP metadata on the shared base transport.
|
||||
// Use mock mode to exercise the path without authentication or network I/O.
|
||||
t.Setenv(envDWSAgentVersion, "2.0.1")
|
||||
refreshRunner := &runtimeRunner{
|
||||
transport: transport.NewClient(nil),
|
||||
globalFlags: &GlobalFlags{Mock: true},
|
||||
auditSink: audit.NopSink{},
|
||||
}
|
||||
refreshInvocation := executor.Invocation{CanonicalProduct: "refresh", Tool: "tool", Params: map[string]any{}}
|
||||
if _, err := refreshRunner.runSingle(context.Background(), refreshInvocation, false); err != nil {
|
||||
t.Fatalf("mock runSingle failed: %v", err)
|
||||
}
|
||||
if hasHeaderFold(refreshRunner.transport.ExtraHeaders, transport.HeaderAgentVersion) {
|
||||
t.Fatalf("runSingle mutated the shared transport Header map: %#v", refreshRunner.transport.ExtraHeaders)
|
||||
}
|
||||
t.Setenv(envDWSAgentVersion, "2.0.0")
|
||||
|
||||
r := &runtimeRunner{
|
||||
transport: transport.NewClient(nil),
|
||||
globalFlags: &GlobalFlags{Token: "test-token"},
|
||||
auditSink: audit.NopSink{},
|
||||
agentMetadata: &agentMetadataSnapshot{
|
||||
version: "2.0.0",
|
||||
ext: `{"ua":"test-agent/2.0"}`,
|
||||
},
|
||||
}
|
||||
builtIn := executor.Invocation{CanonicalProduct: "built-in", Tool: "tool", Params: map[string]any{}}
|
||||
if _, err := r.executeInvocation(context.Background(), "https://example.test", builtIn); err != nil {
|
||||
t.Fatalf("built-in invocation failed: %v", err)
|
||||
}
|
||||
|
||||
pluginDescriptor := mcptypes.ServerDescriptor{
|
||||
Key: "third-party",
|
||||
Endpoint: "https://plugin.example.test",
|
||||
CLI: mcptypes.CLIOverlay{ID: "third-party"},
|
||||
AuthHeaders: map[string]string{
|
||||
"X-Plugin": "yes",
|
||||
"X-Dws-Agent-Ver": "plugin-must-not-forge-version",
|
||||
"X-Dws-Agent-Ext": `{"source":"plugin"}`,
|
||||
},
|
||||
}
|
||||
registerPluginHTTPServer(pluginDescriptor)
|
||||
registeredPlugin, pluginOwned := LookupPluginAuth("third-party")
|
||||
if !pluginOwned || registeredPlugin == nil || registeredPlugin.Token != "" {
|
||||
t.Fatalf("anonymous HTTP plugin ownership = %#v, %v", registeredPlugin, pluginOwned)
|
||||
}
|
||||
registerPluginHTTPServer(mcptypes.ServerDescriptor{
|
||||
Key: "anonymous-empty",
|
||||
Endpoint: "https://anonymous.example.test",
|
||||
CLI: mcptypes.CLIOverlay{ID: "anonymous-empty"},
|
||||
})
|
||||
if emptyPlugin, owned := LookupPluginAuth("anonymous-empty"); !owned || emptyPlugin == nil || emptyPlugin.Token != "" || len(emptyPlugin.ExtraHeaders) != 0 {
|
||||
t.Fatalf("headerless HTTP plugin ownership = %#v, %v", emptyPlugin, owned)
|
||||
}
|
||||
originalPluginHeaders := maps.Clone(registeredPlugin.ExtraHeaders)
|
||||
pluginInvocation := executor.Invocation{CanonicalProduct: "third-party", Tool: "tool", Params: map[string]any{}}
|
||||
if _, err := r.executeInvocation(context.Background(), "https://plugin.example.test", pluginInvocation); err != nil {
|
||||
t.Fatalf("plugin invocation failed: %v", err)
|
||||
}
|
||||
|
||||
if len(captured) != 2 {
|
||||
t.Fatalf("captured %d calls, want 2", len(captured))
|
||||
}
|
||||
if captured[0].headers[transport.HeaderAgentVersion] != "2.0.0" || captured[0].headers[transport.HeaderAgentExt] != `{"ua":"test-agent/2.0"}` {
|
||||
t.Fatalf("built-in MCP metadata = %#v", captured[0].headers)
|
||||
}
|
||||
if hasHeaderFold(captured[1].headers, transport.HeaderAgentVersion) || hasHeaderFold(captured[1].headers, transport.HeaderAgentExt) {
|
||||
t.Fatalf("plugin request leaked Agent metadata: %#v", captured[1].headers)
|
||||
}
|
||||
if got := captured[1].headers["X-Plugin"]; got != "yes" {
|
||||
t.Fatalf("plugin-owned header = %q, want yes", got)
|
||||
}
|
||||
if captured[1].token != "" {
|
||||
t.Fatalf("anonymous plugin unexpectedly received default OAuth token")
|
||||
}
|
||||
if !maps.Equal(registeredPlugin.ExtraHeaders, originalPluginHeaders) {
|
||||
t.Fatalf("plugin Header sanitization mutated registry state: got %#v want %#v", registeredPlugin.ExtraHeaders, originalPluginHeaders)
|
||||
}
|
||||
if got := pluginRequestHeaders(nil); got != nil {
|
||||
t.Fatalf("nil plugin auth produced Headers: %#v", got)
|
||||
}
|
||||
if got := pluginRequestHeaders(&PluginAuth{ExtraHeaders: map[string]string{
|
||||
"X-DWS-AGENT-VER": "forged",
|
||||
"X-DWS-AGENT-EXT": `{"forged":true}`,
|
||||
}}); got != nil {
|
||||
t.Fatalf("reserved-only plugin Headers survived sanitization: %#v", got)
|
||||
}
|
||||
|
||||
// Keep the execution-boundary auth guard independently testable: even if a
|
||||
// future token provider returns an empty token without an error, built-in MCP
|
||||
// calls must fail before preflight or transport while anonymous plugins remain
|
||||
// valid above.
|
||||
resolveCalled := false
|
||||
testseam.Swap(t, &runnerResolveAuthToken, func(*runtimeRunner, context.Context) (string, error) {
|
||||
resolveCalled = true
|
||||
return "", nil
|
||||
})
|
||||
callsBefore := len(captured)
|
||||
unauthenticated := &runtimeRunner{
|
||||
transport: transport.NewClient(nil),
|
||||
globalFlags: &GlobalFlags{},
|
||||
auditSink: audit.NopSink{},
|
||||
}
|
||||
if _, err := unauthenticated.executeInvocation(context.Background(), "https://example.test", executor.Invocation{CanonicalProduct: "built-in-unauthenticated", Tool: "tool"}); err == nil || !isAuthError(err) {
|
||||
t.Fatalf("unauthenticated built-in request = %v, want auth error", err)
|
||||
}
|
||||
if !resolveCalled {
|
||||
t.Fatal("unauthenticated request did not exercise the token resolver")
|
||||
}
|
||||
if len(captured) != callsBefore {
|
||||
t.Fatalf("unauthenticated built-in request reached transport: calls %d -> %d", callsBefore, len(captured))
|
||||
}
|
||||
}
|
||||
|
||||
func hasHeaderFold(headers map[string]string, want string) bool {
|
||||
for key := range headers {
|
||||
if strings.EqualFold(key, want) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
@@ -158,7 +158,7 @@ func TestApplyAgentProductHeader(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestRootRejectsInvalidAgentProductBeforeEditionHook(t *testing.T) {
|
||||
func TestCrossPlatformCoverageRootRejectsInvalidAgentProductBeforeEditionHook(t *testing.T) {
|
||||
t.Setenv("DWS_CONFIG_DIR", t.TempDir())
|
||||
const invalidValue = "DO_NOT ECHO"
|
||||
t.Setenv(agentproduct.EnvName, invalidValue)
|
||||
|
||||
@@ -15,11 +15,10 @@ package app
|
||||
|
||||
import "sync"
|
||||
|
||||
// PluginAuth holds authentication credentials for a plugin-owned
|
||||
// streamable-http MCP server. Each server is keyed by its canonical
|
||||
// product ID (CLI.ID) so that different servers can use independent
|
||||
// tokens without interfering with each other or with the default
|
||||
// DingTalk OAuth token.
|
||||
// PluginAuth marks ownership of a plugin-owned streamable-http MCP server and
|
||||
// holds its optional authentication credentials. Every accepted HTTP plugin,
|
||||
// including an anonymous one, has a non-nil record keyed by canonical product
|
||||
// ID (CLI.ID) so execution never falls back to built-in DingTalk OAuth.
|
||||
type PluginAuth struct {
|
||||
// Token is the Bearer token extracted from the plugin's
|
||||
// "Authorization" header (e.g. a third-party API key).
|
||||
@@ -39,27 +38,25 @@ var (
|
||||
pluginAuthRegistry = make(map[string]*PluginAuth)
|
||||
)
|
||||
|
||||
// RegisterPluginAuth stores authentication credentials for a plugin
|
||||
// server keyed by its canonical product ID. The runner looks up these
|
||||
// credentials at execution time to inject the correct Bearer token
|
||||
// instead of the default DingTalk OAuth token.
|
||||
// RegisterPluginAuth stores ownership and optional authentication credentials
|
||||
// for a plugin server keyed by its canonical product ID.
|
||||
func RegisterPluginAuth(productID string, auth *PluginAuth) {
|
||||
pluginAuthMu.Lock()
|
||||
defer pluginAuthMu.Unlock()
|
||||
pluginAuthRegistry[productID] = auth
|
||||
}
|
||||
|
||||
// ClearPluginAuth removes credentials for a plugin product. Registration uses
|
||||
// this before applying an accepted descriptor so a descriptor without custom
|
||||
// auth cannot inherit stale credentials from an earlier root construction.
|
||||
// ClearPluginAuth removes the ownership and credential record for a plugin
|
||||
// product.
|
||||
func ClearPluginAuth(productID string) {
|
||||
pluginAuthMu.Lock()
|
||||
defer pluginAuthMu.Unlock()
|
||||
delete(pluginAuthRegistry, productID)
|
||||
}
|
||||
|
||||
// LookupPluginAuth returns the authentication credentials registered
|
||||
// for the given product ID, or nil if none exists.
|
||||
// LookupPluginAuth returns plugin ownership and optional authentication
|
||||
// credentials for the product ID. The bool denotes ownership, not whether a
|
||||
// Bearer token is present.
|
||||
func LookupPluginAuth(productID string) (*PluginAuth, bool) {
|
||||
pluginAuthMu.RLock()
|
||||
defer pluginAuthMu.RUnlock()
|
||||
|
||||
@@ -61,11 +61,16 @@ func appRPCServer(t *testing.T, initOK, listOK bool) *httptest.Server {
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoveragePluginAuthCoverage(t *testing.T) {
|
||||
registerPluginAuthFromHeaders(mcptypes.ServerDescriptor{Key: "fallback", Endpoint: "%", AuthHeaders: map[string]string{"Authorization": "token"}})
|
||||
registerPluginAuthFromHeaders(mcptypes.ServerDescriptor{Key: "server", Endpoint: "https://x.test", CLI: mcptypes.CLIOverlay{ID: "cli"}, AuthHeaders: map[string]string{"Authorization": "Bearer token", "X": "Y"}})
|
||||
registerPluginAuthFromHeaders(mcptypes.ServerDescriptor{Key: "none"})
|
||||
if got, ok := LookupPluginAuth("cli"); !ok || got == nil || got.Token != "token" {
|
||||
t.Fatalf("registered plugin auth = %#v, %v", got, ok)
|
||||
fallback := pluginAuthFromServerDescriptor(mcptypes.ServerDescriptor{Key: "fallback", Endpoint: "%", AuthHeaders: map[string]string{"Authorization": "token"}})
|
||||
if fallback == nil || fallback.Token != "token" || len(fallback.TrustedDomains) != 0 {
|
||||
t.Fatalf("fallback plugin auth = %#v", fallback)
|
||||
}
|
||||
got := pluginAuthFromServerDescriptor(mcptypes.ServerDescriptor{Key: "server", Endpoint: "https://x.test", CLI: mcptypes.CLIOverlay{ID: "cli"}, AuthHeaders: map[string]string{"Authorization": "Bearer token", "X": "Y"}})
|
||||
if got == nil || got.Token != "token" || got.ExtraHeaders["X"] != "Y" || len(got.TrustedDomains) != 2 {
|
||||
t.Fatalf("plugin auth = %#v", got)
|
||||
}
|
||||
if anonymous := pluginAuthFromServerDescriptor(mcptypes.ServerDescriptor{Key: "none"}); anonymous == nil || anonymous.Token != "" {
|
||||
t.Fatalf("anonymous plugin ownership = %#v", anonymous)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -13,9 +13,9 @@
|
||||
|
||||
package app
|
||||
|
||||
// MCPIdentityHeaders returns the same header map used for MCP HTTP requests
|
||||
// (agent identity, env trace headers, edition MergeHeaders). Intended for
|
||||
// non-MCP transports such as the A2A gateway client.
|
||||
// MCPIdentityHeaders returns the shared identity header map used by non-MCP
|
||||
// transports such as the A2A gateway client. MCP-only Agent version and
|
||||
// extension metadata are intentionally excluded.
|
||||
func MCPIdentityHeaders() map[string]string {
|
||||
return resolveIdentityHeaders()
|
||||
}
|
||||
|
||||
+128
-24
@@ -79,6 +79,7 @@ var (
|
||||
rootPluginSyncSkills = plugin.SyncSkills
|
||||
rootAuthLoadTokenData = authpkg.LoadTokenData
|
||||
rootNewCommandRunnerWithFlags = newCommandRunnerWithFlags
|
||||
rootEmitResult = output.EmitResult
|
||||
)
|
||||
|
||||
// Execute runs the root command and returns the process exit code.
|
||||
@@ -136,6 +137,16 @@ func Execute() (exitCode int) {
|
||||
restoreArgs := rootNormalizeProcessProfileArgs()
|
||||
defer restoreArgs()
|
||||
|
||||
// Validate MCP Agent metadata before constructing the command tree.
|
||||
// Construction may invoke edition registration/static-server hooks and load
|
||||
// plugin PreParse handlers, so PersistentPreRunE alone is too late for the
|
||||
// process entry point. Retain this exact pair for the eventual invocation.
|
||||
agentMetadata := readAgentMetadataSnapshot()
|
||||
if err := agentMetadata.validationError(); err != nil {
|
||||
emitEarlyAgentMetadataValidationError(err, os.Args[1:])
|
||||
return apperrors.ExitCode(err)
|
||||
}
|
||||
|
||||
timing := NewTimingCollector()
|
||||
defer func() {
|
||||
rootStopAllStdioClients() // Ensure child processes are terminated on exit
|
||||
@@ -147,6 +158,7 @@ func Execute() (exitCode int) {
|
||||
|
||||
// Attach timing collector to context for use by child components
|
||||
ctx := WithTimingCollector(context.Background(), timing)
|
||||
ctx = contextWithAgentMetadataSnapshot(ctx, agentMetadata)
|
||||
ctx, resultStore = output.WithResultStore(ctx)
|
||||
var signalState *processSignalState
|
||||
var stopSignals func()
|
||||
@@ -260,6 +272,70 @@ func Execute() (exitCode int) {
|
||||
return 0
|
||||
}
|
||||
|
||||
// emitEarlyAgentMetadataValidationError preserves each built-in command's
|
||||
// legacy-vs-unified output contract without running extension hooks. The
|
||||
// presentation-only tree contains reviewed open-source commands and flags but
|
||||
// deliberately omits edition registration, plugin loading, and visibility
|
||||
// hooks; callers therefore still fail before any external hook executes.
|
||||
func emitEarlyAgentMetadataValidationError(err error, args []string) {
|
||||
format := processArgsFormat(args)
|
||||
presentationRoot := newRootPresentationCommand()
|
||||
_ = presentationRoot.PersistentFlags().Set("format", format)
|
||||
if target, _, findErr := presentationRoot.Find(args); findErr == nil && target != nil && output.UsesUnifiedResult(target) {
|
||||
target.SetOut(os.Stdout)
|
||||
target.SetErr(os.Stderr)
|
||||
result := output.FailureWithExitCode(errorInfoFromExecutionError(err), apperrors.ExitCode(err))
|
||||
if _, emitErr := rootEmitResult(target, result); emitErr == nil {
|
||||
return
|
||||
}
|
||||
}
|
||||
if strings.EqualFold(strings.TrimSpace(format), "json") {
|
||||
_ = apperrors.PrintJSON(os.Stderr, err)
|
||||
return
|
||||
}
|
||||
_ = apperrors.PrintHumanAt(os.Stderr, err, apperrors.VerbosityNormal)
|
||||
}
|
||||
|
||||
// processArgsRequestJSON preserves the CLI's machine-readable error contract
|
||||
// for validation that must occur before Cobra and its presentation flags exist.
|
||||
// The global format defaults to JSON; an explicit non-JSON format switches to
|
||||
// the human diagnostic path. Last occurrence wins, matching pflag semantics.
|
||||
func processArgsRequestJSON(args []string) bool {
|
||||
return strings.EqualFold(strings.TrimSpace(processArgsFormat(args)), "json")
|
||||
}
|
||||
|
||||
func processArgsFormat(args []string) string {
|
||||
format := "json"
|
||||
for index := 0; index < len(args); index++ {
|
||||
arg := args[index]
|
||||
if arg == "--" {
|
||||
break
|
||||
}
|
||||
if value, ok := strings.CutPrefix(arg, "--format="); ok {
|
||||
format = value
|
||||
continue
|
||||
}
|
||||
if value, ok := strings.CutPrefix(arg, "-f="); ok {
|
||||
format = value
|
||||
continue
|
||||
}
|
||||
if strings.HasPrefix(arg, "-f") && len(arg) > len("-f") {
|
||||
format = strings.TrimPrefix(arg, "-f")
|
||||
continue
|
||||
}
|
||||
if arg != "--format" && arg != "-f" {
|
||||
continue
|
||||
}
|
||||
if index+1 >= len(args) {
|
||||
format = ""
|
||||
break
|
||||
}
|
||||
index++
|
||||
format = args[index]
|
||||
}
|
||||
return format
|
||||
}
|
||||
|
||||
// errorInfoFromExecutionError projects the repository error model into the unified
|
||||
// failure body. Exit code and category are derived from the same error value,
|
||||
// preventing the wire and process status from drifting apart.
|
||||
@@ -633,12 +709,25 @@ func NewRootCommandWithEngine(rootCtx context.Context, engine *pipeline.Engine)
|
||||
}
|
||||
|
||||
func newRootCommandWithEngine(rootCtx context.Context, engine *pipeline.Engine, loadRuntimeExtensions bool, declarationOnly bool) *cobra.Command {
|
||||
return newRootCommandWithMode(rootCtx, engine, loadRuntimeExtensions, declarationOnly, false)
|
||||
}
|
||||
|
||||
func newRootPresentationCommand() *cobra.Command {
|
||||
return newRootCommandWithMode(context.Background(), nil, false, true, true)
|
||||
}
|
||||
|
||||
func newRootCommandWithMode(rootCtx context.Context, engine *pipeline.Engine, loadRuntimeExtensions bool, declarationOnly bool, presentationOnly bool) *cobra.Command {
|
||||
if rootCtx == nil {
|
||||
rootCtx = context.Background()
|
||||
}
|
||||
flags := &GlobalFlags{}
|
||||
authpkg.SetRuntimeProfile(preparseProfileFlag(os.Args[1:]))
|
||||
runner := rootNewCommandRunnerWithFlags(flags)
|
||||
if snapshot, ok := agentMetadataSnapshotFromContext(rootCtx); ok {
|
||||
if runtime, ok := runner.(*runtimeRunner); ok {
|
||||
runtime.agentMetadata = &snapshot
|
||||
}
|
||||
}
|
||||
|
||||
root := &cobra.Command{
|
||||
Use: "dws",
|
||||
@@ -672,15 +761,29 @@ func newRootCommandWithEngine(rootCtx context.Context, engine *pipeline.Engine,
|
||||
// --output sink instead opens at Run entry (after Cobra's own
|
||||
// validation), so validation failures still cannot strand a
|
||||
// temporary file.
|
||||
// Validate caller-provided identity labels before any edition hook
|
||||
// or command network activity can run. Header-only library callers
|
||||
// use the best-effort path in resolveIdentityHeaders instead.
|
||||
// Validate caller-provided identity and MCP metadata before command
|
||||
// execution hooks or network activity. The process entry point additionally
|
||||
// validates Agent metadata before command-tree construction; direct Cobra
|
||||
// embedding retains this execution-boundary guard.
|
||||
if _, err := parseAgentHost(os.Getenv(envDWSAgentHost)); err != nil {
|
||||
return err
|
||||
}
|
||||
if _, err := parseAgentProduct(os.Getenv(agentproduct.EnvName)); err != nil {
|
||||
return err
|
||||
}
|
||||
agentMetadata, cached := agentMetadataSnapshotFromContext(cmd.Context())
|
||||
if !cached {
|
||||
agentMetadata = readAgentMetadataSnapshot()
|
||||
}
|
||||
if err := agentMetadata.validationError(); err != nil {
|
||||
return err
|
||||
}
|
||||
if runtime, ok := runner.(*runtimeRunner); ok {
|
||||
// Retain the exact validated pair for this command execution so a
|
||||
// concurrently mutating embedding environment cannot change what is
|
||||
// later applied after edition and credential hooks.
|
||||
runtime.agentMetadata = &agentMetadata
|
||||
}
|
||||
if shouldDetectNestedSkillLayout(cmd) {
|
||||
if found, err := detectNestedMultiSkillLayout(); err == nil && found {
|
||||
fmt.Fprintln(cmd.ErrOrStderr(), "⚠️ 检测到旧升级器留下的嵌套 Skill;请运行 dws skill setup --mode multi 查看迁移计划并确认")
|
||||
@@ -777,10 +880,12 @@ func newRootCommandWithEngine(rootCtx context.Context, engine *pipeline.Engine,
|
||||
// PAT authorization commands (open-source core)
|
||||
pat.RegisterCommands(root, patCaller)
|
||||
|
||||
if fn := edition.Get().RegisterExtraCommands; fn != nil {
|
||||
caller := newToolCallerAdapter(runner, flags)
|
||||
fn(root, caller)
|
||||
deduplicateCommands(root)
|
||||
if !presentationOnly {
|
||||
if fn := edition.Get().RegisterExtraCommands; fn != nil {
|
||||
caller := newToolCallerAdapter(runner, flags)
|
||||
fn(root, caller)
|
||||
deduplicateCommands(root)
|
||||
}
|
||||
}
|
||||
if loadRuntimeExtensions {
|
||||
// Resolve plugins only after the complete distribution command tree is
|
||||
@@ -791,7 +896,9 @@ func newRootCommandWithEngine(rootCtx context.Context, engine *pipeline.Engine,
|
||||
addPluginCommandsSafe(root, pluginCmds)
|
||||
}
|
||||
}
|
||||
hideNonDirectRuntimeCommands(root)
|
||||
if !presentationOnly {
|
||||
hideNonDirectRuntimeCommands(root)
|
||||
}
|
||||
configureRootHelp(root)
|
||||
// Set custom flag error handler for better UX
|
||||
root.SetFlagErrorFunc(flagErrorWithSuggestions)
|
||||
@@ -1755,17 +1862,16 @@ func distributionRootOwns(root *cobra.Command, name string) bool {
|
||||
func registerPluginHTTPServer(srv mcptypes.ServerDescriptor) {
|
||||
AppendDynamicServer(srv)
|
||||
productID := firstNonEmptyPluginString(srv.CLI.ID, srv.Key)
|
||||
ClearPluginAuth(productID)
|
||||
if len(srv.AuthHeaders) > 0 {
|
||||
registerPluginAuthFromHeaders(srv)
|
||||
}
|
||||
// Register ownership for every accepted HTTP plugin, including anonymous
|
||||
// plugins. Execution must never fall back to the built-in DingTalk OAuth or
|
||||
// Agent-metadata path merely because a plugin has no Authorization Header.
|
||||
RegisterPluginAuth(productID, pluginAuthFromServerDescriptor(srv))
|
||||
}
|
||||
|
||||
// registerPluginAuthFromHeaders extracts authentication credentials from
|
||||
// a server descriptor's AuthHeaders and registers them in the global
|
||||
// PluginAuth registry. The runner uses this registry at execution time
|
||||
// to inject the correct Bearer token for third-party MCP servers.
|
||||
func registerPluginAuthFromHeaders(srv mcptypes.ServerDescriptor) {
|
||||
// pluginAuthFromServerDescriptor extracts plugin-owned credentials and custom
|
||||
// Headers. A non-nil result also acts as the HTTP plugin ownership marker for
|
||||
// anonymous plugins.
|
||||
func pluginAuthFromServerDescriptor(srv mcptypes.ServerDescriptor) *PluginAuth {
|
||||
authToken := ""
|
||||
extraHeaders := make(map[string]string)
|
||||
for key, value := range srv.AuthHeaders {
|
||||
@@ -1776,20 +1882,18 @@ func registerPluginAuthFromHeaders(srv mcptypes.ServerDescriptor) {
|
||||
extraHeaders[key] = value
|
||||
}
|
||||
}
|
||||
if authToken == "" {
|
||||
return
|
||||
}
|
||||
var trustedDomains []string
|
||||
if parsed, err := url.Parse(srv.Endpoint); err == nil {
|
||||
host := parsed.Hostname()
|
||||
trustedDomains = []string{host, "*." + host}
|
||||
if host != "" {
|
||||
trustedDomains = []string{host, "*." + host}
|
||||
}
|
||||
}
|
||||
productID := firstNonEmptyPluginString(srv.CLI.ID, srv.Key)
|
||||
RegisterPluginAuth(productID, &PluginAuth{
|
||||
return &PluginAuth{
|
||||
Token: authToken,
|
||||
ExtraHeaders: extraHeaders,
|
||||
TrustedDomains: trustedDomains,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// newPipelineEngine creates and configures the pipeline engine with
|
||||
|
||||
+26
-13
@@ -137,7 +137,6 @@ func newCommandRunnerWithFlags(flags *GlobalFlags) executor.Runner {
|
||||
httpClient = &http.Client{Timeout: time.Duration(flags.Timeout) * time.Second}
|
||||
}
|
||||
transportClient := transport.NewClient(httpClient)
|
||||
transportClient.ExtraHeaders = resolveIdentityHeaders()
|
||||
transportClient.FileLogger = FileLoggerInstance()
|
||||
return &runtimeRunner{
|
||||
transport: transportClient,
|
||||
@@ -156,12 +155,14 @@ type runtimeRunner struct {
|
||||
enforceContentScan bool
|
||||
includeScanReport bool
|
||||
auditSink audit.Sink
|
||||
agentMetadata *agentMetadataSnapshot
|
||||
}
|
||||
|
||||
var (
|
||||
runnerResolveMultiProfileSelections = resolveMultiProfileSelections
|
||||
runnerResolveProfile = authpkg.ResolveProfile
|
||||
runnerGetCachedRuntimeToken = getCachedRuntimeToken
|
||||
runnerResolveAuthToken = (*runtimeRunner).resolveAuthToken
|
||||
runnerPreflightDocDownload = (*runtimeRunner).preflightDocDownload
|
||||
runnerCallTool = (*transport.Client).CallTool
|
||||
runnerStdioEnsureInitialized = (*transport.StdioClient).EnsureInitialized
|
||||
@@ -237,7 +238,6 @@ func (r *runtimeRunner) runSingle(ctx context.Context, invocation executor.Invoc
|
||||
if r.transport == nil {
|
||||
return r.fallback.Run(ctx, invocation)
|
||||
}
|
||||
r.transport.ExtraHeaders = resolveIdentityHeaders()
|
||||
|
||||
// Mock mode: skip endpoint resolution, use a placeholder endpoint.
|
||||
if r.globalFlags != nil && r.globalFlags.Mock {
|
||||
@@ -543,9 +543,9 @@ func (r *runtimeRunner) executeInvocation(ctx context.Context, endpoint string,
|
||||
emitAudit(auditSink, execID, invokeStart, invocation, endpoint, retErr, version)
|
||||
}()
|
||||
|
||||
// Check if this product has plugin-level auth credentials registered.
|
||||
// If so, use the plugin's token instead of the default DingTalk OAuth token.
|
||||
// This allows third-party MCP servers (e.g. Bailian) to use their own API keys.
|
||||
// Check whether this product belongs to an HTTP plugin. Every accepted
|
||||
// plugin has an ownership record; credentials within that record are
|
||||
// optional. Plugin requests never fall back to the default DingTalk OAuth.
|
||||
pluginAuth, hasPluginAuth := LookupPluginAuth(invocation.CanonicalProduct)
|
||||
|
||||
authToken := ""
|
||||
@@ -553,7 +553,7 @@ func (r *runtimeRunner) executeInvocation(ctx context.Context, endpoint string,
|
||||
authToken = pluginAuth.Token
|
||||
} else if !invocation.DryRun && (r.globalFlags == nil || !r.globalFlags.Mock) {
|
||||
var tokenErr error
|
||||
authToken, tokenErr = r.resolveAuthToken(ctx)
|
||||
authToken, tokenErr = runnerResolveAuthToken(r, ctx)
|
||||
if tokenErr != nil {
|
||||
return executor.Result{}, tokenResolutionError(tokenErr)
|
||||
}
|
||||
@@ -603,9 +603,10 @@ func (r *runtimeRunner) executeInvocation(ctx context.Context, endpoint string,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// Fail-fast: reject unauthenticated requests before making network calls.
|
||||
// This provides a clear error message instead of cryptic HTTP 400 from MCP.
|
||||
if strings.TrimSpace(authToken) == "" {
|
||||
// Preserve a final execution-boundary guard even though the built-in token
|
||||
// resolver normally returns either a non-empty token or an error. HTTP
|
||||
// plugins are ownership-scoped separately and may intentionally be anonymous.
|
||||
if !hasPluginAuth && strings.TrimSpace(authToken) == "" {
|
||||
return executor.Result{}, apperrors.NewAuth(
|
||||
"未登录,请先执行 dws auth login",
|
||||
apperrors.WithReason("not_authenticated"),
|
||||
@@ -616,12 +617,20 @@ func (r *runtimeRunner) executeInvocation(ctx context.Context, endpoint string,
|
||||
|
||||
var tc *transport.Client
|
||||
if hasPluginAuth {
|
||||
// Use plugin-level auth: inject the plugin's token and trust its domains.
|
||||
tc = r.transport.WithAuth(authToken, pluginAuth.ExtraHeaders)
|
||||
// Plugin ownership is authoritative even when the plugin is anonymous.
|
||||
// Copy and sanitize manifest headers so plugins cannot opt themselves into
|
||||
// DWS-owned Agent metadata by declaring the reserved names directly.
|
||||
tc = r.transport.WithAuth(authToken, pluginRequestHeaders(pluginAuth))
|
||||
tc.TrustedDomains = pluginAuth.TrustedDomains
|
||||
} else {
|
||||
// Default path: use DingTalk OAuth token with identity headers.
|
||||
tc = r.transport.WithAuth(authToken, resolveIdentityHeaders())
|
||||
// Default path: use DingTalk OAuth token with identity headers. Agent
|
||||
// metadata is resolved exactly once per invocation and is excluded from
|
||||
// helper-only service-discovery requests.
|
||||
if r.agentMetadata != nil {
|
||||
tc = r.transport.WithAuth(authToken, resolveMCPRequestHeadersForInvocation(invocation, *r.agentMetadata))
|
||||
} else {
|
||||
tc = r.transport.WithAuth(authToken, resolveMCPRequestHeadersForInvocation(invocation))
|
||||
}
|
||||
}
|
||||
|
||||
callCtx := ctx
|
||||
@@ -1069,6 +1078,10 @@ func resolveIdentityHeaders() map[string]string {
|
||||
} else {
|
||||
delete(headers, agentproduct.HeaderName)
|
||||
}
|
||||
// Agent version and extension are intentionally MCP-request-only. Remove
|
||||
// every case variant potentially supplied by an edition or credential hook
|
||||
// so shared consumers such as A2A cannot inherit them.
|
||||
removeAgentMetadataHeaders(headers)
|
||||
return headers
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,80 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package auth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"io"
|
||||
"net/http"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
type agentMetadataHeaderRoundTripFunc func(*http.Request) (*http.Response, error)
|
||||
|
||||
func (f agentMetadataHeaderRoundTripFunc) RoundTrip(req *http.Request) (*http.Response, error) {
|
||||
return f(req)
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageOAuthRequestsExcludeAgentMetadataHeaders(t *testing.T) {
|
||||
t.Setenv("DWS_AGENT_VER", "1.2.3-test")
|
||||
t.Setenv("DWS_AGENT_EXT", `{"umt":"test-umt","miniwua":"test-wua","ua":"test-agent"}`)
|
||||
|
||||
assertExcluded := func(t *testing.T, header http.Header) {
|
||||
t.Helper()
|
||||
for _, name := range []string{"x-dws-agent-ver", "x-dws-agent-ext"} {
|
||||
if values := header.Values(name); len(values) != 0 {
|
||||
t.Fatalf("OAuth request header %q = %q, want absent", name, values)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
newCaptureClient := func(responseBody string) (*http.Client, <-chan http.Header) {
|
||||
seen := make(chan http.Header, 1)
|
||||
client := &http.Client{Transport: agentMetadataHeaderRoundTripFunc(func(req *http.Request) (*http.Response, error) {
|
||||
seen <- req.Header.Clone()
|
||||
return &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Header: http.Header{"Content-Type": {"application/json"}},
|
||||
Body: io.NopCloser(strings.NewReader(responseBody)),
|
||||
Request: req,
|
||||
}, nil
|
||||
})}
|
||||
return client, seen
|
||||
}
|
||||
|
||||
t.Run("JSON token request", func(t *testing.T) {
|
||||
client, seen := newCaptureClient(`{}`)
|
||||
provider := &OAuthProvider{httpClient: client}
|
||||
if _, err := provider.postJSON(context.Background(), "https://oauth.test/token", map[string]string{
|
||||
"code": "test-code",
|
||||
"grantType": "authorization_code",
|
||||
}); err != nil {
|
||||
t.Fatalf("postJSON() error = %v", err)
|
||||
}
|
||||
assertExcluded(t, <-seen)
|
||||
})
|
||||
|
||||
t.Run("device code request", func(t *testing.T) {
|
||||
client, seen := newCaptureClient(`{"success":true,"result":{"deviceCode":"test-device-code","userCode":"TEST-CODE","verificationUri":"https://example.test/device","expiresIn":900,"interval":1}}`)
|
||||
provider := NewDeviceFlowProvider(t.TempDir(), newDeviceFlowTestLogger())
|
||||
provider.clientID = "test-client-id"
|
||||
provider.httpClient = client
|
||||
provider.SetBaseURL("https://oauth.test")
|
||||
if _, err := provider.requestDeviceCode(context.Background()); err != nil {
|
||||
t.Fatalf("requestDeviceCode() error = %v", err)
|
||||
}
|
||||
assertExcluded(t, <-seen)
|
||||
})
|
||||
}
|
||||
@@ -26,6 +26,8 @@ import (
|
||||
var sensitiveKeys = map[string]bool{
|
||||
"authorization": true,
|
||||
"x-user-access-token": true,
|
||||
"dws_agent_ext": true,
|
||||
"x-dws-agent-ext": true,
|
||||
"client_secret": true,
|
||||
"client-secret": true,
|
||||
"token": true,
|
||||
|
||||
@@ -30,6 +30,12 @@ func TestIsSensitiveKey(t *testing.T) {
|
||||
{"authorization", true},
|
||||
{"x-user-access-token", true},
|
||||
{"X-User-Access-Token", true},
|
||||
{"DWS_AGENT_EXT", true},
|
||||
{"dws_agent_ext", true},
|
||||
{"x-dws-agent-ext", true},
|
||||
{"X-Dws-Agent-Ext", true},
|
||||
{"DWS_AGENT_VER", false},
|
||||
{"x-dws-agent-ver", false},
|
||||
{"client_secret", true},
|
||||
{"client-secret", true},
|
||||
{"token", true},
|
||||
@@ -128,17 +134,24 @@ func TestSanitizeArguments_Empty(t *testing.T) {
|
||||
|
||||
func TestRedactHeaders(t *testing.T) {
|
||||
t.Parallel()
|
||||
headers := http.Header{
|
||||
"Authorization": {"Bearer token123456"},
|
||||
"Content-Type": {"application/json"},
|
||||
}
|
||||
headers := make(http.Header)
|
||||
headers.Set("Authorization", "Bearer test-credential-value")
|
||||
headers.Set("Content-Type", "application/json")
|
||||
headers.Set("DWS_AGENT_EXT", `{"umt":"test-umt-value"}`)
|
||||
headers.Set("x-dws-agent-ext", `{"ua":"test-agent-value"}`)
|
||||
attrs := RedactHeaders(headers)
|
||||
if len(attrs) != 2 {
|
||||
t.Fatalf("expected 2 attrs, got %d", len(attrs))
|
||||
if len(attrs) != 4 {
|
||||
t.Fatalf("expected 4 attrs, got %d", len(attrs))
|
||||
}
|
||||
for _, attr := range attrs {
|
||||
if attr.Key == "header.authorization" && !strings.Contains(attr.Value.String(), "***") {
|
||||
t.Fatalf("authorization should be redacted: %s", attr.Value.String())
|
||||
switch attr.Key {
|
||||
case "header.authorization", "header.dws_agent_ext", "header.x-dws-agent-ext":
|
||||
if !strings.Contains(attr.Value.String(), "***") {
|
||||
t.Fatalf("%s should be redacted: %s", attr.Key, attr.Value.String())
|
||||
}
|
||||
if strings.Contains(attr.Value.String(), "test-") {
|
||||
t.Fatalf("%s leaked its original value: %s", attr.Key, attr.Value.String())
|
||||
}
|
||||
}
|
||||
if attr.Key == "header.content-type" && attr.Value.String() != "application/json" {
|
||||
t.Fatalf("content-type should not be redacted: %s", attr.Value.String())
|
||||
|
||||
@@ -0,0 +1,93 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package transport
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestCrossPlatformCoverageAgentMetadataReachesEveryMCPMethod(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
const (
|
||||
agentVersion = "1.2.3-test+7"
|
||||
agentExt = `{"umt":"test-umt","nested":{"enabled":true}}`
|
||||
cliVersion = "9.8.7-cli"
|
||||
userAgent = "dws-user-agent-test/1.0"
|
||||
)
|
||||
|
||||
seen := make(map[string]int)
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
var request requestEnvelope
|
||||
if err := json.NewDecoder(r.Body).Decode(&request); err != nil {
|
||||
http.Error(w, err.Error(), http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
seen[request.Method]++
|
||||
if got := r.Header.Get(HeaderAgentVersion); got != agentVersion {
|
||||
t.Errorf("%s %s = %q, want %q", request.Method, HeaderAgentVersion, got, agentVersion)
|
||||
}
|
||||
if got := r.Header.Get(HeaderAgentExt); got != agentExt {
|
||||
t.Errorf("%s %s = %q, want %q", request.Method, HeaderAgentExt, got, agentExt)
|
||||
}
|
||||
if got := r.Header.Get(HeaderVersion); got != cliVersion {
|
||||
t.Errorf("%s %s = %q, want %q", request.Method, HeaderVersion, got, cliVersion)
|
||||
}
|
||||
if got := r.Header.Get("User-Agent"); got != userAgent {
|
||||
t.Errorf("%s User-Agent = %q, want preset value %q", request.Method, got, userAgent)
|
||||
}
|
||||
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
switch request.Method {
|
||||
case "initialize":
|
||||
_, _ = fmt.Fprintf(w, `{"jsonrpc":"2.0","id":1,"result":{"protocolVersion":"%s","capabilities":{}}}`, supportedProtocolVersions[0])
|
||||
case "tools/list":
|
||||
_, _ = fmt.Fprint(w, `{"jsonrpc":"2.0","id":2,"result":{"tools":[]}}`)
|
||||
case "tools/call":
|
||||
_, _ = fmt.Fprint(w, `{"jsonrpc":"2.0","id":3,"result":{"content":{"success":true}}}`)
|
||||
default:
|
||||
http.Error(w, "unexpected method", http.StatusBadRequest)
|
||||
}
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
client := NewClient(server.Client())
|
||||
client.ExtraHeaders = map[string]string{
|
||||
HeaderAgentVersion: agentVersion,
|
||||
HeaderAgentExt: agentExt,
|
||||
HeaderVersion: cliVersion,
|
||||
"User-Agent": userAgent,
|
||||
}
|
||||
ctx := context.Background()
|
||||
if _, err := client.Initialize(ctx, server.URL); err != nil {
|
||||
t.Fatalf("Initialize() error = %v", err)
|
||||
}
|
||||
if _, err := client.ListTools(ctx, server.URL); err != nil {
|
||||
t.Fatalf("ListTools() error = %v", err)
|
||||
}
|
||||
if _, err := client.CallTool(ctx, server.URL, "test_tool", map[string]any{"value": "safe"}); err != nil {
|
||||
t.Fatalf("CallTool() error = %v", err)
|
||||
}
|
||||
|
||||
for _, method := range []string{"initialize", "tools/list", "tools/call"} {
|
||||
if got := seen[method]; got != 1 {
|
||||
t.Errorf("%s request count = %d, want 1", method, got)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -68,10 +68,12 @@ const (
|
||||
defaultRetryMaxDelay = 5 * time.Second
|
||||
|
||||
// Security headers
|
||||
HeaderSource = "X-Cli-Source"
|
||||
HeaderVersion = "X-Cli-Version"
|
||||
HeaderExecutionId = "X-Cli-Execution-Id"
|
||||
SourceValue = "dws-cli"
|
||||
HeaderSource = "X-Cli-Source"
|
||||
HeaderVersion = "X-Cli-Version"
|
||||
HeaderExecutionId = "X-Cli-Execution-Id"
|
||||
HeaderAgentVersion = "x-dws-agent-ver"
|
||||
HeaderAgentExt = "x-dws-agent-ext"
|
||||
SourceValue = "dws-cli"
|
||||
)
|
||||
|
||||
// Supported MCP protocol versions, ordered from newest to oldest.
|
||||
@@ -256,20 +258,43 @@ func NewClient(httpClient *http.Client) *Client {
|
||||
|
||||
// safeRedirectPolicy prevents credential headers from being forwarded
|
||||
// when a response redirects to a different host (e.g. API 302 → CDN).
|
||||
// Strips Authorization, x-user-access-token on cross-host redirects;
|
||||
// other headers like X-Cli-* pass through.
|
||||
// Credentials and Agent extension context are bound to the initial origin:
|
||||
// once a redirect chain leaves that origin they remain stripped for every
|
||||
// subsequent hop, including a redirect back to the initial origin.
|
||||
// Non-sensitive headers like X-Cli-* and x-dws-agent-ver pass through.
|
||||
func safeRedirectPolicy(req *http.Request, via []*http.Request) error {
|
||||
if len(via) >= 10 {
|
||||
return fmt.Errorf("too many redirects")
|
||||
}
|
||||
if len(via) > 0 && req.URL.Host != via[0].URL.Host {
|
||||
// Cross-host redirect: strip sensitive headers to prevent credential leakage
|
||||
if redirectChainLeftInitialOrigin(req, via) {
|
||||
req.Header.Del(HeaderAgentExt)
|
||||
req.Header.Del("Authorization")
|
||||
req.Header.Del("x-user-access-token")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func redirectChainLeftInitialOrigin(req *http.Request, via []*http.Request) bool {
|
||||
if len(via) == 0 {
|
||||
return false
|
||||
}
|
||||
initialURL := via[0].URL
|
||||
if !sameOrigin(req.URL, initialURL) {
|
||||
return true
|
||||
}
|
||||
for _, previous := range via[1:] {
|
||||
if !sameOrigin(previous.URL, initialURL) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func sameOrigin(left, right *url.URL) bool {
|
||||
return strings.EqualFold(left.Scheme, right.Scheme) &&
|
||||
strings.EqualFold(left.Host, right.Host)
|
||||
}
|
||||
|
||||
// WithAuth returns a shallow copy of c with the given auth token and extra
|
||||
// headers. The returned client shares the underlying HTTP client but is safe
|
||||
// to use concurrently with the original.
|
||||
|
||||
@@ -0,0 +1,140 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package transport
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestCrossPlatformCoverageSafeRedirectPolicyAgentMetadataHeaders(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
newRequest := func(t *testing.T, rawURL string) *http.Request {
|
||||
t.Helper()
|
||||
req, err := http.NewRequest(http.MethodGet, rawURL, nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
req.Header.Set(HeaderAgentVersion, "1.2.3-test")
|
||||
req.Header.Set(HeaderAgentExt, `{"ua":"test-agent-value"}`)
|
||||
req.Header.Set("Authorization", "Bearer test-token")
|
||||
req.Header.Set("x-user-access-token", "test-token")
|
||||
return req
|
||||
}
|
||||
|
||||
assertHeaders := func(t *testing.T, request *http.Request, wantExt bool) {
|
||||
t.Helper()
|
||||
if got := request.Header.Get(HeaderAgentVersion); got != "1.2.3-test" {
|
||||
t.Fatalf("agent version = %q, want retained", got)
|
||||
}
|
||||
gotExt := request.Header.Get(HeaderAgentExt)
|
||||
if wantExt && gotExt == "" {
|
||||
t.Fatal("agent extension was removed")
|
||||
}
|
||||
if !wantExt && gotExt != "" {
|
||||
t.Fatalf("agent extension leaked across origins: %q", gotExt)
|
||||
}
|
||||
for _, key := range []string{"Authorization", "x-user-access-token"} {
|
||||
got := request.Header.Get(key)
|
||||
if wantExt && got == "" {
|
||||
t.Fatalf("same-origin redirect removed %s", key)
|
||||
}
|
||||
if !wantExt && got != "" {
|
||||
t.Fatalf("credential %s leaked across origins", key)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
t.Run("same origin retains version and extension", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
previous := newRequest(t, "https://api.example.test/start")
|
||||
redirected := newRequest(t, "https://api.example.test/next")
|
||||
|
||||
if err := safeRedirectPolicy(redirected, []*http.Request{previous}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
assertHeaders(t, redirected, true)
|
||||
})
|
||||
|
||||
t.Run("cross host strips extension and retains version", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
previous := newRequest(t, "https://api.example.test/start")
|
||||
redirected := newRequest(t, "https://cdn.example.test/asset")
|
||||
|
||||
if err := safeRedirectPolicy(redirected, []*http.Request{previous}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
assertHeaders(t, redirected, false)
|
||||
})
|
||||
|
||||
t.Run("scheme downgrade on same host strips extension", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
previous := newRequest(t, "https://api.example.test/start")
|
||||
redirected := newRequest(t, "http://api.example.test/next")
|
||||
|
||||
if err := safeRedirectPolicy(redirected, []*http.Request{previous}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
assertHeaders(t, redirected, false)
|
||||
})
|
||||
|
||||
t.Run("extension remains stripped after returning to initial origin", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
initial := newRequest(t, "https://api.example.test/start")
|
||||
crossOrigin := newRequest(t, "https://cdn.example.test/asset")
|
||||
redirectedBack := newRequest(t, "https://api.example.test/final")
|
||||
|
||||
if err := safeRedirectPolicy(redirectedBack, []*http.Request{initial, crossOrigin}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
assertHeaders(t, redirectedBack, false)
|
||||
})
|
||||
|
||||
t.Run("extension remains stripped on later cross-origin hop", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
initial := newRequest(t, "https://api.example.test/start")
|
||||
crossOrigin := newRequest(t, "https://cdn.example.test/asset")
|
||||
redirected := newRequest(t, "https://cdn.example.test/final")
|
||||
|
||||
if err := safeRedirectPolicy(redirected, []*http.Request{initial, crossOrigin}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
assertHeaders(t, redirected, false)
|
||||
})
|
||||
|
||||
t.Run("initial request is unchanged", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
request := newRequest(t, "https://api.example.test/start")
|
||||
|
||||
if err := safeRedirectPolicy(request, nil); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
assertHeaders(t, request, true)
|
||||
})
|
||||
|
||||
t.Run("redirect limit is enforced", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
request := newRequest(t, "https://api.example.test/final")
|
||||
via := make([]*http.Request, 10)
|
||||
for index := range via {
|
||||
via[index] = newRequest(t, "https://api.example.test/"+strings.Repeat("x", index+1))
|
||||
}
|
||||
|
||||
if err := safeRedirectPolicy(request, via); err == nil {
|
||||
t.Fatal("safeRedirectPolicy() error = nil, want redirect limit error")
|
||||
}
|
||||
})
|
||||
}
|
||||
+3
-2
@@ -15,8 +15,9 @@ package cli
|
||||
|
||||
import "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/app"
|
||||
|
||||
// MCPIdentityHeaders returns HTTP headers aligned with MCP tool calls
|
||||
// (identity + edition merge). Overlays may pass this to auxiliary clients.
|
||||
// MCPIdentityHeaders returns the shared identity and edition headers that
|
||||
// overlays may pass to auxiliary clients. MCP-request-only metadata such as
|
||||
// DWS_AGENT_VER and DWS_AGENT_EXT is intentionally excluded.
|
||||
func MCPIdentityHeaders() map[string]string {
|
||||
return app.MCPIdentityHeaders()
|
||||
}
|
||||
|
||||
@@ -0,0 +1,35 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package cli
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestCrossPlatformCoverageMCPIdentityHeadersExcludeMCPOnlyMetadata(t *testing.T) {
|
||||
t.Setenv("DWS_CONFIG_DIR", t.TempDir())
|
||||
t.Setenv("DWS_AGENT_VER", "1.2.3")
|
||||
t.Setenv("DWS_AGENT_EXT", `{"ua":"test-agent/1.2.3"}`)
|
||||
|
||||
headers := MCPIdentityHeaders()
|
||||
if headers == nil {
|
||||
t.Fatal("MCPIdentityHeaders() returned nil")
|
||||
}
|
||||
for key := range headers {
|
||||
if strings.EqualFold(key, "x-dws-agent-ver") || strings.EqualFold(key, "x-dws-agent-ext") {
|
||||
t.Fatalf("shared identity export leaked MCP-only Header %q", key)
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user