Compare commits

...
Author SHA1 Message Date
wxianfeng 0137a1950b fix(event): harden personal stream lifecycle 2026-07-10 14:27:17 +08:00
wxianfeng ed0ec3103a Merge remote-tracking branch 'upstream/main' into feat/dws-event 2026-07-10 11:10:03 +08:00
wxianfeng 8f0bb91ddb chore(event): prepare official release 2026-07-10 11:09:47 +08:00
wxianfeng b5be003330 remove external event reference comments 2026-07-09 22:05:01 +08:00
wxianfeng dd29805063 hide incomplete personal sender event 2026-07-09 21:43:20 +08:00
wxianfeng 497b627080 hide app event public entrypoints 2026-07-09 21:11:08 +08:00
wxianfeng 4233105963 fix event stop and status visibility 2026-07-09 18:24:50 +08:00
wxianfeng 2541081f97 feat: align personal event consume flags 2026-07-09 16:38:25 +08:00
wxianfeng 3a6e5fb4b3 docs: refine dingtalk event skill references 2026-07-09 15:47:21 +08:00
wxianfeng 83783972a3 feat: simplify event schema output 2026-07-09 15:11:46 +08:00
wxianfeng 919b9d698f Merge remote-tracking branch 'upstream/main' into feat/dws-event 2026-07-09 14:24:36 +08:00
wxianfeng ea77270030 Merge remote-tracking branch 'upstream/main' into feat/dws-event 2026-07-08 16:26:11 +08:00
wxianfeng 1a62ef3616 feat: simplify personal event schemas 2026-07-08 15:54:21 +08:00
wxianfeng 7d37c15a46 fix: avoid duplicate app helper name 2026-07-08 14:54:30 +08:00
wxianfeng 9fa763e213 Merge remote-tracking branch 'upstream/main' into feat/dws-event
# Conflicts:
#	skills/mono/SKILL.md
2026-07-08 14:46:36 +08:00
wxianfeng ee0e629c72 fix: align personal event schema with stream payload 2026-07-07 22:04:17 +08:00
wxianfeng 2c8f970703 filter subId 2026-07-07 21:25:18 +08:00
wxianfeng 20d0eaa489 more im event 2026-07-06 21:24:26 +08:00
xianfeng wang 67c88828f0 Merge pull request #25 from sczheng189/feat/dws-event
修复event 描述错误
2026-07-06 20:42:29 +08:00
zhengyubai b015c27034 fix(personal): 修正指定发送人消息描述错误 2026-07-06 21:27:37 +09:00
玉澜 604ec5f50a Merge remote-tracking branch 'origin/feat/dws-event' into feat/dws-event 2026-07-06 10:04:19 +08:00
玉澜 27296ec426 fix: remove subscribe id event fanout filter 2026-07-06 10:04:12 +08:00
wxianfeng 6a38a168dd install script event 2026-07-02 20:41:46 +08:00
wxianfeng 9771053d81 default value 2026-07-02 20:14:30 +08:00
wxianfeng 10c0c5083e event skill 2026-07-02 19:42:10 +08:00
wxianfeng a0187b5297 Merge branch 'feat/dws-event' of github.com:wxianfeng/dingtalk-workspace-cli into feat/dws-event 2026-07-02 17:15:48 +08:00
wxianfeng 81991f1c07 opt 2026-07-02 17:14:55 +08:00
xianfeng wang 836670ef50 Merge pull request #24 from sczheng189/feat/dws-event
fix(event): unix socket 路径超长时 fallback 到短路径,修复深层配置目录下 bus 无法启动
2026-07-02 16:54:58 +08:00
zhengyubai 53ce0a8303 refactor(event): 优化IPC端点路径处理和改进相关测试
- 用dwsevent.IPCEndpoint替代原先根据GOOS判断的路径逻辑
- 新增event包实现Unix socket路径长度限制及长路径fallback机制
- 添加endpoint_test.go覆盖路径短长及唯一性的单元测试
- 修改busctl模块使用统一的IPC端点获取方法,避免重复实现
- transport_unix.go新增checkSocketPath函数检查路径长度,防止EINVAL错误
- 在监听和连接Unix socket时加入路径限制检查,提升错误明晰度
- 去除多个文件中无用的runtime导入,简化代码依赖
2026-07-02 17:25:56 +09:00
wxianfeng 78867f3601 eventType filter 2026-07-02 16:19:28 +08:00
wxianfeng 37438659e6 user event 2026-07-01 15:58:09 +08:00
wxianfeng 3c12c835a3 Merge branch 'feat/dws-event' of github.com:wxianfeng/dingtalk-workspace-cli into feat/dws-event 2026-07-01 14:22:06 +08:00
wxianfeng 389f83241f event 2026-07-01 14:21:36 +08:00
玉澜 3714adc2db Merge remote-tracking branch 'origin/feat/dws-event' into feat/dws-event
# Conflicts:
#	internal/app/event_command.go
2026-07-01 14:20:17 +08:00
玉澜 f37d0569a1 fix: allow portal ticket normal without app secret 2026-07-01 14:17:12 +08:00
wxianfeng 798b58bf3c fix conflict 2026-07-01 11:15:48 +08:00
wxianfeng d926bed3cc user event 2026-07-01 11:07:57 +08:00
玉澜 5ac180d3dd feat: add portal ticket stream mode 2026-06-30 20:38:27 +08:00
玉澜 35e60407d3 test: add stream ticket injection probe 2026-06-30 15:33:17 +08:00
wxianfeng ea46132cf6 merge upstream main 2026-06-29 16:15:13 +08:00
wxianfeng 478dc155e8 fix consume fail 2026-06-04 10:41:50 +08:00
wxianfeng 08ecb38a42 dws event 2026-06-03 19:12:23 +08:00
114 changed files with 18987 additions and 182 deletions
+45
View File
@@ -406,6 +406,51 @@ Env vars: `DWS_SKILL_MODE=mono|multi` (also honored by `install.sh` / `install.p
## Features
<details>
<summary><strong>Personal Event Subscription</strong> — real-time DingTalk messages for event-driven agents</summary>
`dws event consume` subscribes as the currently logged-in user over a managed Stream WebSocket and emits each event as one NDJSON line on stdout. The public catalog currently covers messages that mention the current user, one-to-one messages with a specified user, and messages in a specified group.
> **Prerequisite**: run `dws auth login`. Personal identity is resolved from the OAuth token and cannot be supplied through command-line identity flags.
For an event-focused installation, use the official convenience installer:
```bash
curl -fsSL https://raw.githubusercontent.com/DingTalk-Real-AI/dingtalk-workspace-cli/main/scripts/install-event.sh | sh
```
```bash
# Inspect the public personal event catalog and schema
dws event list
dws event schema user_im_message_receive_o2o
# Listen for messages that mention the current user
dws event consume user_im_message_receive_at -f ndjson
# Listen for one-to-one messages with a specified user
dws event consume user_im_message_receive_o2o --user <userId> -f ndjson
# Listen for messages in a specified group
dws event consume user_im_message_receive_group --group <openConversationId> -f ndjson
# Inspect local consumers and cancel a subscription
dws event status
dws event stop <subscribe_id>
```
| Feature | Details |
|---------|---------|
| Managed lifecycle | `consume` creates or reuses the personal subscription; `stop` cancels it and cleans local state |
| Shared connection | Consumers for the same user share one local bus and cloud connection |
| Subscription isolation | Normal consumers match both event type and `subscribe_id` |
| Agent-friendly output | Stream events are written to stdout as NDJSON; status and diagnostics use stderr |
| Observability | `status` shows remote subscriptions, the personal bus, and local consumers |
| Cross-platform | Unix Socket on macOS/Linux, Windows Named Pipe on Windows |
See `skills/multi/dingtalk-event/SKILL.md` for the Agent workflow and supported event parameters.
</details>
<details>
<summary><strong>Raw API Access</strong> — call any DingTalk OpenAPI directly</summary>
+45
View File
@@ -403,6 +403,51 @@ DWS_SKILL_SOURCE=/path/to/skills dws skill setup --mode multi
## 功能特性
<details>
<summary><strong>个人事件订阅</strong> — 实时接收钉钉消息,驱动事件触发的 Agent</summary>
`dws event consume` 使用当前 OAuth 登录用户建立托管的 Stream WebSocket 长连接,并把每条事件以 NDJSON 一行输出到 stdout。当前公开目录包括:当前用户被 @ 的消息、与指定用户的单聊消息、指定群的消息。
> **前置条件**:先运行 `dws auth login`。个人身份从 OAuth token 解析,不允许通过命令行伪造。
只需要 event 能力时,可以使用官方便捷安装脚本:
```bash
curl -fsSL https://raw.githubusercontent.com/DingTalk-Real-AI/dingtalk-workspace-cli/main/scripts/install-event.sh | sh
```
```bash
# 查看公开个人事件目录和 schema
dws event list
dws event schema user_im_message_receive_o2o
# 监听当前用户被 @ 的消息
dws event consume user_im_message_receive_at -f ndjson
# 监听与指定用户的单聊消息
dws event consume user_im_message_receive_o2o --user <userId> -f ndjson
# 监听指定群的消息
dws event consume user_im_message_receive_group --group <openConversationId> -f ndjson
# 查看本地 consume,并取消指定订阅
dws event status
dws event stop <subscribe_id>
```
| 特性 | 说明 |
|------|------|
| 自动编排 | `consume` 创建或复用个人订阅,`stop` 取消订阅并清理本地状态 |
| 共享连接 | 同一用户的多个 consumer 共享本地 bus 和云端长连接 |
| 订阅隔离 | 正常 consumer 同时按事件类型和 `subscribe_id` 匹配 |
| Agent 友好输出 | Stream 事件写入 stdout,连接状态和诊断信息写入 stderr |
| 状态可观测 | `status` 同时显示服务端订阅、personal bus 和本地 consumers |
| 跨平台 | macOS/Linux 使用 Unix Socket,Windows 使用 Named Pipe |
Agent 工作流和事件参数详见 `skills/multi/dingtalk-event/SKILL.md`。
</details>
<details>
<summary><strong>Raw API 调用</strong> — 直接调用钉钉 OpenAPI</summary>
+2 -1
View File
@@ -3,12 +3,14 @@ module github.com/DingTalk-Real-AI/dingtalk-workspace-cli
go 1.25.9
require (
github.com/Microsoft/go-winio v0.6.2
github.com/RealAlexandreAI/json-repair v0.0.15
github.com/charmbracelet/bubbletea v1.3.6
github.com/charmbracelet/huh v1.0.0
github.com/charmbracelet/lipgloss v1.1.0
github.com/fatih/color v1.18.0
github.com/google/uuid v1.6.0
github.com/gorilla/websocket v1.5.0
github.com/itchyny/gojq v0.12.18
github.com/muesli/termenv v0.16.0
github.com/open-dingtalk/dingtalk-stream-sdk-go v0.9.1
@@ -35,7 +37,6 @@ require (
github.com/dustin/go-humanize v1.0.1 // indirect
github.com/erikgeiser/coninput v0.0.0-20211004153227-1c3628e74d0f // indirect
github.com/godbus/dbus/v5 v5.2.2 // indirect
github.com/gorilla/websocket v1.5.0 // indirect
github.com/itchyny/timefmt-go v0.1.7 // indirect
github.com/lucasb-eyer/go-colorful v1.2.0 // indirect
github.com/mattn/go-colorable v0.1.13 // indirect
+2
View File
@@ -1,5 +1,7 @@
github.com/MakeNowJust/heredoc v1.0.0 h1:cXCdzVdstXyiTqTvfqk9SDHpKNjxuom+DOlyEeQ4pzQ=
github.com/MakeNowJust/heredoc v1.0.0/go.mod h1:mG5amYoWBHf8vpLOuehzbGGw0EHxpZZ6lCpQ4fNJ8LE=
github.com/Microsoft/go-winio v0.6.2 h1:F2VQgta7ecxGYO8k3ZZz3RS8fVIXVxONVUPlNERoyfY=
github.com/Microsoft/go-winio v0.6.2/go.mod h1:yd8OoFMLzJbo9gZq8j5qaps8bJ9aShtEA8Ipt1oGCvU=
github.com/RealAlexandreAI/json-repair v0.0.15 h1:AN8/yt8rcphwQrIs/FZeki+cKaIERUNr25zf1flirIs=
github.com/RealAlexandreAI/json-repair v0.0.15/go.mod h1:GKJi5borR78O8c7HCVbgqjhoiVibZ6hJldxbc6dGrAI=
github.com/atotto/clipboard v0.1.4 h1:EH0zSVneZPSuFR11BlR9YppQTVDbh5+16AmcJi4g1z4=
File diff suppressed because it is too large Load Diff
+69
View File
@@ -0,0 +1,69 @@
// 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 (
"encoding/json"
"errors"
"os"
"testing"
authpkg "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/auth"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/source"
)
func writeEventTestAppConfig(t *testing.T, dir string, cfg authpkg.AppConfig) {
t.Helper()
raw, err := json.MarshalIndent(cfg, "", " ")
if err != nil {
t.Fatalf("marshal app config: %v", err)
}
if err := os.WriteFile(authpkg.GetAppConfigPath(dir), raw, 0o600); err != nil {
t.Fatalf("write app config: %v", err)
}
}
func TestResolveEventCredentials_PortalNormalAllowsMissingClientSecret(t *testing.T) {
t.Setenv(authpkg.EnvClientID, "")
t.Setenv(authpkg.EnvClientSecret, "")
dir := t.TempDir()
clientID, clientSecret, err := resolveEventCredentials(dir, eventStreamTicketOptions{
Mode: source.PortalTicketModeNormal,
SourceID: "pre_open_source",
})
if err != nil {
t.Fatalf("resolveEventCredentials: %v", err)
}
if clientID != "portal-ticket-normal:pre_open_source" {
t.Fatalf("clientID = %q, want portal-ticket-normal:pre_open_source", clientID)
}
if clientSecret != "" {
t.Fatalf("clientSecret = %q, want empty", clientSecret)
}
}
func TestResolveEventCredentials_PortalCustomStillRequiresClientSecret(t *testing.T) {
t.Setenv(authpkg.EnvClientID, "")
t.Setenv(authpkg.EnvClientSecret, "")
dir := t.TempDir()
writeEventTestAppConfig(t, dir, authpkg.AppConfig{ClientID: "ding-custom"})
_, _, err := resolveEventCredentials(dir, eventStreamTicketOptions{
Mode: source.PortalTicketModeCustom,
})
if !errors.Is(err, authpkg.ErrClientSecretEmpty) {
t.Fatalf("err = %v, want ErrClientSecretEmpty", err)
}
}
+845
View File
@@ -0,0 +1,845 @@
// 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"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"io"
"net/http"
"os"
"path/filepath"
"sort"
"strings"
"text/tabwriter"
"time"
authpkg "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/auth"
dwsevent "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/bus"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/busctl"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/consume"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/personal"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/source"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
"github.com/spf13/cobra"
)
type commonConsumeOptions struct {
EventTypes []string
Filter string
Compact bool
FormatRaw string
OutputDir string
RoutesRaw []string
MaxEvents int
Duration time.Duration
Quiet bool
Force bool
DryRun bool
Foreground bool
}
type personalConsumeOptions struct {
Common commonConsumeOptions
EventKey string
DebugRawEvents bool
SubscribeID string
Rule string
Name string
FilterJSON string
QueryCSV string
TTL time.Duration
Ephemeral bool
UserID string
GroupID string
ControlBaseURL string
StreamTicketMode string
StreamTicketURL string
StreamSourceID string
}
type personalListOptions struct {
Category string
EnabledOnly bool
IncludePending bool
Format string
}
type personalStatusOptions struct {
EventKey string
Status string
SubscribeID string
Format string
ControlBaseURL string
StreamSourceID string
}
type personalStopOptions struct {
SubscribeID string
All bool
ControlBaseURL string
StreamSourceID string
}
type personalStreamSourceOptions struct {
ConfigDir string
Identity personal.Identity
TicketMode string
TicketURL string
ClientIDOverride string
}
func newEventSchemaCommand() *cobra.Command {
var asIdentity string
var formatRaw string
cmd := &cobra.Command{
Use: "schema <event_key>",
Short: "显示事件 schema",
Args: cobra.ExactArgs(1),
DisableAutoGenTag: true,
RunE: func(c *cobra.Command, args []string) error {
as, err := normalizeEventAs(asIdentity)
if err != nil {
return err
}
if as != "user" {
return fmt.Errorf("event schema is only supported with --as user")
}
def, ok := personal.Lookup(args[0])
if !ok {
return fmt.Errorf("unknown personal event key %q", args[0])
}
if !def.Public {
return personal.PublicAvailabilityError(args[0])
}
return renderPersonalSchema(c.OutOrStdout(), def, formatRaw)
},
}
cmd.Flags().StringVar(&asIdentity, "as", "user", "事件身份: user")
cmd.Flags().StringVarP(&formatRaw, "format", "f", "json", "输出格式: json")
hideEventInternalFlags(cmd, "as")
return cmd
}
func runPersonalEventList(c *cobra.Command, opts personalListOptions) error {
items := personal.Catalog(opts.Category, opts.EnabledOnly, opts.IncludePending)
if opts.Format == "json" {
enc := json.NewEncoder(c.OutOrStdout())
enc.SetIndent("", " ")
return enc.Encode(items)
}
tw := tabwriter.NewWriter(c.OutOrStdout(), 0, 0, 2, ' ', 0)
fmt.Fprintln(tw, "EVENT_KEY\tRULE\tSTATUS\tDESCRIPTION")
for _, it := range items {
fmt.Fprintf(tw, "%s\t%s\t%s\t%s\n",
it.EventKey, it.RuleType, it.Status, it.Description)
}
return tw.Flush()
}
func renderPersonalSchema(w io.Writer, def personal.Definition, format string) error {
format = strings.ToLower(strings.TrimSpace(format))
if format == "" {
format = "json"
}
if format != "json" {
return fmt.Errorf("event schema only supports json output")
}
enc := json.NewEncoder(w)
enc.SetIndent("", " ")
return enc.Encode(personal.BuildSchemaDocument(def))
}
func runPersonalEventConsume(c *cobra.Command, opts personalConsumeOptions) error {
ctx := c.Context()
if err := ensurePublicPersonalEvent(opts.EventKey); err != nil {
return err
}
configDir := defaultConfigDir()
identity, err := resolvePersonalEventIdentity(ctx, configDir, opts.StreamSourceID)
if err != nil {
return fmt.Errorf("event consume --as user: %w", err)
}
identityHash := dwsevent.IdentityHash(identity.Key())
editionName := editionNameOrDefault()
workDir := eventWorkDir(configDir, editionName, dwsevent.SourceKindPersonalStream, identityHash)
ipcEndpoint := defaultIPCEndpoint(workDir, editionName, dwsevent.SourceKindPersonalStream, identityHash)
routes, err := consume.ParseRoutes(opts.Common.RoutesRaw)
if err != nil {
return fmt.Errorf("event consume --as user: %w", err)
}
rawFormat := ""
if f := c.Flags().Lookup("format"); f != nil && f.Changed {
rawFormat = opts.Common.FormatRaw
}
normalised, fellback := consume.NormalizeFormat(rawFormat)
if fellback && !opts.Common.Quiet {
fmt.Fprintf(c.ErrOrStderr(), "WARN: --format %q has no meaning for event stream; using ndjson\n", rawFormat)
}
if opts.Common.DryRun {
cfg := consume.Config{
WorkDir: workDir,
IPCEndpoint: ipcEndpoint,
ClientID: identity.ClientID,
SpawnExtraArgs: personalBusSpawnArgs(identity, opts.StreamTicketMode, personalEventStreamTicketURL(opts.StreamTicketURL, configDir)),
Compact: opts.Common.Compact,
MaxEvents: opts.Common.MaxEvents,
Duration: opts.Common.Duration,
Format: normalised,
OutputDir: opts.Common.OutputDir,
Routes: routes,
Stderr: c.ErrOrStderr(),
Quiet: opts.Common.Quiet,
Foreground: opts.Common.Foreground,
Force: opts.Common.Force,
DryRun: true,
}
applyPersonalConsumeFilters(&cfg, opts, strings.TrimSpace(opts.SubscribeID), opts.EventKey)
return consume.Run(ctx, cfg)
}
client := personal.NewClient(personalEventControlBaseURL(opts.ControlBaseURL, configDir), identity)
sub, eventKey, ruleType, err := ensurePersonalSubscription(ctx, client, identity, opts)
if err != nil {
return fmt.Errorf("event consume --as user: %w", err)
}
if sub.SubscribeID == "" {
return fmt.Errorf("event consume --as user: server returned empty subscribe_id")
}
if err := personal.UpsertRunState(workDir, personal.RunState{
SubscribeID: sub.SubscribeID,
EventKey: eventKey,
RuleType: ruleType,
ClientID: identity.ClientID,
SourceID: identity.SourceID,
IdentityHash: identityHash,
}); err != nil {
return fmt.Errorf("event consume --as user: save run state: %w", err)
}
cleanup := func() {
_ = client.DeleteSubscription(context.Background(), sub.SubscribeID)
_ = personal.RemoveRunStates(workDir, []string{sub.SubscribeID})
}
if opts.Ephemeral {
defer cleanup()
}
cfg := consume.Config{
WorkDir: workDir,
IPCEndpoint: ipcEndpoint,
ClientID: identity.ClientID,
SpawnExtraArgs: personalBusSpawnArgs(identity, opts.StreamTicketMode, opts.StreamTicketURL),
Compact: opts.Common.Compact,
MaxEvents: opts.Common.MaxEvents,
Duration: opts.Common.Duration,
Format: normalised,
OutputDir: opts.Common.OutputDir,
Routes: routes,
Stdout: c.OutOrStdout(),
Stderr: c.ErrOrStderr(),
Quiet: opts.Common.Quiet,
Foreground: opts.Common.Foreground,
Force: opts.Common.Force,
}
applyPersonalConsumeFilters(&cfg, opts, sub.SubscribeID, eventKey)
if opts.DebugRawEvents && !opts.Common.Quiet {
fmt.Fprintf(c.ErrOrStderr(), "debug raw events enabled: local event filters disabled\nworkdir: %s\nbus_log: %s\n",
workDir, filepath.Join(workDir, "bus.log"))
}
if err := consume.ValidateConfig(cfg); err != nil {
return err
}
if o := c.Flags().Lookup("output"); o != nil && o.Changed {
if err := consume.ValidateNoOutputConflict(cfg, o.Value.String()); err != nil {
return err
}
}
if opts.Common.Foreground {
src, err := newPersonalStreamSource(ctx, personalStreamSourceOptions{
ConfigDir: configDir,
Identity: identity,
TicketMode: opts.StreamTicketMode,
TicketURL: opts.StreamTicketURL,
})
if err != nil {
if !opts.Ephemeral {
cleanup()
}
return err
}
busCfg := bus.Config{
WorkDir: workDir,
IPCEndpoint: ipcEndpoint,
ClientID: identity.ClientID,
Edition: editionName,
SourceKind: dwsevent.SourceKindPersonalStream,
IdentityHash: identityHash,
SourceID: identity.SourceID,
Source: src,
}
bus.ApplyEnvTuning(&busCfg)
err = bus.Run(ctx, busCfg)
if err != nil && !opts.Ephemeral {
cleanup()
}
return err
}
err = consume.Run(ctx, cfg)
if err != nil && !opts.Ephemeral {
cleanup()
}
return err
}
func applyPersonalConsumeFilters(cfg *consume.Config, opts personalConsumeOptions, subscribeID, eventKey string) {
if cfg == nil {
return
}
if opts.DebugRawEvents {
cfg.EventTypes = nil
cfg.Filter = ""
cfg.SubscribeID = ""
return
}
cfg.EventTypes = personalEventTypes(eventKey, opts.Common.EventTypes)
cfg.Filter = opts.Common.Filter
cfg.SubscribeID = strings.TrimSpace(subscribeID)
}
func ensurePersonalSubscription(ctx context.Context, client *personal.Client, identity personal.Identity, opts personalConsumeOptions) (*personal.Subscription, string, string, error) {
if strings.TrimSpace(opts.SubscribeID) != "" {
sub, err := client.GetSubscription(ctx, opts.SubscribeID)
if err != nil {
return nil, "", "", err
}
eventKey := firstNonEmptyPersonalString(opts.EventKey, sub.EventKey)
if eventKey == "" {
return nil, "", "", fmt.Errorf("event_key is required when --subscribe-id lookup returns no event_key")
}
if err := ensurePublicPersonalEvent(eventKey); err != nil {
return nil, "", "", err
}
ruleType := firstNonEmptyPersonalString(sub.RuleType, opts.Rule)
if ruleType == "" {
if def, ok := personal.Lookup(eventKey); ok {
ruleType = def.RuleType
}
}
sub.SubscribeID = strings.TrimSpace(opts.SubscribeID)
return sub, eventKey, ruleType, nil
}
if strings.TrimSpace(opts.EventKey) == "" {
return nil, "", "", fmt.Errorf("event_key is required unless --subscribe-id is provided")
}
if err := ensurePublicPersonalEvent(opts.EventKey); err != nil {
return nil, "", "", err
}
ruleType, ruleParam, err := personal.BuildRuleParam(opts.EventKey, personal.RuleOptions{
RuleType: opts.Rule,
UserID: opts.UserID,
GroupID: opts.GroupID,
})
if err != nil {
return nil, "", "", err
}
filter, filterCanonical, err := personal.BuildFilter(opts.FilterJSON, opts.QueryCSV)
if err != nil {
return nil, "", "", err
}
req := personal.CreateSubscriptionRequest{
EventKey: opts.EventKey,
RuleType: ruleType,
Name: opts.Name,
RuleParam: ruleParam,
Filter: filter,
Delivery: map[string]any{"mode": "stream"},
IdempotencyKey: personal.IdempotencyKey(identity, opts.EventKey, ruleType, ruleParam, filterCanonical),
}
if opts.TTL > 0 {
req.TTLSeconds = int64(opts.TTL.Seconds())
}
sub, err := client.CreateSubscription(ctx, req)
if err != nil {
return nil, "", "", err
}
return sub, opts.EventKey, ruleType, nil
}
func runPersonalEventStatus(c *cobra.Command, opts personalStatusOptions) error {
ctx := c.Context()
if err := ensurePublicPersonalEvent(opts.EventKey); err != nil {
return err
}
configDir := defaultConfigDir()
identity, err := resolvePersonalEventIdentity(ctx, configDir, opts.StreamSourceID)
if err != nil {
return fmt.Errorf("event status --as user: %w", err)
}
identityHash := dwsevent.IdentityHash(identity.Key())
editionName := editionNameOrDefault()
workDir := eventWorkDir(configDir, editionName, dwsevent.SourceKindPersonalStream, identityHash)
entry := busctl.FindBusByIdentity(configDir, editionName, dwsevent.SourceKindPersonalStream, identityHash)
var qs busctl.EntryStatus
if entry != nil {
qs = busctl.QueryEntry(*entry)
} else {
qs = busctl.EntryStatus{Entry: busctl.BusEntry{
WorkDir: workDir,
Edition: editionName,
SourceKind: dwsevent.SourceKindPersonalStream,
ClientIDHash: identityHash,
IdentityHash: identityHash,
State: busctl.BusStateNotRunning,
Meta: &bus.Meta{
ClientID: identity.ClientID,
Edition: editionName,
SourceKind: dwsevent.SourceKindPersonalStream,
IdentityHash: identityHash,
SourceID: identity.SourceID,
},
}}
}
status := opts.Status
if status == "" || status == "all" {
status = ""
}
subs, err := personal.NewClient(personalEventControlBaseURL(opts.ControlBaseURL, configDir), identity).ListSubscriptions(ctx, personal.ListOptions{
Status: status,
EventKey: opts.EventKey,
SubscribeID: opts.SubscribeID,
})
if err != nil {
return fmt.Errorf("event status --as user: %w", err)
}
if opts.Format == "json" {
enc := json.NewEncoder(c.OutOrStdout())
enc.SetIndent("", " ")
return enc.Encode(map[string]any{
"identity": redactedPersonalIdentity(identity, identityHash),
"subscriptions": subs,
"bus": qs,
})
}
renderPersonalStatusText(c.OutOrStdout(), identity, identityHash, subs, qs)
return nil
}
func ensurePublicPersonalEvent(eventKey string) error {
eventKey = strings.TrimSpace(eventKey)
if eventKey == "" {
return nil
}
if def, ok := personal.Lookup(eventKey); ok && !def.Public {
return personal.PublicAvailabilityError(eventKey)
}
return nil
}
func renderPersonalStatusText(w io.Writer, identity personal.Identity, identityHash string, subs []personal.Subscription, qs busctl.EntryStatus) {
fmt.Fprintf(w, "Personal identity: corp=%s user=%s client=%s source=%s hash=%s\n",
displayIdentityPart(identity.CorpID), displayIdentityPart(identity.UserID), identity.ClientID, identity.SourceID, identityHash)
fmt.Fprintf(w, "Bus: %s", qs.Entry.State)
if qs.Entry.HolderPID > 0 {
fmt.Fprintf(w, " pid=%d", qs.Entry.HolderPID)
}
fmt.Fprintf(w, "\nWorkdir: %s\n", qs.Entry.WorkDir)
if len(subs) == 0 {
fmt.Fprintln(w, "Subscriptions: none")
} else {
tw := tabwriter.NewWriter(w, 0, 0, 2, ' ', 0)
fmt.Fprintln(tw, "SUBSCRIBE_ID\tEVENT_KEY\tRULE\tSTATUS\tSOURCE")
for _, sub := range subs {
fmt.Fprintf(tw, "%s\t%s\t%s\t%s\t%s\n",
sub.SubscribeID, sub.EventKey, sub.RuleType, sub.Status, sub.SourceID)
}
_ = tw.Flush()
}
renderPersonalConsumers(w, qs)
}
func renderPersonalConsumers(w io.Writer, qs busctl.EntryStatus) {
if qs.Entry.State != busctl.BusStateRunning {
fmt.Fprintln(w, "Consumers: none")
return
}
if qs.Live == nil {
fmt.Fprintln(w, "Consumers: unavailable (status RPC failed)")
return
}
if len(qs.Live.Consumers) == 0 {
fmt.Fprintln(w, "Consumers: none")
return
}
fmt.Fprintln(w, "Consumers:")
tw := tabwriter.NewWriter(w, 0, 0, 2, ' ', 0)
fmt.Fprintln(tw, "PID\tEVENT_KEYS\tSUBSCRIBE_ID\tFILTER\tRECEIVED\tDROPPED")
for _, cs := range qs.Live.Consumers {
eventKeys := strings.Join(cs.EventTypes, ",")
if eventKeys == "" {
eventKeys = "(catch-all)"
}
subscribeID := displayPersonalStatusValue(cs.SubscribeID)
filter := displayPersonalStatusValue(cs.Filter)
fmt.Fprintf(tw, "%d\t%s\t%s\t%s\t%d\t%d\n",
cs.PID, eventKeys, subscribeID, filter, cs.Received, cs.Dropped)
}
_ = tw.Flush()
}
func displayPersonalStatusValue(v string) string {
v = strings.TrimSpace(v)
if v == "" {
return "-"
}
return v
}
func runPersonalEventStop(c *cobra.Command, opts personalStopOptions) error {
ctx := c.Context()
explicitSubscribeID := strings.TrimSpace(opts.SubscribeID)
isSingleTarget := explicitSubscribeID != ""
if explicitSubscribeID != "" && opts.All {
return fmt.Errorf("event stop --as user: subscribe_id and --all are mutually exclusive")
}
if explicitSubscribeID == "" && !opts.All {
return fmt.Errorf("event stop --as user: subscribe_id is required unless --all is set")
}
configDir := defaultConfigDir()
identity, err := resolvePersonalEventIdentity(ctx, configDir, opts.StreamSourceID)
if err != nil {
return fmt.Errorf("event stop --as user: %w", err)
}
identityHash := dwsevent.IdentityHash(identity.Key())
editionName := editionNameOrDefault()
workDir := eventWorkDir(configDir, editionName, dwsevent.SourceKindPersonalStream, identityHash)
ipcEndpoint := defaultIPCEndpoint(workDir, editionName, dwsevent.SourceKindPersonalStream, identityHash)
subscribeIDs, err := personalStopTargets(workDir, explicitSubscribeID, opts.All)
if err != nil {
return fmt.Errorf("event stop --as user: %w", err)
}
client := personal.NewClient(personalEventControlBaseURL(opts.ControlBaseURL, configDir), identity)
for _, id := range subscribeIDs {
if err := client.DeleteSubscription(ctx, id); err != nil {
return fmt.Errorf("event stop --as user: cancel subscription %s: %w", id, err)
}
}
if err := personal.RemoveRunStates(workDir, subscribeIDs); err != nil {
return fmt.Errorf("event stop --as user: update local state: %w", err)
}
if err := interruptPersonalConsumers(ipcEndpoint, subscribeIDs); err != nil {
fmt.Fprintf(c.ErrOrStderr(), "WARN: failed to stop matching local consume process: %v\n", err)
}
remaining, err := personal.LoadRunStates(workDir)
if err != nil {
return fmt.Errorf("event stop --as user: load remaining local state: %w", err)
}
if len(remaining) > 0 {
printPersonalStopResult(c.OutOrStdout(), subscribeIDs, isSingleTarget, "personal bus still running")
return nil
}
busState := "personal bus stopped"
if err := busctl.Stop(busctl.StopConfig{WorkDir: workDir}); err != nil {
if errors.Is(err, busctl.ErrNotRunning) {
busState = "personal bus is not running"
} else {
return err
}
}
printPersonalStopResult(c.OutOrStdout(), subscribeIDs, isSingleTarget, busState)
return nil
}
func personalStopTargets(workDir, explicit string, all bool) ([]string, error) {
explicit = strings.TrimSpace(explicit)
if explicit != "" && all {
return nil, fmt.Errorf("subscribe_id and --all are mutually exclusive")
}
if explicit != "" {
return []string{explicit}, nil
}
if !all {
return nil, fmt.Errorf("subscribe_id is required unless --all is set")
}
states, err := personal.LoadRunStates(workDir)
if err != nil {
return nil, err
}
ids := make([]string, 0, len(states))
for _, st := range states {
if st.SubscribeID != "" {
ids = append(ids, st.SubscribeID)
}
}
sort.Strings(ids)
return ids, nil
}
func interruptPersonalConsumers(ipcEndpoint string, subscribeIDs []string) error {
targets := make(map[string]struct{}, len(subscribeIDs))
for _, id := range subscribeIDs {
id = strings.TrimSpace(id)
if id != "" {
targets[id] = struct{}{}
}
}
if ipcEndpoint == "" || len(targets) == 0 {
return nil
}
status, err := busctl.QueryStatus(ipcEndpoint)
if err != nil {
return nil
}
signalled := make(map[int]struct{})
for _, consumer := range status.Consumers {
if _, ok := targets[strings.TrimSpace(consumer.SubscribeID)]; !ok {
continue
}
if consumer.PID <= 0 || consumer.PID == os.Getpid() {
continue
}
if _, ok := signalled[consumer.PID]; ok {
continue
}
proc, err := os.FindProcess(consumer.PID)
if err != nil {
return fmt.Errorf("find consume pid=%d: %w", consumer.PID, err)
}
if err := proc.Signal(os.Interrupt); err != nil && !errors.Is(err, os.ErrProcessDone) {
return fmt.Errorf("signal consume pid=%d: %w", consumer.PID, err)
}
signalled[consumer.PID] = struct{}{}
}
return nil
}
func printPersonalStopResult(w io.Writer, subscribeIDs []string, single bool, busState string) {
if single && len(subscribeIDs) == 1 {
fmt.Fprintf(w, "cancelled personal subscription %s; %s\n", subscribeIDs[0], busState)
return
}
fmt.Fprintf(w, "cancelled %d personal subscription(s); %s\n", len(subscribeIDs), busState)
}
func resolvePersonalEventIdentity(ctx context.Context, configDir string, sourceIDOverride string) (personal.Identity, error) {
accessToken, err := ResolveAuxiliaryAccessToken(ctx, configDir, "")
if err != nil {
return personal.Identity{}, err
}
tokenData, _ := authpkg.LoadTokenData(configDir)
var corpID, userID, clientID, refreshToken string
if tokenData != nil {
corpID = tokenData.CorpID
userID = tokenData.UserID
clientID = tokenData.ClientID
refreshToken = tokenData.RefreshToken
}
if corpID == "" {
corpID = resolveRuntimeDefault(ctx, "$corpId")
}
if userID == "" {
userID = resolveRuntimeDefault(ctx, "$currentUserId")
}
if clientID == "" {
clientID = authpkg.ClientID()
}
if clientID == "" {
if id, _, _, _, err := authpkg.ResolveAppCredentialsStrict(configDir); err == nil {
clientID = id
}
}
if clientID == "" {
return personal.Identity{}, fmt.Errorf("cannot resolve OAuth client_id for personal events")
}
sourceID := strings.TrimSpace(sourceIDOverride)
if sourceID == "" {
sourceID = personalEventStreamSourceID("")
}
localSubject := ""
if strings.TrimSpace(corpID) == "" || strings.TrimSpace(userID) == "" {
localSubject = personalTokenSubject("refresh", refreshToken)
if localSubject == "" {
localSubject = personalTokenSubject("access", accessToken)
}
}
return personal.Identity{
AccessToken: accessToken,
LocalSubject: localSubject,
CorpID: corpID,
UserID: userID,
ClientID: clientID,
SourceID: sourceID,
}, nil
}
func personalTokenSubject(kind, token string) string {
token = strings.TrimSpace(token)
if token == "" {
return ""
}
sum := sha256.Sum256([]byte(token))
return strings.TrimSpace(kind) + ":" + hex.EncodeToString(sum[:])
}
func resolveRuntimeDefault(ctx context.Context, key string) string {
if fnMap := edition.Get().RuntimeDefaults; fnMap != nil {
if fn := fnMap()[key]; fn != nil {
if v, ok := fn(ctx); ok {
return strings.TrimSpace(v)
}
}
}
return ""
}
func newPersonalStreamSource(ctx context.Context, opts personalStreamSourceOptions) (*source.PersonalSource, error) {
mode := strings.TrimSpace(opts.TicketMode)
if mode == "" {
mode = "normal"
}
if mode != "normal" && mode != "custom" {
return nil, fmt.Errorf("stream ticket mode must be normal or custom")
}
ticketURL := strings.TrimSpace(opts.TicketURL)
if ticketURL == "" {
ticketURL = personalEventStreamTicketURL("", opts.ConfigDir)
}
clientID := opts.Identity.ClientID
clientSecret := ""
if mode == "custom" {
resolvedID, secret, _, _, err := authpkg.ResolveAppCredentialsStrict(opts.ConfigDir)
if err != nil {
return nil, err
}
if opts.ClientIDOverride != "" {
clientID = opts.ClientIDOverride
} else if clientID == "" {
clientID = resolvedID
}
clientSecret = secret
}
_ = ctx
return source.NewPersonal(source.PersonalConfig{
AccessToken: opts.Identity.AccessToken,
ClientID: clientID,
ClientSecret: clientSecret,
SourceID: opts.Identity.SourceID,
TicketURL: ticketURL,
TicketMode: mode,
HTTPClient: &http.Client{Timeout: 30 * time.Second},
})
}
func personalBusSpawnArgs(identity personal.Identity, ticketMode, ticketURL string) []string {
args := []string{
"--source-kind", string(dwsevent.SourceKindPersonalStream),
"--stream-source-id", identity.SourceID,
}
if strings.TrimSpace(ticketMode) != "" {
args = append(args, "--stream-ticket-mode", ticketMode)
}
if strings.TrimSpace(ticketURL) != "" {
args = append(args, "--stream-ticket-url", ticketURL)
}
return args
}
func personalEventTypes(eventKey string, explicit []string) []string {
if len(explicit) > 0 {
return explicit
}
if strings.TrimSpace(eventKey) == "" {
return nil
}
return []string{eventKey}
}
func redactedPersonalIdentity(identity personal.Identity, identityHash string) map[string]string {
return map[string]string{
"corp_id": displayIdentityPart(identity.CorpID),
"user_id": displayIdentityPart(identity.UserID),
"client_id": identity.ClientID,
"source_id": identity.SourceID,
"identity_hash": identityHash,
}
}
func displayIdentityPart(v string) string {
v = strings.TrimSpace(v)
if v == "" {
return "unknown"
}
return v
}
func firstNonEmptyPersonalString(values ...string) string {
for _, v := range values {
if strings.TrimSpace(v) != "" {
return strings.TrimSpace(v)
}
}
return ""
}
func personalEventControlBaseURL(raw, configDir string) string {
if v := strings.TrimSpace(raw); v != "" {
return strings.TrimRight(v, "/")
}
return personalEventMCPBaseURL(configDir) + personal.DefaultBasePath
}
func personalEventStreamTicketURL(raw, configDir string) string {
if v := strings.TrimSpace(raw); v != "" {
return strings.TrimRight(v, "/")
}
return personalEventMCPBaseURL(configDir) + "/stream/connections/ticket"
}
func personalEventStreamSourceID(raw string) string {
if v := strings.TrimSpace(raw); v != "" {
return v
}
if v := strings.TrimSpace(edition.PersonalEventSourceID()); v != "" {
return v
}
return "open"
}
func personalEventMCPBaseURL(configDir string) string {
if v := configuredMCPBaseURL(configDir); v != "" {
return strings.TrimRight(v, "/")
}
return config.DefaultMCPBaseURL
}
func configuredMCPBaseURL(configDir string) string {
if strings.TrimSpace(configDir) == "" {
configDir = defaultConfigDir()
}
data, err := os.ReadFile(filepath.Join(configDir, "mcp_url"))
if err != nil {
return ""
}
return strings.TrimSpace(string(data))
}
+131
View File
@@ -0,0 +1,131 @@
// 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 (
"strings"
"testing"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/consume"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/personal"
)
func TestApplyPersonalConsumeFiltersDebugRawEvents(t *testing.T) {
cfg := consume.Config{}
opts := personalConsumeOptions{
DebugRawEvents: true,
Common: commonConsumeOptions{
EventTypes: []string{"should-not-survive"},
Filter: "^should-not-survive$",
},
}
applyPersonalConsumeFilters(&cfg, opts, "sub-1", "user_im_message_receive_o2o")
if cfg.EventTypes != nil || cfg.Filter != "" || cfg.SubscribeID != "" {
t.Fatalf("raw debug filters = eventTypes=%#v filter=%q subscribeID=%q, want catch-all", cfg.EventTypes, cfg.Filter, cfg.SubscribeID)
}
}
func TestApplyPersonalConsumeFiltersDefault(t *testing.T) {
cfg := consume.Config{}
opts := personalConsumeOptions{Common: commonConsumeOptions{Filter: "^user_im_"}}
applyPersonalConsumeFilters(&cfg, opts, "sub-1", "user_im_message_receive_o2o")
if len(cfg.EventTypes) != 1 || cfg.EventTypes[0] != "user_im_message_receive_o2o" {
t.Fatalf("eventTypes = %#v", cfg.EventTypes)
}
if cfg.Filter != "^user_im_" || cfg.SubscribeID != "sub-1" {
t.Fatalf("filter=%q subscribeID=%q", cfg.Filter, cfg.SubscribeID)
}
}
func TestEventConsumeDebugRawEventsRequiresUserMode(t *testing.T) {
cmd := newEventConsumeCommand()
cmd.SilenceUsage = true
cmd.SilenceErrors = true
cmd.SetArgs([]string{"--as", "app", "--debug-raw-events"})
err := cmd.Execute()
if err == nil || !strings.Contains(err.Error(), "app event is not publicly available yet") {
t.Fatalf("Execute() error = %v, want public availability guard", err)
}
}
func TestEventConsumeAsAppRejectedBeforeEventKeyValidation(t *testing.T) {
cmd := newEventConsumeCommand()
cmd.SilenceUsage = true
cmd.SilenceErrors = true
cmd.SetArgs([]string{"--as", "app", personal.EventSingleChat})
err := cmd.Execute()
if err == nil || !strings.Contains(err.Error(), "app event is not publicly available yet") {
t.Fatalf("Execute() error = %v, want public availability guard", err)
}
}
func TestEventConsumePersonalParamSpecFlags(t *testing.T) {
cmd := newEventConsumeCommand()
for _, name := range []string{"user", "group", "query"} {
if cmd.Flags().Lookup(name) == nil {
t.Fatalf("flag --%s is not registered", name)
}
}
for _, name := range []string{
"peer-user-id",
"peer-union-id",
"sender-user-id",
"sender-union-id",
"open-conversation-id",
"keyword",
} {
if cmd.Flags().Lookup(name) != nil {
t.Fatalf("retired flag --%s is still registered", name)
}
}
}
func TestEventConsumeRetiredPersonalFlagsAreUnknown(t *testing.T) {
for _, name := range []string{
"peer-user-id",
"peer-union-id",
"sender-user-id",
"sender-union-id",
"open-conversation-id",
"keyword",
} {
t.Run(name, func(t *testing.T) {
cmd := newEventConsumeCommand()
cmd.SilenceUsage = true
cmd.SilenceErrors = true
cmd.SetArgs([]string{personal.EventSingleChat, "--" + name, "x"})
err := cmd.Execute()
if err == nil || !strings.Contains(err.Error(), "unknown flag: --"+name) {
t.Fatalf("Execute() error = %v, want unknown flag", err)
}
})
}
}
func TestEventConsumeAsAppRejectedBeforePersonalParamSpecFlags(t *testing.T) {
for _, args := range [][]string{
{"--as", "app", "--user", "507971"},
{"--as", "app", "--group", "cid"},
{"--as", "app", "--query", "报警"},
} {
cmd := newEventConsumeCommand()
cmd.SilenceUsage = true
cmd.SilenceErrors = true
cmd.SetArgs(args)
err := cmd.Execute()
if err == nil || !strings.Contains(err.Error(), "app event is not publicly available yet") {
t.Fatalf("Execute(%v) error = %v, want public availability guard", args, err)
}
}
}
@@ -0,0 +1,211 @@
// 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"
"path/filepath"
"strings"
"testing"
"time"
authpkg "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/auth"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/busctl"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
)
func TestResolvePersonalEventIdentityUsesCorpUserWhenAvailable(t *testing.T) {
configDir := setupPersonalIdentityToken(t, &authpkg.TokenData{
AccessToken: "access-1",
RefreshToken: "refresh-1",
ExpiresAt: time.Now().Add(time.Hour),
RefreshExpAt: time.Now().Add(24 * time.Hour),
CorpID: "corp-1",
UserID: "user-1",
ClientID: "client-1",
})
identity, err := resolvePersonalEventIdentity(context.Background(), configDir, "pre_open_source")
if err != nil {
t.Fatalf("resolvePersonalEventIdentity() error = %v", err)
}
if identity.LocalSubject != "" {
t.Fatalf("LocalSubject = %q, want empty when corp/user are available", identity.LocalSubject)
}
wantKey := "corp_user\x00corp-1\x00user-1\x00client-1\x00pre_open_source"
if got := identity.Key(); got != wantKey {
t.Fatalf("identity key = %q, want %q", got, wantKey)
}
}
func TestResolvePersonalEventIdentityFallsBackToRefreshTokenSubject(t *testing.T) {
configDir := setupPersonalIdentityToken(t, &authpkg.TokenData{
AccessToken: "access-1",
RefreshToken: "refresh-1",
ExpiresAt: time.Now().Add(time.Hour),
RefreshExpAt: time.Now().Add(24 * time.Hour),
ClientID: "client-1",
})
identity, err := resolvePersonalEventIdentity(context.Background(), configDir, "pre_open_source")
if err != nil {
t.Fatalf("resolvePersonalEventIdentity() error = %v", err)
}
wantSubject := personalTokenSubject("refresh", "refresh-1")
if identity.LocalSubject != wantSubject {
t.Fatalf("LocalSubject = %q, want %q", identity.LocalSubject, wantSubject)
}
if strings.Contains(identity.Key(), "refresh-1") || strings.Contains(identity.Key(), "access-1") {
t.Fatalf("identity key leaked raw token: %q", identity.Key())
}
body, err := json.Marshal(redactedPersonalIdentity(identity, "identity-hash-1"))
if err != nil {
t.Fatalf("marshal redacted identity: %v", err)
}
if strings.Contains(string(body), wantSubject) || strings.Contains(string(body), "refresh-1") || strings.Contains(string(body), "access-1") {
t.Fatalf("redacted identity leaked local subject/token: %s", string(body))
}
if !strings.Contains(string(body), "unknown") {
t.Fatalf("redacted identity should mark missing corp/user as unknown: %s", string(body))
}
}
func TestResolvePersonalEventIdentityFallsBackToAccessTokenSubject(t *testing.T) {
configDir := setupPersonalIdentityToken(t, &authpkg.TokenData{
AccessToken: "access-1",
ExpiresAt: time.Now().Add(time.Hour),
ClientID: "client-1",
})
identity, err := resolvePersonalEventIdentity(context.Background(), configDir, "pre_open_source")
if err != nil {
t.Fatalf("resolvePersonalEventIdentity() error = %v", err)
}
wantSubject := personalTokenSubject("access", "access-1")
if identity.LocalSubject != wantSubject {
t.Fatalf("LocalSubject = %q, want %q", identity.LocalSubject, wantSubject)
}
var out bytes.Buffer
renderPersonalStatusText(&out, identity, "identity-hash-1", nil, busctl.EntryStatus{
Entry: busctl.BusEntry{WorkDir: "wd", State: busctl.BusStateNotRunning},
})
rendered := out.String()
if !strings.Contains(rendered, "corp=unknown user=unknown") {
t.Fatalf("status output = %q, want unknown corp/user", rendered)
}
if strings.Contains(rendered, wantSubject) || strings.Contains(rendered, "access-1") {
t.Fatalf("status output leaked local subject/token: %q", rendered)
}
}
func TestResolvePersonalEventIdentityDefaultsSourceIDToOpen(t *testing.T) {
configDir := setupPersonalIdentityToken(t, &authpkg.TokenData{
AccessToken: "access-1",
ExpiresAt: time.Now().Add(time.Hour),
CorpID: "corp-1",
UserID: "user-1",
ClientID: "client-1",
})
identity, err := resolvePersonalEventIdentity(context.Background(), configDir, "")
if err != nil {
t.Fatalf("resolvePersonalEventIdentity() error = %v", err)
}
if identity.SourceID != "open" {
t.Fatalf("SourceID = %q, want open", identity.SourceID)
}
}
func TestPersonalEventDefaultsUseProductionWithoutMCPConfig(t *testing.T) {
dir := t.TempDir()
t.Setenv("DWS_CONFIG_DIR", dir)
prev := edition.Get()
edition.Override(&edition.Hooks{})
t.Cleanup(func() { edition.Override(prev) })
if got := personalEventControlBaseURL("", dir); got != "https://mcp.dingtalk.com/dws" {
t.Fatalf("personalEventControlBaseURL() = %q, want production control URL", got)
}
if got := personalEventStreamTicketURL("", dir); got != "https://mcp.dingtalk.com/stream/connections/ticket" {
t.Fatalf("personalEventStreamTicketURL() = %q, want production ticket URL", got)
}
if got := personalEventStreamSourceID(""); got != "open" {
t.Fatalf("personalEventStreamSourceID() = %q, want open", got)
}
if got := config.GetMCPBaseURL(); got != "https://mcp.dingtalk.com" {
t.Fatalf("config.GetMCPBaseURL() = %q, want production MCP URL", got)
}
}
func TestPersonalEventDefaultsRespectExplicitAndMCPConfig(t *testing.T) {
dir := t.TempDir()
if err := os.WriteFile(filepath.Join(dir, "mcp_url"), []byte("https://custom-mcp.example.com\n"), 0o600); err != nil {
t.Fatalf("write mcp_url: %v", err)
}
if got := personalEventControlBaseURL("", dir); got != "https://custom-mcp.example.com/dws" {
t.Fatalf("personalEventControlBaseURL() = %q, want configured control URL", got)
}
if got := personalEventStreamTicketURL("", dir); got != "https://custom-mcp.example.com/stream/connections/ticket" {
t.Fatalf("personalEventStreamTicketURL() = %q, want configured ticket URL", got)
}
if got := personalEventControlBaseURL(" https://override.example.com/dws/ ", dir); got != "https://override.example.com/dws" {
t.Fatalf("explicit control URL = %q, want trimmed override", got)
}
if got := personalEventStreamTicketURL(" https://override.example.com/ticket/ ", dir); got != "https://override.example.com/ticket" {
t.Fatalf("explicit ticket URL = %q, want trimmed override", got)
}
if got := personalEventStreamSourceID("flag_source"); got != "flag_source" {
t.Fatalf("explicit sourceID = %q, want flag_source", got)
}
}
func TestPersonalEventSourceIDPrefersEditionOverride(t *testing.T) {
prev := edition.Get()
edition.Override(&edition.Hooks{PersonalEventSourceID: "edition_source"})
t.Cleanup(func() { edition.Override(prev) })
if got := personalEventStreamSourceID(""); got != "edition_source" {
t.Fatalf("personalEventStreamSourceID() = %q, want edition_source", got)
}
if got := personalEventStreamSourceID("flag_source"); got != "flag_source" {
t.Fatalf("explicit sourceID = %q, want flag_source", got)
}
}
func setupPersonalIdentityToken(t *testing.T, data *authpkg.TokenData) string {
t.Helper()
configDir := t.TempDir()
raw, err := json.Marshal(data)
if err != nil {
t.Fatalf("marshal token data: %v", err)
}
prev := edition.Get()
edition.Override(&edition.Hooks{
LoadToken: func(dir string) ([]byte, error) {
if filepath.Clean(dir) != filepath.Clean(configDir) {
t.Fatalf("LoadToken dir = %q, want %q", dir, configDir)
}
return raw, nil
},
})
t.Cleanup(func() { edition.Override(prev) })
return configDir
}
+349
View File
@@ -0,0 +1,349 @@
// 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"
"encoding/json"
"strings"
"testing"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/personal"
"github.com/spf13/cobra"
)
func TestPersonalEventListHidesSchemaIDs(t *testing.T) {
for _, tc := range []struct {
name string
args []string
}{
{name: "table", args: []string{"--as", "user"}},
{name: "json", args: []string{"--as", "user", "--format", "json"}},
} {
t.Run(tc.name, func(t *testing.T) {
cmd := newEventListCommand()
cmd.SilenceUsage = true
cmd.SilenceErrors = true
var out bytes.Buffer
cmd.SetOut(&out)
cmd.SetArgs(tc.args)
if err := cmd.Execute(); err != nil {
t.Fatalf("Execute() error = %v", err)
}
got := out.String()
assertPersonalOutputHidesSchemaIDs(t, got)
if strings.Contains(got, personal.EventFromUser) {
t.Fatalf("list output exposed hidden event %s: %s", personal.EventFromUser, got)
}
})
}
}
func TestEventListDefaultsToUser(t *testing.T) {
t.Setenv("DWS_CONFIG_DIR", t.TempDir())
cmd := newEventListCommand()
cmd.SilenceUsage = true
cmd.SilenceErrors = true
var out bytes.Buffer
cmd.SetOut(&out)
if err := cmd.Execute(); err != nil {
t.Fatalf("Execute() error = %v", err)
}
got := out.String()
if !strings.Contains(got, personal.EventSingleChat) || !strings.Contains(got, "EVENT_KEY") {
t.Fatalf("list output = %s, want personal event catalog", got)
}
if strings.Contains(got, personal.EventFromUser) {
t.Fatalf("list output exposed hidden event %s: %s", personal.EventFromUser, got)
}
if strings.Contains(got, "CLIENT_ID") || strings.Contains(got, "ClientSecret") {
t.Fatalf("list default appears to use legacy application output: %s", got)
}
}
func TestEventPublicHelpHidesAppMode(t *testing.T) {
for _, tc := range []struct {
name string
cmd *cobra.Command
}{
{name: "consume", cmd: newEventConsumeCommand()},
{name: "list", cmd: newEventListCommand()},
{name: "schema", cmd: newEventSchemaCommand()},
{name: "status", cmd: newEventStatusCommand()},
{name: "stop", cmd: newEventStopCommand()},
} {
t.Run(tc.name, func(t *testing.T) {
var out bytes.Buffer
tc.cmd.SetOut(&out)
tc.cmd.SetArgs([]string{"--help"})
if tc.name == "schema" {
tc.cmd.SetArgs([]string{personal.EventSingleChat, "--help"})
}
if err := tc.cmd.Execute(); err != nil {
t.Fatalf("Execute() error = %v", err)
}
got := out.String()
for _, hidden := range []string{"--as", "user|app", "应用事件" + " Stream"} {
if strings.Contains(got, hidden) {
t.Fatalf("%s help leaked %q:\n%s", tc.name, hidden, got)
}
}
})
}
}
func TestEventListAppOnlyFlagsRejectedForPersonalEvents(t *testing.T) {
cmd := newEventListCommand()
cmd.SilenceUsage = true
cmd.SilenceErrors = true
cmd.SetArgs([]string{"--all"})
err := cmd.Execute()
if err == nil || !strings.Contains(err.Error(), "--all are not supported for personal events") {
t.Fatalf("Execute() error = %v, want unsupported flag validation", err)
}
}
func TestEventAsAppRejected(t *testing.T) {
t.Setenv("DWS_CONFIG_DIR", t.TempDir())
for _, cmd := range []*cobra.Command{
newEventListCommand(),
newEventStatusCommand(),
newEventConsumeCommand(),
newEventStopCommand(),
newEventSchemaCommand(),
} {
cmd.SilenceUsage = true
cmd.SilenceErrors = true
cmd.SetArgs([]string{"--as", "app"})
if cmd.Use == "schema <event_key>" {
cmd.SetArgs([]string{personal.EventSingleChat, "--as", "app"})
}
err := cmd.Execute()
if err == nil || !strings.Contains(err.Error(), "app event is not publicly available yet") {
t.Fatalf("%s Execute() error = %v, want public availability guard", cmd.Use, err)
}
}
}
func TestEventStatusAppOnlyFlagsRejectedForPersonalEvents(t *testing.T) {
cmd := newEventStatusCommand()
cmd.SilenceUsage = true
cmd.SilenceErrors = true
cmd.SetArgs([]string{"--all", "--fail-on-orphan"})
err := cmd.Execute()
if err == nil ||
!strings.Contains(err.Error(), "--all") ||
!strings.Contains(err.Error(), "--fail-on-orphan") ||
!strings.Contains(err.Error(), "not supported for personal events") {
t.Fatalf("Execute() error = %v, want unsupported flag validation", err)
}
}
func TestPersonalEventSchemaHidesSchemaIDs(t *testing.T) {
for _, tc := range []struct {
name string
args []string
}{
{name: "default", args: []string{personal.EventSingleChat, "--as", "user"}},
{name: "json", args: []string{personal.EventSingleChat, "--as", "user", "--format", "json"}},
} {
t.Run(tc.name, func(t *testing.T) {
cmd := newEventSchemaCommand()
cmd.SilenceUsage = true
cmd.SilenceErrors = true
var out bytes.Buffer
cmd.SetOut(&out)
cmd.SetArgs(tc.args)
if err := cmd.Execute(); err != nil {
t.Fatalf("Execute() error = %v", err)
}
assertPersonalOutputHidesSchemaIDs(t, out.String())
if strings.Contains(out.String(), "Schemas") {
t.Fatalf("schema output contains Schemas line: %s", out.String())
}
})
}
}
func TestPersonalEventSchemaUsesSingleJSONSchema(t *testing.T) {
for _, eventKey := range []string{
personal.EventMention,
personal.EventSingleChat,
personal.EventInChat,
} {
t.Run(eventKey, func(t *testing.T) {
cmd := newEventSchemaCommand()
cmd.SilenceUsage = true
cmd.SilenceErrors = true
var out bytes.Buffer
cmd.SetOut(&out)
cmd.SetArgs([]string{eventKey})
if err := cmd.Execute(); err != nil {
t.Fatalf("Execute() error = %v", err)
}
got := out.String()
var doc map[string]any
if err := json.Unmarshal(out.Bytes(), &doc); err != nil {
t.Fatalf("schema output for %s is not JSON: %v\n%s", eventKey, err, got)
}
for _, want := range []string{
"event_key",
"display_name",
"description",
"category",
"rule_type",
"required_params",
"jq_root_path",
"schema",
"event_id",
"timestamp",
"subscribe_id",
"content",
"sender",
"sender_open_dingtalk_id",
"conversation_id",
"message_id",
"create_time",
"event_time",
} {
if !strings.Contains(got, want) {
t.Fatalf("schema output for %s missing %q: %s", eventKey, want, got)
}
}
for _, leaked := range []string{
"message.text",
"chat.openConversationId",
"sender.userId",
"sender.unionId",
"auth",
"resolved_output_schema",
"decoded_data_schema",
"filter_schema",
"payload_schema",
"output_schema",
"data_json_path",
"headers",
"audit",
"tenant",
"subject",
"traceId",
"msgIdMetaq",
"at_users",
"sender_user_id",
} {
if strings.Contains(got, leaked) {
t.Fatalf("schema output for %s leaked %q: %s", eventKey, leaked, got)
}
}
if doc["jq_root_path"] != ".data | fromjson" {
t.Fatalf("jq_root_path = %#v, want .data | fromjson", doc["jq_root_path"])
}
schema, ok := doc["schema"].(map[string]any)
if !ok {
t.Fatalf("schema = %#v, want object", doc["schema"])
}
props, ok := schema["properties"].(map[string]any)
if !ok {
t.Fatalf("schema.properties = %#v, want object", schema["properties"])
}
if _, ok := props["content"].(map[string]any); !ok {
t.Fatalf("schema.properties.content = %#v, want object", props["content"])
}
})
}
}
func TestEventSchemaDefaultsToUser(t *testing.T) {
cmd := newEventSchemaCommand()
cmd.SilenceUsage = true
cmd.SilenceErrors = true
var out bytes.Buffer
cmd.SetOut(&out)
cmd.SetArgs([]string{personal.EventSingleChat})
if err := cmd.Execute(); err != nil {
t.Fatalf("Execute() error = %v", err)
}
var doc map[string]any
if err := json.Unmarshal(out.Bytes(), &doc); err != nil {
t.Fatalf("schema output is not JSON: %v\n%s", err, out.String())
}
if doc["event_key"] != personal.EventSingleChat {
t.Fatalf("event_key = %#v, want %s", doc["event_key"], personal.EventSingleChat)
}
}
func TestPersonalEventFromUserIsNotPubliclyAvailable(t *testing.T) {
for _, tc := range []struct {
name string
cmd *cobra.Command
args []string
}{
{
name: "schema",
cmd: newEventSchemaCommand(),
args: []string{personal.EventFromUser},
},
{
name: "consume",
cmd: newEventConsumeCommand(),
args: []string{personal.EventFromUser, "--user", "507971", "--dry-run"},
},
{
name: "status",
cmd: newEventStatusCommand(),
args: []string{"--event", personal.EventFromUser},
},
} {
t.Run(tc.name, func(t *testing.T) {
tc.cmd.SilenceUsage = true
tc.cmd.SilenceErrors = true
tc.cmd.SetArgs(tc.args)
err := tc.cmd.Execute()
if err == nil || !strings.Contains(err.Error(), "event "+personal.EventFromUser+" is not publicly available yet") {
t.Fatalf("Execute() error = %v, want not publicly available", err)
}
})
}
}
func TestPersonalEventSchemaRejectsTableFormat(t *testing.T) {
cmd := newEventSchemaCommand()
cmd.SilenceUsage = true
cmd.SilenceErrors = true
cmd.SetArgs([]string{personal.EventSingleChat, "--format", "table"})
err := cmd.Execute()
if err == nil || !strings.Contains(err.Error(), "event schema only supports json output") {
t.Fatalf("Execute() error = %v, want json-only format validation", err)
}
}
func TestEventAsBotRejected(t *testing.T) {
cmd := newEventListCommand()
cmd.SilenceUsage = true
cmd.SilenceErrors = true
cmd.SetArgs([]string{"--as", "bot"})
err := cmd.Execute()
if err == nil || !strings.Contains(err.Error(), "app event is not publicly available yet") {
t.Fatalf("Execute() error = %v, want public availability guard", err)
}
}
func assertPersonalOutputHidesSchemaIDs(t *testing.T, out string) {
t.Helper()
for _, leaked := range []string{"SCHEMA_IDS", "schema_ids", "im_msg_23", "im_msg_29"} {
if strings.Contains(out, leaked) {
t.Fatalf("output leaked %q: %s", leaked, out)
}
}
}
+105
View File
@@ -0,0 +1,105 @@
// 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"
"strings"
"testing"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/busctl"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/personal"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/transport"
)
func TestRenderPersonalStatusTextShowsConsumersWithoutSubscriptions(t *testing.T) {
var out bytes.Buffer
renderPersonalStatusText(&out, personal.Identity{
CorpID: "corp-1",
UserID: "user-1",
ClientID: "client-1",
SourceID: "source-1",
}, "identity-hash-1", nil, busctl.EntryStatus{
Entry: busctl.BusEntry{
WorkDir: "wd",
State: busctl.BusStateRunning,
HolderPID: 100,
},
Live: &transport.StatusResp{
Consumers: []transport.StatusConsumer{
{
PID: 12345,
EventTypes: []string{"user_im_message_receive_o2o"},
SubscribeID: "subId-1",
Filter: "content",
Received: 3,
Dropped: 1,
},
{
PID: 12346,
Received: 5,
},
},
},
})
got := out.String()
for _, want := range []string{
"Subscriptions: none",
"Consumers:",
"PID",
"EVENT_KEYS",
"SUBSCRIBE_ID",
"RECEIVED",
"DROPPED",
"12345",
"user_im_message_receive_o2o",
"subId-1",
"content",
"3",
"1",
"(catch-all)",
"-",
} {
if !strings.Contains(got, want) {
t.Fatalf("status output missing %q:\n%s", want, got)
}
}
}
func TestRenderPersonalStatusTextConsumersUnavailableWhenRPCFails(t *testing.T) {
var out bytes.Buffer
renderPersonalStatusText(&out, personal.Identity{ClientID: "client-1", SourceID: "source-1"}, "identity-hash-1", nil, busctl.EntryStatus{
Entry: busctl.BusEntry{
WorkDir: "wd",
State: busctl.BusStateRunning,
HolderPID: 100,
},
})
if got := out.String(); !strings.Contains(got, "Consumers: unavailable (status RPC failed)") {
t.Fatalf("status output = %q, want unavailable consumers", got)
}
}
func TestRenderPersonalStatusTextConsumersNoneWhenBusNotRunning(t *testing.T) {
var out bytes.Buffer
renderPersonalStatusText(&out, personal.Identity{ClientID: "client-1", SourceID: "source-1"}, "identity-hash-1", nil, busctl.EntryStatus{
Entry: busctl.BusEntry{
WorkDir: "wd",
State: busctl.BusStateNotRunning,
},
})
if got := out.String(); !strings.Contains(got, "Consumers: none") {
t.Fatalf("status output = %q, want no consumers", got)
}
}
+127
View File
@@ -0,0 +1,127 @@
// 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"
"reflect"
"strings"
"testing"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/personal"
)
func TestEventStopHelpDescribesPersonalSubscription(t *testing.T) {
cmd := newEventStopCommand()
var out bytes.Buffer
cmd.SetOut(&out)
cmd.SetArgs([]string{"--help"})
if err := cmd.Execute(); err != nil {
t.Fatalf("Execute() error = %v", err)
}
got := out.String()
for _, want := range []string{
"stop [subscribe_id]",
"取消个人事件订阅并停止本地消费",
"取消个人事件订阅并停止本地消费,清理对应本地消费状态",
} {
if !strings.Contains(got, want) {
t.Fatalf("help missing %q:\n%s", want, got)
}
}
for _, stale := range []string{"优雅停止 bus 守护进程", strings.Join([]string{"--as", "app"}, " "), "应用事件"} {
if strings.Contains(got, stale) {
t.Fatalf("help still contains stale public app wording %q:\n%s", stale, got)
}
}
}
func TestEventStopRequiresSubscribeIDOrAll(t *testing.T) {
cmd := newEventStopCommand()
cmd.SilenceUsage = true
cmd.SilenceErrors = true
err := cmd.Execute()
if err == nil || !strings.Contains(err.Error(), "subscribe_id is required unless --all is set") {
t.Fatalf("Execute() error = %v, want subscribe_id requirement", err)
}
}
func TestEventStopSubscribeIDAndAllAreMutuallyExclusive(t *testing.T) {
cmd := newEventStopCommand()
cmd.SilenceUsage = true
cmd.SilenceErrors = true
cmd.SetArgs([]string{"subId-1", "--all"})
err := cmd.Execute()
if err == nil || !strings.Contains(err.Error(), "subscribe_id and --all are mutually exclusive") {
t.Fatalf("Execute() error = %v, want mutual exclusion", err)
}
}
func TestEventStopAsAppRejectsSubscribeID(t *testing.T) {
cmd := newEventStopCommand()
cmd.SilenceUsage = true
cmd.SilenceErrors = true
cmd.SetArgs([]string{"--as", "app", "subId-1"})
err := cmd.Execute()
if err == nil || !strings.Contains(err.Error(), "app event is not publicly available yet") {
t.Fatalf("Execute() error = %v, want public availability guard", err)
}
}
func TestPersonalStopTargets(t *testing.T) {
workDir := t.TempDir()
if err := personal.UpsertRunState(workDir, personal.RunState{SubscribeID: "sub-b"}); err != nil {
t.Fatalf("UpsertRunState() error = %v", err)
}
if err := personal.UpsertRunState(workDir, personal.RunState{SubscribeID: "sub-a"}); err != nil {
t.Fatalf("UpsertRunState() error = %v", err)
}
got, err := personalStopTargets(workDir, "sub-explicit", false)
if err != nil {
t.Fatalf("personalStopTargets(explicit) error = %v", err)
}
if want := []string{"sub-explicit"}; !reflect.DeepEqual(got, want) {
t.Fatalf("explicit targets = %#v, want %#v", got, want)
}
got, err = personalStopTargets(workDir, "", true)
if err != nil {
t.Fatalf("personalStopTargets(all) error = %v", err)
}
if want := []string{"sub-a", "sub-b"}; !reflect.DeepEqual(got, want) {
t.Fatalf("all targets = %#v, want %#v", got, want)
}
if _, err := personalStopTargets(workDir, "", false); err == nil || !strings.Contains(err.Error(), "subscribe_id is required unless --all is set") {
t.Fatalf("personalStopTargets(no target) error = %v, want required error", err)
}
if _, err := personalStopTargets(workDir, "sub-explicit", true); err == nil || !strings.Contains(err.Error(), "mutually exclusive") {
t.Fatalf("personalStopTargets(explicit+all) error = %v, want mutual exclusion", err)
}
}
func TestPrintPersonalStopResult(t *testing.T) {
var out bytes.Buffer
printPersonalStopResult(&out, []string{"sub-1"}, true, "personal bus stopped")
if got := out.String(); got != "cancelled personal subscription sub-1; personal bus stopped\n" {
t.Fatalf("single output = %q", got)
}
out.Reset()
printPersonalStopResult(&out, []string{"sub-1", "sub-2"}, false, "personal bus still running")
if got := out.String(); got != "cancelled 2 personal subscription(s); personal bus still running\n" {
t.Fatalf("multi output = %q", got)
}
}
+1
View File
@@ -356,6 +356,7 @@ func NewRootCommandWithEngine(rootCtx context.Context, engine *pipeline.Engine)
newCatalogCommand(loader),
newConfigCommand(),
newDoctorCommand(),
newEventCommand(),
newCompletionCommand(root),
newRecoveryCommand(rootCtx, loader, flags),
newUpgradeCommand(),
+1 -2
View File
@@ -216,8 +216,7 @@ func normalizeMultiSkillName(name string) string {
return multiSkillPrefix + n
}
// filterMultiSkillNames narrows `all` by include / exclude lists.
// Semantics mirror lark-cli's `npx skills add -s lark-calendar`:
// filterMultiSkillNames narrows `all` by include / exclude lists:
//
// - include + exclude are mutually exclusive (both → error)
// - names accept short or full form; normalized before matching
+1 -2
View File
@@ -445,8 +445,7 @@ func TestFilterMultiSkillNames(t *testing.T) {
// TestSkillSetupMultiAdditivePreservesSiblings verifies the key UX promise of
// `dws skill setup --mode multi -s aitable`: installing a subset must NOT
// touch already-installed dingtalk-* siblings (additive semantics, matches
// lark-cli `npx skills add -s lark-calendar`).
// touch already-installed dingtalk-* siblings (additive semantics).
func TestSkillSetupMultiAdditivePreservesSiblings(t *testing.T) {
src := writeMultiSkillSource(t, []string{
"dingtalk-aitable", "dingtalk-calendar", "dingtalk-doc",
+157
View File
@@ -0,0 +1,157 @@
// 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 (
"errors"
"fmt"
"os"
)
// CredentialSource identifies where a particular credential field
// (ClientID or ClientSecret) was loaded from. It is exposed in
// `dws event status` and the HelloAck IPC frame so users can verify which
// credential channel is actually in use — important because env vars,
// keychain, and config file can all coexist and silently override each
// other (see plan §1 决策 "凭证来源拆字段").
type CredentialSource string
const (
CredentialSourceUnknown CredentialSource = "unknown"
CredentialSourceEnv CredentialSource = "env"
CredentialSourceAppConfig CredentialSource = "app_config" // value pulled from app config (plain or SecretRef metadata)
CredentialSourceKeychain CredentialSource = "keychain" // SecretRef resolved through OS keychain
CredentialSourcePlainConfig CredentialSource = "plain_config" // SecretInput stored as plaintext in config file (insecure but supported)
)
// Strict resolver error sentinels. Use errors.Is to distinguish failure
// modes; see plan §8 strict resolver decision (4 classes).
var (
// ErrAppConfigMissing — no app config file on disk AND no env-var
// credentials present. Prompt the user to either `dws config init` or
// set DWS_CLIENT_ID + DWS_CLIENT_SECRET.
ErrAppConfigMissing = errors.New("app config missing: run `dws config init` or set DWS_CLIENT_ID/DWS_CLIENT_SECRET env vars")
// ErrClientIDEmpty — neither env nor config supplies a non-empty ClientID.
ErrClientIDEmpty = errors.New("ClientID is empty")
// ErrClientSecretEmpty — there's a ClientID but ClientSecret resolved to "".
ErrClientSecretEmpty = errors.New("ClientSecret is empty")
// ErrSecretResolve — the secret-resolution backend (keychain) failed
// unrecoverably. Typically headless Linux without gnome-keyring, locked
// macOS keychain, or CI sandboxes. Suggest the env-var fallback.
ErrSecretResolve = errors.New("ClientSecret resolution failed (keychain unavailable?); try DWS_CLIENT_ID/DWS_CLIENT_SECRET env vars")
)
// Env var names used by the env fallback channel. Must be set as a pair —
// any single-variable configuration is rejected so users cannot accidentally
// "set the env half-way" and silently fall back to keychain.
const (
EnvClientID = "DWS_CLIENT_ID"
EnvClientSecret = "DWS_CLIENT_SECRET"
)
// ResolveAppCredentialsStrict is the credentials channel used by the event
// subsystem (and by future commands that need fine-grained failure
// reporting). It distinguishes 4 failure classes and reports the source of
// each successfully-resolved field separately.
//
// Resolution order:
// 1. Env var override: if BOTH DWS_CLIENT_ID and DWS_CLIENT_SECRET are
// set non-empty, use them as a pair and skip keychain/config entirely.
// Single-variable configuration is detected and reported via the
// EnvHalfSet flag in the warning channel (callers MAY log a warning).
// 2. App config from disk:
// - ClientID from cfg.ClientID
// - ClientSecret from ResolveSecret(cfg.ClientSecret):
// - SecretInput.IsPlain() → CredentialSourcePlainConfig
// - SecretRef → CredentialSourceKeychain (or whatever Ref.Source says)
//
// Empty returns: clientID and secret may be empty when err is non-nil;
// callers must NOT use them in that case.
func ResolveAppCredentialsStrict(configDir string) (
clientID, secret string,
clientIDSource, secretSource CredentialSource,
err error,
) {
// Step 1: env var fallback (atomic pair)
envID := os.Getenv(EnvClientID)
envSecret := os.Getenv(EnvClientSecret)
if envID != "" && envSecret != "" {
return envID, envSecret, CredentialSourceEnv, CredentialSourceEnv, nil
}
// Note: if only one of the two is set we explicitly do NOT use it.
// The half-set warning is surfaced via EnvHalfSet() so the CLI can
// stderr-warn the user during preflight.
// Step 2: app config from disk
cfg, loadErr := LoadAppConfig(configDir)
if loadErr != nil {
return "", "", CredentialSourceUnknown, CredentialSourceUnknown,
fmt.Errorf("load app config: %w", loadErr)
}
if cfg == nil {
return "", "", CredentialSourceUnknown, CredentialSourceUnknown, ErrAppConfigMissing
}
if cfg.ClientID == "" {
return "", "", CredentialSourceUnknown, CredentialSourceUnknown, ErrClientIDEmpty
}
clientID = cfg.ClientID
clientIDSource = CredentialSourceAppConfig
// Resolve secret. Source depends on the SecretInput shape:
// - IsPlain (no Ref) → it's stored as plaintext in the config file
// - has Ref → it's a SecretRef pointing at keychain/file
wasPlain := cfg.ClientSecret.IsPlain()
resolved, resolveErr := ResolveSecret(cfg.ClientSecret)
if resolveErr != nil {
return "", "", CredentialSourceUnknown, CredentialSourceUnknown,
fmt.Errorf("%w: %v", ErrSecretResolve, resolveErr)
}
if resolved == "" {
return "", "", clientIDSource, CredentialSourceUnknown, ErrClientSecretEmpty
}
secret = resolved
if wasPlain {
secretSource = CredentialSourcePlainConfig
} else {
// For SecretRef we map Source verbatim (keychain / file / future)
switch cfg.ClientSecret.Ref.Source {
case "keychain":
secretSource = CredentialSourceKeychain
default:
// File-backed secrets share the "plain_config" category from
// the consumer's perspective: stored as readable bytes outside
// keychain. Status output renders them as "plain_config" so
// users see "secret is not in keychain".
secretSource = CredentialSourcePlainConfig
}
}
return clientID, secret, clientIDSource, secretSource, nil
}
// EnvHalfSet reports whether exactly one of (DWS_CLIENT_ID, DWS_CLIENT_SECRET)
// is set. Used by CLI preflight to emit a clear stderr warning of the form:
//
// WARN: DWS_CLIENT_ID is set but DWS_CLIENT_SECRET is not — env fallback
// disabled; using keychain/app config. Set both or unset both to
// avoid this warning.
//
// The strict resolver itself does NOT log; logging is the caller's job.
func EnvHalfSet() bool {
id := os.Getenv(EnvClientID) != ""
secret := os.Getenv(EnvClientSecret) != ""
return id != secret
}
+165
View File
@@ -0,0 +1,165 @@
// 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 (
"encoding/json"
"errors"
"os"
"path/filepath"
"testing"
"time"
)
// resetStrictResolverState clears caches the strict resolver shares with
// the existing legacy resolver. Tests must call this between scenarios
// because GetCachedAppConfig and the resolved-credential cache outlive
// individual t.TempDir setups.
func resetStrictResolverState(t *testing.T) {
t.Helper()
cachedAppConfigMu.Lock()
cachedAppConfig = nil
cachedAppConfigMu.Unlock()
cachedResolvedMu.Lock()
cachedResolvedValid = false
cachedResolvedID = ""
cachedResolvedSecret = ""
cachedResolvedMu.Unlock()
}
// writeAppConfig drops a config JSON into dir. clientSecret == "" produces
// the legacy "no SecretInput field" shape (treated as empty).
func writeAppConfig(t *testing.T, dir, clientID, clientSecret string) {
t.Helper()
cfg := AppConfig{
ClientID: clientID,
CreatedAt: time.Now(),
}
if clientSecret != "" {
cfg.ClientSecret = PlainSecret(clientSecret)
}
path := GetAppConfigPath(dir)
b, err := json.MarshalIndent(cfg, "", " ")
if err != nil {
t.Fatalf("marshal: %v", err)
}
if err := os.WriteFile(path, b, 0o600); err != nil {
t.Fatalf("write: %v", err)
}
}
func unsetEnv(t *testing.T) {
t.Helper()
t.Setenv(EnvClientID, "")
t.Setenv(EnvClientSecret, "")
_ = os.Unsetenv(EnvClientID)
_ = os.Unsetenv(EnvClientSecret)
}
func TestResolveStrict_AppConfigMissing(t *testing.T) {
resetStrictResolverState(t)
unsetEnv(t)
dir := t.TempDir()
_, _, _, _, err := ResolveAppCredentialsStrict(dir)
if !errors.Is(err, ErrAppConfigMissing) {
t.Fatalf("err = %v, want ErrAppConfigMissing", err)
}
}
func TestResolveStrict_ClientIDEmpty(t *testing.T) {
resetStrictResolverState(t)
unsetEnv(t)
dir := t.TempDir()
writeAppConfig(t, dir, "", "some-secret")
_, _, _, _, err := ResolveAppCredentialsStrict(dir)
if !errors.Is(err, ErrClientIDEmpty) {
t.Fatalf("err = %v, want ErrClientIDEmpty", err)
}
}
func TestResolveStrict_ClientSecretEmpty(t *testing.T) {
resetStrictResolverState(t)
unsetEnv(t)
dir := t.TempDir()
writeAppConfig(t, dir, "ding_abc", "")
_, _, _, _, err := ResolveAppCredentialsStrict(dir)
if !errors.Is(err, ErrClientSecretEmpty) {
t.Fatalf("err = %v, want ErrClientSecretEmpty", err)
}
}
func TestResolveStrict_PlainConfigSuccess(t *testing.T) {
resetStrictResolverState(t)
unsetEnv(t)
dir := t.TempDir()
writeAppConfig(t, dir, "ding_abc", "supersecret123")
id, secret, idSrc, secretSrc, err := ResolveAppCredentialsStrict(dir)
if err != nil {
t.Fatalf("unexpected err: %v", err)
}
if id != "ding_abc" {
t.Errorf("id = %q", id)
}
if secret != "supersecret123" {
t.Errorf("secret = %q", secret)
}
if idSrc != CredentialSourceAppConfig {
t.Errorf("idSrc = %s, want app_config", idSrc)
}
if secretSrc != CredentialSourcePlainConfig {
t.Errorf("secretSrc = %s, want plain_config (PlainSecret was used)", secretSrc)
}
}
func TestResolveStrict_SecretRefFileSuccess(t *testing.T) {
resetStrictResolverState(t)
unsetEnv(t)
dir := t.TempDir()
// Write secret file
secretPath := filepath.Join(dir, "secret.txt")
if err := os.WriteFile(secretPath, []byte("via-file-secret\n"), 0o600); err != nil {
t.Fatal(err)
}
// Write app config with file SecretRef
cfg := AppConfig{
ClientID: "ding_abc",
ClientSecret: SecretInput{
Ref: &SecretRef{Source: "file", ID: secretPath},
},
CreatedAt: time.Now(),
}
b, _ := json.MarshalIndent(cfg, "", " ")
if err := os.WriteFile(GetAppConfigPath(dir), b, 0o600); err != nil {
t.Fatal(err)
}
id, secret, idSrc, secretSrc, err := ResolveAppCredentialsStrict(dir)
if err != nil {
t.Fatalf("unexpected err: %v", err)
}
if id != "ding_abc" || secret != "via-file-secret" {
t.Errorf("id/secret = %q/%q", id, secret)
}
if idSrc != CredentialSourceAppConfig {
t.Errorf("idSrc = %s", idSrc)
}
// File-backed secrets are reported as plain_config (not in keychain).
if secretSrc != CredentialSourcePlainConfig {
t.Errorf("secretSrc = %s, want plain_config for file-backed secret", secretSrc)
}
}
+525
View File
@@ -0,0 +1,525 @@
// 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 bus
import (
"context"
"errors"
"fmt"
"io"
"log/slog"
"net"
"os"
"path/filepath"
"sync"
"sync/atomic"
"time"
dwsevent "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/dedup"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/transport"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
)
// SourceAdapter is the interface daemon.go uses to talk to the cloud Source
// (in practice internal/event/source.DingtalkSource). Kept abstract so the
// bus daemon can be tested without spinning up the real Stream SDK; the
// integration test substitutes a fake.
type SourceAdapter interface {
// Start opens the cloud connection and blocks until ctx is cancelled
// or a fatal error occurs. emit is called for each incoming event.
Start(ctx context.Context, emit dwsevent.EmitFn) error
}
// Config bundles everything Run needs to start a daemon. All paths and
// identifiers come from busctl/Source so the bus itself stays oblivious
// to ConfigDir / edition rules.
type Config struct {
// WorkDir is the bus working directory:
// <ConfigDir>/events/<edition>/<source_kind>/<identity_hash>/
// The caller MUST mkdir this with pkg/config.DirPerm before calling Run.
WorkDir string
// IPCEndpoint is the Unix socket path or Windows pipe name. Caller
// computes this from WorkDir (Unix) or edition/clientIDHash (Windows).
IPCEndpoint string
// ClientID is the human-readable identifier written into bus.meta and
// status output. NOT used in any path.
ClientID string
// SourceKind/IdentityHash/SourceID are diagnostic identity fields used by
// list/status. Empty SourceKind is interpreted as app_stream for backward
// compatibility.
SourceKind dwsevent.SourceKind
IdentityHash string
SourceID string
// Edition is written into bus.meta. Comes from edition.Get().Name with
// "open" fallback applied by the caller.
Edition string
// SDKVersion is recorded in bus.meta for diagnostics.
SDKVersion string
// Source is the cloud adapter. Required.
Source SourceAdapter
// IdleTimeout: bus self-exits after this long with zero consumers.
// Zero disables (bus runs until SIGTERM).
IdleTimeout time.Duration
// ConsumerBuffer overrides per-consumer sendCh capacity. Zero uses
// DefaultSendBuffer.
ConsumerBuffer int
// DedupCapacity overrides event_id LRU size. Zero uses dedup.DefaultCapacity.
DedupCapacity int
// DropWarnPercent is the per-event-type drop-rate threshold (whole
// percentage points) that triggers a slog WARN in bus.log. Zero or
// out-of-range values fall back to DefaultDropWarnPercent. Overridable
// via env DWS_EVENT_DROP_WARN_PCT (read by the cobra layer).
DropWarnPercent int
// ReadyPipe receives a single byte ('R' on success, 'E' on failure)
// once the bus has either come up or failed startup, so the parent
// process forked by busctl/spawn can stop polling and either dial or
// surface the error. nil disables (foreground mode).
ReadyPipe *os.File
// Logger sink. Nil → slog.Default.
Logger *slog.Logger
}
// Run starts the bus daemon. Lifecycle (plan §4 invariant #6):
// 1. Acquire bus.lock (single-instance enforcement)
// 2. Write bus.meta
// 3. Listen IPC (so consumers can connect before SDK starts pushing)
// 4. Signal readiness via ReadyPipe
// 5. Start the Source (cloud SDK); concurrent with consumer accept loop
// 6. Wait on ctx for shutdown signal
// 7. Graceful: broadcast Bye → close listener → close source → release lock
//
// Run blocks until ctx is cancelled, the Source returns an error, or a fatal
// startup error occurs.
func Run(ctx context.Context, cfg Config) error {
if cfg.Source == nil {
return errors.New("bus: Source is required")
}
if cfg.Logger == nil {
cfg.Logger = slog.Default()
}
log := cfg.Logger.With("component", "bus", "client_id", cfg.ClientID, "edition", cfg.Edition)
// 1. Acquire bus.lock
if err := os.MkdirAll(cfg.WorkDir, config.DirPerm); err != nil {
return failReady(cfg.ReadyPipe, fmt.Errorf("bus: mkdir workdir: %w", err))
}
lockPath := filepath.Join(cfg.WorkDir, LockFileName)
lock, err := Acquire(lockPath)
if err != nil {
return failReady(cfg.ReadyPipe, fmt.Errorf("bus: acquire lock: %w", err))
}
defer lock.Close()
// 2. Write bus.meta
meta := Meta{
ClientID: cfg.ClientID,
Edition: cfg.Edition,
SourceKind: cfg.SourceKind,
IdentityHash: cfg.IdentityHash,
SourceID: cfg.SourceID,
StartedAt: time.Now().UTC(),
SDKVersion: cfg.SDKVersion,
BusPID: os.Getpid(),
}
if err := WriteMeta(cfg.WorkDir, meta); err != nil {
return failReady(cfg.ReadyPipe, fmt.Errorf("bus: write meta: %w", err))
}
// 3. IPC listen
listener, err := transport.Listen(cfg.IPCEndpoint)
if err != nil {
return failReady(cfg.ReadyPipe, fmt.Errorf("bus: ipc listen: %w", err))
}
defer listener.Close()
hub := NewHub(cfg.ConsumerBuffer)
dd := dedup.NewWithCapacity(cfg.DedupCapacity)
d := &daemon{
cfg: cfg,
log: log,
lock: lock,
listener: listener,
hub: hub,
dedup: dd,
started: time.Now().UTC(),
idleStop: make(chan struct{}),
}
// 4. Signal ready BEFORE accepting consumers (avoids a slow-fork
// scenario where the parent thinks bus is dead but it's actually mid-
// startup). Source.Start hasn't yet pulled events from the cloud, but
// any consumer that connects gets queued for the first events.
signalReady(cfg.ReadyPipe)
// runCtx is a child of the caller's ctx that ALL background goroutines
// (acceptLoop / idleWatch / dropWarnWatcher / source.Start) listen on.
// On idle-timeout shutdown the parent ctx is never cancelled, so we
// cancel runCtx ourselves before waiting for the goroutines — otherwise
// dropWarnWatcher (which only exits on ctx.Done) hangs forever.
runCtx, cancelRun := context.WithCancel(ctx)
defer cancelRun()
// 5. Start accept loop and Source concurrently. runCtx cancellation
// propagates to both.
acceptDone := make(chan struct{})
go func() {
defer close(acceptDone)
d.acceptLoop(runCtx)
}()
idleDone := make(chan struct{})
go func() {
defer close(idleDone)
d.idleWatch(runCtx)
}()
// Drop-rate WARN: per-event-type back-pressure monitoring. Runs in
// the background until runCtx cancellation; emits one WARN per scan
// when a type's drop rate first crosses the threshold (hysteresis
// suppresses repeats unless the rate jumps further).
dropWarnDone := make(chan struct{})
go func() {
defer close(dropWarnDone)
dropWarnWatcher(runCtx, hub.Counters(), log, cfg.DropWarnPercent)
}()
srcErr := make(chan error, 1)
go func() {
// emit is called from inside the SDK callback goroutine. It MUST
// NOT block (plan invariant #1) — dedup + Hub.Deliver are both
// non-blocking by construction.
emit := func(raw *dwsevent.RawEvent) {
if raw == nil {
return
}
if dd.Seen(raw.DedupKey()) {
return // duplicate redelivery from the cloud
}
hub.Deliver(raw)
}
srcErr <- cfg.Source.Start(runCtx, emit)
}()
// 6. Wait for shutdown trigger.
var exitErr error
select {
case <-ctx.Done():
log.Info("bus: shutdown requested by ctx", "reason", ctx.Err())
case err := <-srcErr:
log.Error("bus: source exited", "err", err)
exitErr = err
case <-d.idleStop:
log.Info("bus: idle timeout reached, shutting down")
}
// 7. Graceful shutdown — cancel runCtx first so all background
// goroutines wake up, then close listener / drain consumers.
cancelRun()
d.shutdown()
<-acceptDone
<-idleDone
<-dropWarnDone
return exitErr
}
// daemon is the in-memory state of one bus run. Lifetime equals one Run() call.
type daemon struct {
cfg Config
log *slog.Logger
lock *Lock
listener transport.Listener
hub *Hub
dedup *dedup.LRU
started time.Time
consumerWG sync.WaitGroup // tracks live connection handler goroutines
conns sync.Map // map[net.Conn]struct{} for forced shutdown close
shutdownMu sync.Mutex
shuttingDown atomic.Bool
idleStop chan struct{}
}
// acceptLoop drives the IPC accept goroutine. Each accepted connection is
// passed to handleConnection in its own goroutine; the accept loop returns
// when the listener Close()s (typically during shutdown).
func (d *daemon) acceptLoop(ctx context.Context) {
for {
conn, err := d.listener.Accept()
if err != nil {
if d.shuttingDown.Load() {
return
}
// Transient errors: log + continue. Non-transient (listener
// closed) shows up as net.ErrClosed → also a clean exit.
if errors.Is(err, net.ErrClosed) {
return
}
d.log.Warn("bus: accept error", "err", err)
continue
}
d.consumerWG.Add(1)
go func() {
defer d.consumerWG.Done()
d.handleConnection(ctx, conn)
}()
}
}
// handleConnection processes one IPC connection's full lifecycle: read
// Hello → register with Hub → spawn writer goroutine → read until EOF/Bye.
// Always Unregisters and Closes on exit (plan invariant #5).
func (d *daemon) handleConnection(ctx context.Context, conn net.Conn) {
d.conns.Store(conn, struct{}{})
defer func() {
d.conns.Delete(conn)
conn.Close()
}()
r := transport.NewReader(conn)
w := transport.NewWriter(conn)
// Expect Hello first.
var hello transport.Hello
if err := r.ReadJSON(&hello); err != nil {
d.log.Warn("bus: hello read failed", "err", err)
return
}
if hello.Type != transport.FrameTypeHello {
d.log.Warn("bus: first frame not hello", "type", hello.Type)
return
}
// Ad-hoc tooling (status/list/stop) — short-lived RPC, no Hub register.
if hello.Role == transport.HelloRoleStatus {
d.handleStatusRPC(w, r)
return
}
if hello.Role == transport.HelloRoleStop {
// Signal shutdown by cancelling our parent ctx via shutdown().
// For now, return after acking — daemon.shutdown is wired through
// daemon's exit path (the busctl/stop command sends SIGTERM in
// addition to this RPC for v1).
_ = w.WriteJSON(transport.Bye{Type: transport.FrameTypeBye, Reason: "stop_request"})
go d.triggerShutdown("stop_request")
return
}
// Regular consumer registration
c, err := d.hub.Register(hello)
if err != nil {
d.log.Warn("bus: register failed", "err", err, "pid", hello.ConsumerPID)
_ = w.WriteJSON(transport.Bye{Type: transport.FrameTypeBye, Reason: "register_failed: " + err.Error()})
return
}
// HelloAck — credentials_source fields are filled in by the daemon
// runner (which knows from the strict resolver) and exposed via the
// adapter for forward-compat. v1 leaves them empty here; daemon.Run
// passes them through future config if the caller wishes.
idleSecs := int(d.cfg.IdleTimeout / time.Second)
if err := w.WriteJSON(transport.HelloAck{
Type: transport.FrameTypeHelloAck,
BusPID: os.Getpid(),
SourceState: "connected", // best-effort; full state machine pushed via SourceState frames
StateSource: "inferred",
IdleTimeoutSecs: idleSecs,
}); err != nil {
d.log.Warn("bus: helloack write failed", "err", err)
return
}
// Writer goroutine pulls from SendCh and writes to the wire.
writerDone := make(chan struct{})
go func() {
defer close(writerDone)
for frame := range c.SendCh {
if err := w.WriteJSON(frame); err != nil {
// Wire error: peer dead. Returning here will let the
// reader goroutine notice EOF and Unregister.
return
}
}
}()
// Reader loop: wait for Bye or EOF. Both trigger Unregister.
for {
raw, err := r.Read()
if err != nil {
if !errors.Is(err, io.EOF) {
d.log.Debug("bus: consumer read error", "pid", c.PID, "err", err)
}
break
}
typ, err := transport.PeekType(raw)
if err != nil {
continue
}
if typ == transport.FrameTypeBye {
break
}
// Heartbeat / future client→bus frames: ignored for v1.
}
// Order matters: Unregister first (closes SendCh), then wait for the
// writer goroutine to drain. The reverse order would deadlock because
// the writer loops on `range SendCh` until close, but only the Hub
// can close that channel via Unregister.
d.hub.Unregister(c.ID)
<-writerDone
_ = ctx // for future use (writer ctx-cancel propagation)
}
// handleStatusRPC services a single status_req and returns. The connection
// is closed by the caller's defer.
func (d *daemon) handleStatusRPC(w *transport.Writer, r *transport.Reader) {
var req transport.StatusReq
if err := r.ReadJSON(&req); err != nil {
return
}
resp := transport.StatusResp{
Type: transport.FrameTypeStatusResp,
Bus: transport.StatusBus{
PID: os.Getpid(),
UptimeSecs: int64(time.Since(d.started).Seconds()),
IdleTimeoutSec: int(d.cfg.IdleTimeout / time.Second),
ClientID: d.cfg.ClientID,
Edition: d.cfg.Edition,
SourceKind: d.cfg.SourceKind,
IdentityHash: d.cfg.IdentityHash,
SourceID: d.cfg.SourceID,
},
SourceState: transport.StatusSource{
State: "connected", // v1: source state plumbed in P3+
Source: "inferred",
},
Consumers: d.hub.Snapshot(),
PerEventTypeCounters: d.hub.Counters().Snapshot(),
}
_ = w.WriteJSON(resp)
}
// idleWatch fires d.idleStop when IdleTimeout passes with zero registered
// consumers. Disabled when IdleTimeout <= 0 (returns immediately; idleStop
// is then never closed and Run's select branch on it is effectively dead).
//
// Pre-condition: d.idleStop has already been allocated by Run so the parent
// select never races on a nil channel (which would block forever).
func (d *daemon) idleWatch(ctx context.Context) {
if d.cfg.IdleTimeout <= 0 {
return
}
tick := time.NewTicker(d.cfg.IdleTimeout / 4)
defer tick.Stop()
emptySince := time.Time{}
for {
select {
case <-ctx.Done():
return
case <-tick.C:
if d.hub.Len() == 0 {
if emptySince.IsZero() {
emptySince = time.Now()
} else if time.Since(emptySince) >= d.cfg.IdleTimeout {
close(d.idleStop)
return
}
} else {
emptySince = time.Time{}
}
}
}
}
// triggerShutdown is called from RPC handlers that want to end the bus.
// It works by closing the listener (which unblocks Run's select via the
// source error path, indirectly). For v1 a full ctx-cancellation hook is
// out of scope; busctl/stop also sends SIGTERM which is the authoritative
// shutdown path.
func (d *daemon) triggerShutdown(reason string) {
d.log.Info("bus: shutdown triggered via IPC", "reason", reason)
_ = d.listener.Close() // unblocks Accept(), but doesn't kill Source
// Best-effort: a future version wires a context.CancelFunc here.
}
// shutdown performs the graceful tear-down sequence:
// 1. mark shuttingDown so acceptLoop exits cleanly
// 2. broadcast Bye to all consumers
// 3. close listener (interrupts pending Accept)
// 4. wait for all per-connection goroutines to drain
// 5. lock + meta cleanup via Run's defers
func (d *daemon) shutdown() {
d.shutdownMu.Lock()
defer d.shutdownMu.Unlock()
if !d.shuttingDown.CompareAndSwap(false, true) {
return
}
d.hub.Broadcast(transport.Bye{Type: transport.FrameTypeBye, Reason: "shutdown"})
_ = d.listener.Close()
// Force-close all open IPC connections so any reader goroutine blocked
// on Read() returns with a network error and exits cleanly. Without
// this the consumerWG never drains and Run hangs forever.
d.conns.Range(func(k, _ any) bool {
if c, ok := k.(net.Conn); ok {
_ = c.Close()
}
return true
})
// Give consumers a brief moment to drain final frames before we tear
// down their channels.
doneCh := make(chan struct{})
go func() {
d.consumerWG.Wait()
close(doneCh)
}()
select {
case <-doneCh:
case <-time.After(2 * time.Second):
d.log.Warn("bus: shutdown: consumer goroutines did not drain within 2s")
}
}
// signalReady writes a single 'R' byte to the ready pipe (if provided) and
// closes it. The parent process (busctl/spawn) reads one byte and proceeds.
func signalReady(p *os.File) {
if p == nil {
return
}
_, _ = p.Write([]byte{'R'})
_ = p.Close()
}
// failReady writes 'E' to the ready pipe (if provided) and returns err.
// Used by the startup-failure paths so the parent can distinguish "still
// starting up" from "failed to start".
func failReady(p *os.File, err error) error {
if p != nil {
_, _ = p.Write([]byte{'E'})
_ = p.Close()
}
return err
}
+358
View File
@@ -0,0 +1,358 @@
// 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 bus
import (
"context"
"errors"
"io"
"os"
"path/filepath"
"runtime"
"testing"
"time"
dwsevent "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/transport"
)
// fakeSource is a minimal SourceAdapter that emits a configured slice of
// events then blocks on ctx until cancel. Used by daemon integration tests
// in place of the real Stream SDK.
//
// If trigger is non-nil, fakeSource waits for it to close before emitting.
// Tests use this to register a consumer first (avoids the race where events
// flow before Hello completes and end up dropped by Hub for lack of a
// matching consumer).
type fakeSource struct {
events []dwsevent.RawEvent
delay time.Duration // optional delay between emits to let the consumer drain
trigger <-chan struct{} // optional gate; nil = emit immediately
}
func (f *fakeSource) Start(ctx context.Context, emit dwsevent.EmitFn) error {
if f.trigger != nil {
select {
case <-f.trigger:
case <-ctx.Done():
return ctx.Err()
}
}
for i := range f.events {
select {
case <-ctx.Done():
return ctx.Err()
default:
}
ev := f.events[i]
ev.ReceivedAt = time.Now().UTC()
emit(&ev)
if f.delay > 0 {
time.Sleep(f.delay)
}
}
<-ctx.Done()
return ctx.Err()
}
func skipOnWindows(t *testing.T, reason string) {
t.Helper()
if runtime.GOOS == "windows" {
t.Skipf("skip on windows: %s", reason)
}
}
// shortTempDir returns a temp dir under /tmp so the resulting unix socket
// path stays under the macOS 104-byte sun_path limit. t.TempDir() lives in
// $TMPDIR (/var/folders/.../T/...) which can easily exceed that.
func shortTempDir(t *testing.T) string {
t.Helper()
dir, err := os.MkdirTemp("/tmp", "dws-bus-")
if err != nil {
t.Fatalf("mktemp: %v", err)
}
t.Cleanup(func() { _ = os.RemoveAll(dir) })
return dir
}
// waitForFile polls until path exists or timeout elapses.
func waitForFile(t *testing.T, path string, timeout time.Duration) {
t.Helper()
deadline := time.Now().Add(timeout)
for time.Now().Before(deadline) {
if _, err := os.Stat(path); err == nil {
return
}
time.Sleep(10 * time.Millisecond)
}
t.Fatalf("file %q did not appear within %s", path, timeout)
}
func TestDaemon_RunStartsAndShutsDownCleanly(t *testing.T) {
skipOnWindows(t, "Unix socket path semantics differ; covered by transport_windows_test")
workDir := shortTempDir(t)
sockPath := filepath.Join(workDir, "bus.sock")
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
src := &fakeSource{}
runDone := make(chan error, 1)
go func() {
runDone <- Run(ctx, Config{
WorkDir: workDir,
IPCEndpoint: sockPath,
ClientID: "ding_test",
Edition: "open",
Source: src,
})
}()
waitForFile(t, sockPath, 2*time.Second)
if pid := ReadHolderPID(filepath.Join(workDir, LockFileName)); pid != os.Getpid() {
t.Errorf("bus.lock pid = %d, want %d", pid, os.Getpid())
}
if _, err := ReadMeta(workDir); err != nil {
t.Errorf("bus.meta missing: %v", err)
}
cancel()
select {
case err := <-runDone:
if err != nil && !errors.Is(err, context.Canceled) {
t.Fatalf("Run returned %v, want nil or canceled", err)
}
case <-time.After(3 * time.Second):
t.Fatal("Run did not return after ctx cancel")
}
}
func TestDaemon_ConsumerReceivesEvents(t *testing.T) {
skipOnWindows(t, "uses Unix socket dial")
workDir := shortTempDir(t)
sockPath := filepath.Join(workDir, "bus.sock")
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
trigger := make(chan struct{})
src := &fakeSource{
events: []dwsevent.RawEvent{
{EventID: "1", EventType: "im.message.receive_v1", Data: `{"text":"hi"}`},
{EventID: "2", EventType: "approval.task", Data: `{"task":"x"}`},
{EventID: "3", EventType: "im.message.at_v1", Data: `{"at":1}`},
},
delay: 5 * time.Millisecond,
trigger: trigger,
}
runDone := make(chan error, 1)
go func() {
runDone <- Run(ctx, Config{
WorkDir: workDir,
IPCEndpoint: sockPath,
ClientID: "ding_test",
Edition: "open",
Source: src,
})
}()
waitForFile(t, sockPath, 2*time.Second)
conn, err := transport.Dial(sockPath)
if err != nil {
t.Fatalf("Dial: %v", err)
}
defer conn.Close()
w := transport.NewWriter(conn)
r := transport.NewReader(conn)
if err := w.WriteJSON(transport.Hello{
Type: transport.FrameTypeHello,
ConsumerPID: os.Getpid(),
EventTypes: []string{"im.*"},
}); err != nil {
t.Fatalf("write hello: %v", err)
}
var ack transport.HelloAck
if err := r.ReadJSON(&ack); err != nil {
t.Fatalf("read ack: %v", err)
}
if ack.Type != transport.FrameTypeHelloAck {
t.Fatalf("ack type = %s", ack.Type)
}
// Consumer is now registered; trigger fakeSource to emit.
close(trigger)
received := 0
deadline := time.After(3 * time.Second)
for received < 2 {
select {
case <-deadline:
t.Fatalf("only received %d events, want 2", received)
default:
}
raw, err := r.Read()
if err != nil {
t.Fatalf("read frame: %v", err)
}
typ, _ := transport.PeekType(raw)
if typ == transport.FrameTypeEvent {
received++
}
}
cancel()
<-runDone
}
func TestDaemon_LockBusyOnSecondRun(t *testing.T) {
skipOnWindows(t, "uses Unix socket / flock semantics")
workDir := shortTempDir(t)
sockPath := filepath.Join(workDir, "bus.sock")
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
go func() {
_ = Run(ctx, Config{
WorkDir: workDir,
IPCEndpoint: sockPath,
ClientID: "ding_test",
Edition: "open",
Source: &fakeSource{},
})
}()
waitForFile(t, sockPath, 2*time.Second)
err := Run(context.Background(), Config{
WorkDir: workDir,
IPCEndpoint: sockPath + ".other",
ClientID: "ding_test",
Edition: "open",
Source: &fakeSource{},
})
if !errors.Is(err, ErrBusy) {
t.Fatalf("second Run = %v, want ErrBusy", err)
}
}
func TestDaemon_IdleTimeoutSelfStops(t *testing.T) {
skipOnWindows(t, "uses Unix socket")
workDir := shortTempDir(t)
sockPath := filepath.Join(workDir, "bus.sock")
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
runDone := make(chan error, 1)
go func() {
runDone <- Run(ctx, Config{
WorkDir: workDir,
IPCEndpoint: sockPath,
ClientID: "ding_test",
Edition: "open",
Source: &fakeSource{},
IdleTimeout: 200 * time.Millisecond,
})
}()
waitForFile(t, sockPath, 2*time.Second)
select {
case <-runDone:
// success
case <-time.After(3 * time.Second):
t.Fatal("bus did not idle-stop within deadline")
}
}
func TestDaemon_ReadyPipeSignalsR(t *testing.T) {
skipOnWindows(t, "uses Unix socket + pipe")
workDir := shortTempDir(t)
sockPath := filepath.Join(workDir, "bus.sock")
pr, pw, err := os.Pipe()
if err != nil {
t.Fatal(err)
}
defer pr.Close()
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
go Run(ctx, Config{
WorkDir: workDir,
IPCEndpoint: sockPath,
ClientID: "ding_test",
Edition: "open",
Source: &fakeSource{},
ReadyPipe: pw,
})
buf := make([]byte, 1)
pr.SetReadDeadline(time.Now().Add(2 * time.Second))
n, err := pr.Read(buf)
if err != nil && !errors.Is(err, io.EOF) {
t.Fatalf("read ready pipe: %v", err)
}
if n != 1 || buf[0] != 'R' {
t.Fatalf("ready byte = %q n=%d, want 'R'", buf, n)
}
cancel()
}
func TestDaemon_ConsumerEOFAutoUnregisters(t *testing.T) {
skipOnWindows(t, "uses Unix socket")
workDir := shortTempDir(t)
sockPath := filepath.Join(workDir, "bus.sock")
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
go Run(ctx, Config{
WorkDir: workDir,
IPCEndpoint: sockPath,
ClientID: "ding_test",
Edition: "open",
Source: &fakeSource{},
})
waitForFile(t, sockPath, 2*time.Second)
conn, err := transport.Dial(sockPath)
if err != nil {
t.Fatal(err)
}
w := transport.NewWriter(conn)
r := transport.NewReader(conn)
_ = w.WriteJSON(transport.Hello{Type: transport.FrameTypeHello, ConsumerPID: 12345})
var ack transport.HelloAck
_ = r.ReadJSON(&ack)
// Slam the connection shut without sending Bye.
conn.Close()
time.Sleep(200 * time.Millisecond)
c2, err := transport.Dial(sockPath)
if err != nil {
t.Fatal(err)
}
defer c2.Close()
w2 := transport.NewWriter(c2)
r2 := transport.NewReader(c2)
_ = w2.WriteJSON(transport.Hello{Type: transport.FrameTypeHello, Role: transport.HelloRoleStatus})
_ = w2.WriteJSON(transport.StatusReq{Type: transport.FrameTypeStatusReq})
var resp transport.StatusResp
if err := r2.ReadJSON(&resp); err != nil {
t.Fatalf("read status: %v", err)
}
for _, c := range resp.Consumers {
if c.PID == 12345 {
t.Fatalf("dead consumer 12345 still present in status: %+v", resp.Consumers)
}
}
}
+34
View File
@@ -0,0 +1,34 @@
// 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 bus implements the daemon side of the dws event subsystem: one
// long-lived process per ClientID that holds the single cloud connection
// and fans out events to N local consumers over IPC.
//
// Files (mirroring plan §5 layout):
//
// daemon.go main loop: lock → meta → IPC listen → ready → source.Start
// hub.go consumer registry, per-consumer sendCh, drop-oldest backpressure
// metrics.go per-event-type + per-consumer received/dropped counters
// lockfile.go single bus.lock (flock + PID content + stale recovery)
// meta.go bus.meta JSON (clientID/edition/started_at) for list/status reverse mapping
//
// Lifecycle invariants (plan §4 invariants 1–7):
// 1. emit non-blocking (drop-oldest, never block SDK callback)
// 2. dedup on event_id (LRU) to absorb cloud-side redelivery
// 3. single bus per ClientID (bus.lock enforces, all FS paths use clientIDHash)
// 4. upstream always full subscription; consumer filter only affects bus→consume
// 5. dead-consumer auto-reap on socket EOF
// 6. startup order: lock → meta → IPC listen → ready pipe → Source.Start
// 7. stdio detach when fork'd by busctl/spawn
package bus
+81
View File
@@ -0,0 +1,81 @@
// 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 bus
import (
"context"
"log/slog"
"time"
)
// DefaultDropWarnPercent is the threshold above which per-event-type drop
// rate triggers a slog WARN line in bus.log. 5% is the plan default
// (§15 已决项); overridable via DWS_EVENT_DROP_WARN_PCT.
//
// We use whole-percentage granularity (int) because the counter math is
// integer; sub-percent precision would just add noise.
const DefaultDropWarnPercent = 5
// dropWarnTickInterval is how often the watcher samples counters. 30s
// balances responsiveness ("see the warning while the burst is still
// happening") with log noise (one warning per scan, not one per drop).
const dropWarnTickInterval = 30 * time.Second
// dropWarnState memoises the last warned drop rate per event type so we
// only emit a fresh WARN when the situation worsens by at least
// dropWarnHysteresis percentage points. Without hysteresis a steady-state
// burst would re-warn every tick.
const dropWarnHysteresis = 5
// dropWarnWatcher periodically samples per-event-type counters; for any
// type whose drop rate crosses the threshold (and hasn't recently been
// warned at the same level), it emits a slog WARN. Runs as a daemon
// goroutine spawned from bus.Run.
//
// Lifecycle: returns when ctx is cancelled (bus shutdown). Never holds
// any external lock — uses the counters' own concurrency-safe Snapshot.
func dropWarnWatcher(ctx context.Context, counters *PerTypeCounters, log *slog.Logger, threshold int) {
if threshold <= 0 || threshold > 100 {
threshold = DefaultDropWarnPercent
}
tick := time.NewTicker(dropWarnTickInterval)
defer tick.Stop()
lastWarnedPct := make(map[string]int)
for {
select {
case <-ctx.Done():
return
case <-tick.C:
for _, et := range counters.SortedTypes() {
pct := counters.DropRatePercent(et)
if pct < threshold {
// drop rate is healthy → forget any prior warning so
// a future spike re-triggers a fresh WARN
delete(lastWarnedPct, et)
continue
}
prev, warned := lastWarnedPct[et]
if warned && pct < prev+dropWarnHysteresis {
continue // not significantly worse; suppress
}
log.Warn("bus: event type backpressure",
"event_type", et,
"drop_pct", pct,
"threshold_pct", threshold,
)
lastWarnedPct[et] = pct
}
}
}
}
+51
View File
@@ -0,0 +1,51 @@
// 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 bus
import (
"testing"
)
// dropWarnWatcher is exercised end-to-end by daemon integration tests
// (it runs as a goroutine inside bus.Run). At the unit level we just
// verify the threshold-clamping behaviour in isolation so a misconfigured
// env var doesn't disable the safety net silently.
func TestDropWarnWatcher_ThresholdClamp(t *testing.T) {
// The function applies its own clamp before using the threshold;
// we re-derive the expected effective value through the same path
// (calling helper-style code below).
for _, in := range []int{-1, 0, 101, 1000} {
got := clampDropWarnPctForTest(in)
if got != DefaultDropWarnPercent {
t.Errorf("invalid threshold %d should clamp to %d, got %d",
in, DefaultDropWarnPercent, got)
}
}
for _, in := range []int{1, 5, 10, 50, 100} {
if got := clampDropWarnPctForTest(in); got != in {
t.Errorf("valid threshold %d should pass through, got %d", in, got)
}
}
}
// clampDropWarnPctForTest mirrors the clamp logic in dropWarnWatcher.
// Kept separate so the public Watcher signature does not have to expose
// internal validation as a method — the watcher reads threshold from
// closure, and tests want to assert on the validation contract.
func clampDropWarnPctForTest(threshold int) int {
if threshold <= 0 || threshold > 100 {
return DefaultDropWarnPercent
}
return threshold
}
+399
View File
@@ -0,0 +1,399 @@
// 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 bus
import (
"regexp"
"sort"
"strings"
"sync"
"sync/atomic"
"time"
dwsevent "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/transport"
)
// DefaultSendBuffer is the per-consumer channel capacity used when the Hub
// is constructed via NewHub. Sized to absorb a short event burst (~100ms at
// ~1k evt/s) without backpressure. Overridable via DWS_EVENT_CONSUMER_BUFFER
// at daemon start (plan §15 已决但暴露方式).
const DefaultSendBuffer = 100
// Consumer represents one registered IPC connection to the bus. Wire-side
// reader/writer goroutines are owned by daemon.go; the Hub holds the
// metadata + sendCh.
type Consumer struct {
ID int // monotonic, assigned by Hub
PID int // from Hello.ConsumerPID
EventTypes []string // raw wildcard patterns from Hello
Filter string // raw regex from Hello (for status display)
SubscribeID string // optional personal subscription label and local isolation key
SubscribedAt time.Time
SendCh chan any // bus → consume frames (Event/SourceState/Heartbeat/Bye)
matcher consumerMatcher
sendMu sync.Mutex // serialises Deliver/Broadcast with SendCh close
closed bool // guarded by sendMu
seq atomic.Uint64 // monotonic per-consumer sequence, starts at 1
received atomic.Uint64
dropped atomic.Uint64
}
// consumerMatcher pre-compiles EventTypes wildcard patterns and the optional
// Filter regex into a fast checker invoked once per delivered event per
// consumer. Empty EventTypes means catch-all (everything matches except
// what Filter excludes). Non-empty SubscribeID is an additional exact-match
// constraint used by personal_stream consumers so same event_type subscriptions
// do not fan out to each other.
type consumerMatcher struct {
catchAll bool
exact map[string]struct{} // patterns without '*'
prefixes []string // patterns ending in ".*" or "*" — store the prefix only
filter *regexp.Regexp // nil if no filter
subscribeID string // empty = no subscribe_id filtering
}
func compileMatcher(eventTypes []string, filter string, subscribeID string) (consumerMatcher, error) {
m := consumerMatcher{
exact: make(map[string]struct{}),
subscribeID: strings.TrimSpace(subscribeID),
}
if len(eventTypes) == 0 {
m.catchAll = true
}
for _, raw := range eventTypes {
raw = strings.TrimSpace(raw)
if raw == "" {
continue
}
// Treat trailing '*' or '.*' as a prefix wildcard (e.g. "im.*",
// "im.message.*"). Middle/leading wildcards are rare for event_type
// strings — if a real use case appears we can switch to regex
// compilation here without breaking the wire format.
switch {
case raw == "*":
m.catchAll = true
case strings.HasSuffix(raw, ".*"):
m.prefixes = append(m.prefixes, raw[:len(raw)-1]) // keep the dot, drop the *
case strings.HasSuffix(raw, "*"):
m.prefixes = append(m.prefixes, raw[:len(raw)-1])
default:
m.exact[raw] = struct{}{}
}
}
if filter != "" {
re, err := regexp.Compile(filter)
if err != nil {
return consumerMatcher{}, err
}
m.filter = re
}
return m, nil
}
func (m *consumerMatcher) matches(raw *dwsevent.RawEvent) bool {
if raw == nil {
return false
}
if m.subscribeID != "" && raw.SubscribeID != m.subscribeID {
return false
}
eventType := raw.EventType
// First the include rules: catchAll OR exact-list OR prefix-list.
included := m.catchAll
if !included {
if _, ok := m.exact[eventType]; ok {
included = true
}
}
if !included {
for _, p := range m.prefixes {
if strings.HasPrefix(eventType, p) {
included = true
break
}
}
}
if !included {
return false
}
// Filter is an additional AND constraint (regex on the event_type
// string). Empty filter is no-op.
if m.filter != nil && !m.filter.MatchString(eventType) {
return false
}
return true
}
// Hub is the bus's fan-out engine. It owns the set of registered consumers
// and the bus-wide per-event-type counters. The Hub is concurrency-safe
// across Register/Unregister/Deliver/Snapshot; Deliver is the hot path and
// is RLock-only.
type Hub struct {
mu sync.RWMutex
consumers map[int]*Consumer
nextID int
bufferSize int
counters *PerTypeCounters
}
// NewHub returns a Hub with the given per-consumer channel buffer size.
// Zero or negative uses DefaultSendBuffer.
func NewHub(bufferSize int) *Hub {
if bufferSize <= 0 {
bufferSize = DefaultSendBuffer
}
return &Hub{
consumers: make(map[int]*Consumer),
bufferSize: bufferSize,
counters: NewPerTypeCounters(),
}
}
// Counters exposes the bus-wide per-event-type counter set for daemon-side
// rendering (status RPC, drop-rate warning).
func (h *Hub) Counters() *PerTypeCounters { return h.counters }
// RegisterError wraps the matcher compile error so the daemon can refuse
// the Hello and return a clean error to the consume client (instead of
// silently accepting a bad filter regex).
type RegisterError struct{ Err error }
func (e *RegisterError) Error() string { return "bus: register consumer: " + e.Err.Error() }
func (e *RegisterError) Unwrap() error { return e.Err }
// Register adds a consumer derived from a Hello frame. Returns a new
// Consumer with the populated ID + sendCh ready to use, or a RegisterError
// if the Hello's Filter regex is invalid.
func (h *Hub) Register(hello transport.Hello) (*Consumer, error) {
m, err := compileMatcher(hello.EventTypes, hello.Filter, hello.SubscribeID)
if err != nil {
return nil, &RegisterError{Err: err}
}
h.mu.Lock()
defer h.mu.Unlock()
h.nextID++
c := &Consumer{
ID: h.nextID,
PID: hello.ConsumerPID,
EventTypes: append([]string(nil), hello.EventTypes...),
Filter: hello.Filter,
SubscribeID: strings.TrimSpace(hello.SubscribeID),
SubscribedAt: time.Now().UTC(),
SendCh: make(chan any, h.bufferSize),
matcher: m,
}
h.consumers[c.ID] = c
return c, nil
}
// Unregister removes a consumer by ID and closes its sendCh. Idempotent —
// calling twice or on an unknown ID is a no-op. closeSend shares the same
// per-consumer lock as Deliver/Broadcast, so a stale Hub snapshot cannot send
// to the channel after it has been closed.
func (h *Hub) Unregister(id int) {
h.mu.Lock()
c, ok := h.consumers[id]
if !ok {
h.mu.Unlock()
return
}
delete(h.consumers, id)
h.mu.Unlock()
c.closeSend()
}
// Deliver fans the raw event out to every matching consumer. Updates
// bus-wide and per-consumer counters. Always non-blocking — drops the
// oldest entry in any full sendCh (plan invariant #1).
//
// Called from the bus daemon's main loop after dedup; safe for concurrent
// callers (Hub uses an RLock, Consumer drop-oldest is single-producer-safe
// because the daemon serialises Deliver per event).
func (h *Hub) Deliver(raw *dwsevent.RawEvent) {
if raw == nil {
return
}
h.counters.AddReceived(raw.EventType)
h.mu.RLock()
matched := make([]*Consumer, 0, len(h.consumers))
for _, c := range h.consumers {
if c.matcher.matches(raw) {
matched = append(matched, c)
}
}
h.mu.RUnlock()
for _, c := range matched {
c.deliver(raw, h.counters)
}
}
// deliver builds the per-consumer Event frame (assigning seq) and pushes
// it onto sendCh with drop-oldest semantics. Updates per-consumer and
// bus-wide drop counters.
func (c *Consumer) deliver(raw *dwsevent.RawEvent, hubCounters *PerTypeCounters) {
c.sendMu.Lock()
defer c.sendMu.Unlock()
if c.closed {
return
}
seq := c.seq.Add(1)
frame := transport.Event{
Type: transport.FrameTypeEvent,
Seq: seq,
EventID: raw.EventID,
EventBornTime: raw.EventBornTime,
EventCorpID: raw.EventCorpID,
EventType: raw.EventType,
EventUnifiedAppID: raw.EventUnifiedAppID,
EventScope: raw.EventScope,
SubscribeID: raw.SubscribeID,
SourceID: raw.SourceID,
RuleType: raw.RuleType,
Data: raw.Data,
Headers: raw.Headers,
ReceivedAtUnixMS: raw.ReceivedAt.UnixMilli(),
}
success, evicted := c.tryPushOrDropOldestLocked(frame)
if evicted {
// An older event we had previously enqueued is gone.
c.dropped.Add(1)
c.received.Add(^uint64(0)) // -1
hubCounters.AddDropped(raw.EventType)
}
if success {
c.received.Add(1)
} else {
// New event also didn't make it (rare: lost the race after eviction)
c.dropped.Add(1)
if !evicted {
// No eviction happened but push still failed → unexpected; only
// reached via concurrent reader-then-something. Count the bus-wide
// drop too so the metric matches per-consumer drops.
hubCounters.AddDropped(raw.EventType)
}
}
}
// tryPushOrDropOldestLocked tries to push frame to SendCh non-blockingly. If
// full, it pops the oldest entry to make room and tries once more. Caller must
// hold c.sendMu, which makes drop-oldest a true single-producer operation even
// when Deliver and Broadcast run concurrently.
//
// Returns:
//
// success: true if the new frame is now in the channel
// evicted: true if we removed an older queued frame to make room
//
// "received" semantics in deliver() interpret these:
//
// success=true, evicted=false → +1 received (normal push)
// success=true, evicted=true → net 0 (lost 1, gained 1); +1 dropped
// success=false, evicted=true → -1 received, +2 dropped (both lost; rare)
// success=false, evicted=false → only possible if reader drained between
// try and we still missed (extremely rare); +1 dropped only.
//
// The per-consumer send lock makes this block the sole SendCh producer, so the
// second push cannot race another producer. The transport writer may receive
// between our two operations, which only makes more room — never less.
func (c *Consumer) tryPushOrDropOldestLocked(frame any) (success bool, evicted bool) {
select {
case c.SendCh <- frame:
return true, false
default:
}
// Drop oldest, then retry.
select {
case <-c.SendCh:
evicted = true
default:
// Reader took it, room is back.
}
select {
case c.SendCh <- frame:
return true, evicted
default:
return false, evicted
}
}
func (c *Consumer) enqueue(frame any) bool {
c.sendMu.Lock()
defer c.sendMu.Unlock()
if c.closed {
return false
}
success, _ := c.tryPushOrDropOldestLocked(frame)
return success
}
func (c *Consumer) closeSend() {
c.sendMu.Lock()
defer c.sendMu.Unlock()
if c.closed {
return
}
c.closed = true
close(c.SendCh)
}
// Broadcast sends the same frame (e.g. SourceState change, Bye on shutdown)
// to every consumer using drop-oldest semantics. Returns the number of
// consumers the frame was successfully enqueued for.
func (h *Hub) Broadcast(frame any) int {
h.mu.RLock()
cs := make([]*Consumer, 0, len(h.consumers))
for _, c := range h.consumers {
cs = append(cs, c)
}
h.mu.RUnlock()
ok := 0
for _, c := range cs {
if c.enqueue(frame) {
ok++
}
}
return ok
}
// Snapshot returns a deterministic StatusConsumer slice (sorted by PID)
// for status RPC encoding. Caller MUST NOT mutate the returned slice.
func (h *Hub) Snapshot() []transport.StatusConsumer {
h.mu.RLock()
defer h.mu.RUnlock()
out := make([]transport.StatusConsumer, 0, len(h.consumers))
for _, c := range h.consumers {
out = append(out, transport.StatusConsumer{
PID: c.PID,
EventTypes: append([]string(nil), c.EventTypes...),
Filter: c.Filter,
SubscribeID: c.SubscribeID,
SubscribedAtMS: c.SubscribedAt.UnixMilli(),
Received: c.received.Load(),
Dropped: c.dropped.Load(),
})
}
sort.Slice(out, func(i, j int) bool { return out[i].PID < out[j].PID })
return out
}
// Len returns the current number of registered consumers.
func (h *Hub) Len() int {
h.mu.RLock()
defer h.mu.RUnlock()
return len(h.consumers)
}
+411
View File
@@ -0,0 +1,411 @@
// 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 bus
import (
"errors"
"sync"
"testing"
"time"
dwsevent "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/transport"
)
func mkEvent(typ, id string) *dwsevent.RawEvent {
return &dwsevent.RawEvent{
EventID: id,
EventType: typ,
Data: `{}`,
ReceivedAt: time.Now().UTC(),
}
}
func drain(c *Consumer, n int, t *testing.T) []*transport.Event {
t.Helper()
out := make([]*transport.Event, 0, n)
for i := 0; i < n; i++ {
select {
case f := <-c.SendCh:
ev, ok := f.(transport.Event)
if !ok {
t.Fatalf("frame %d is not Event: %T", i, f)
}
out = append(out, &ev)
case <-time.After(time.Second):
t.Fatalf("timed out waiting for event %d", i)
}
}
return out
}
func assertNoEvent(c *Consumer, t *testing.T) {
t.Helper()
select {
case f := <-c.SendCh:
t.Fatalf("unexpected frame: %#v", f)
case <-time.After(50 * time.Millisecond):
}
}
func TestHub_RegisterAssignsMonotonicID(t *testing.T) {
h := NewHub(10)
c1, err := h.Register(transport.Hello{ConsumerPID: 1})
if err != nil {
t.Fatal(err)
}
c2, err := h.Register(transport.Hello{ConsumerPID: 2})
if err != nil {
t.Fatal(err)
}
if c1.ID >= c2.ID {
t.Fatalf("IDs not monotonic: %d, %d", c1.ID, c2.ID)
}
}
func TestHub_DeliverMatchesPrefix(t *testing.T) {
h := NewHub(10)
c, err := h.Register(transport.Hello{EventTypes: []string{"im.message.*"}})
if err != nil {
t.Fatal(err)
}
h.Deliver(mkEvent("im.message.receive_v1", "1"))
h.Deliver(mkEvent("approval.task", "2")) // no match
h.Deliver(mkEvent("im.message.at_v1", "3"))
got := drain(c, 2, t)
if got[0].EventID != "1" || got[1].EventID != "3" {
t.Fatalf("expected 1,3 got %s,%s", got[0].EventID, got[1].EventID)
}
if got[0].Seq != 1 || got[1].Seq != 2 {
t.Fatalf("seq mismatch: %d %d", got[0].Seq, got[1].Seq)
}
}
func TestHub_DeliverCatchAll(t *testing.T) {
h := NewHub(10)
c, _ := h.Register(transport.Hello{}) // empty == catch-all
h.Deliver(mkEvent("im.message.receive_v1", "1"))
h.Deliver(mkEvent("approval.task", "2"))
h.Deliver(mkEvent("foo.bar", "3"))
got := drain(c, 3, t)
if got[0].EventID != "1" || got[1].EventID != "2" || got[2].EventID != "3" {
t.Fatalf("catch-all missed events: %+v", got)
}
}
func TestHub_DeliverFilterRegex(t *testing.T) {
h := NewHub(10)
c, err := h.Register(transport.Hello{Filter: `^im\.`})
if err != nil {
t.Fatal(err)
}
h.Deliver(mkEvent("im.message.receive_v1", "1"))
h.Deliver(mkEvent("approval.task", "2")) // filtered out
h.Deliver(mkEvent("im.chat.member.bot.added_v1", "3"))
got := drain(c, 2, t)
if got[0].EventID != "1" || got[1].EventID != "3" {
t.Fatalf("filter regex missed: %+v", got)
}
}
func TestHub_DeliverFiltersBySubscribeID(t *testing.T) {
h := NewHub(10)
c, err := h.Register(transport.Hello{
EventTypes: []string{"user_im_message_receive_o2o"},
SubscribeID: "sub-b",
})
if err != nil {
t.Fatal(err)
}
wrong := mkEvent("user_im_message_receive_o2o", "1")
wrong.SubscribeID = "sub-c"
h.Deliver(wrong)
assertNoEvent(c, t)
right := mkEvent("user_im_message_receive_o2o", "2")
right.SubscribeID = "sub-b"
h.Deliver(right)
got := drain(c, 1, t)
if got[0].EventID != "2" || got[0].SubscribeID != "sub-b" {
t.Fatalf("event = %#v, want sub-b event 2", got[0])
}
}
func TestHub_DeliverSubscribeIDSeparatesSameEventType(t *testing.T) {
h := NewHub(10)
b, err := h.Register(transport.Hello{
EventTypes: []string{"user_im_message_receive_o2o"},
SubscribeID: "sub-b",
})
if err != nil {
t.Fatal(err)
}
c, err := h.Register(transport.Hello{
EventTypes: []string{"user_im_message_receive_o2o"},
SubscribeID: "sub-c",
})
if err != nil {
t.Fatal(err)
}
evB := mkEvent("user_im_message_receive_o2o", "b-msg")
evB.SubscribeID = "sub-b"
evC := mkEvent("user_im_message_receive_o2o", "c-msg")
evC.SubscribeID = "sub-c"
h.Deliver(evB)
h.Deliver(evC)
gotB := drain(b, 1, t)
gotC := drain(c, 1, t)
if gotB[0].EventID != "b-msg" || gotB[0].SubscribeID != "sub-b" {
t.Fatalf("B consumer got %#v", gotB[0])
}
if gotC[0].EventID != "c-msg" || gotC[0].SubscribeID != "sub-c" {
t.Fatalf("C consumer got %#v", gotC[0])
}
assertNoEvent(b, t)
assertNoEvent(c, t)
}
func TestHub_DeliverDropsMissingSubscribeIDForSpecificConsumer(t *testing.T) {
h := NewHub(10)
c, err := h.Register(transport.Hello{
EventTypes: []string{"user_im_message_receive_group"},
SubscribeID: "sub-group",
})
if err != nil {
t.Fatal(err)
}
h.Deliver(mkEvent("user_im_message_receive_group", "missing-sub"))
assertNoEvent(c, t)
}
func TestHub_DeliverEmptyConsumerSubscribeIDReceivesAnySubscribeID(t *testing.T) {
h := NewHub(10)
c, err := h.Register(transport.Hello{
EventTypes: []string{"user_im_message_receive_o2o"},
})
if err != nil {
t.Fatal(err)
}
ev := mkEvent("user_im_message_receive_o2o", "1")
ev.SubscribeID = "sub-any"
h.Deliver(ev)
got := drain(c, 1, t)
if got[0].SubscribeID != "sub-any" {
t.Fatalf("event subscribe_id = %q, want sub-any", got[0].SubscribeID)
}
}
func TestHub_DeliverEventTypeMismatchEvenWithSubscribeID(t *testing.T) {
h := NewHub(10)
c, err := h.Register(transport.Hello{
EventTypes: []string{"user_im_message_receive_o2o"},
SubscribeID: "sub-1",
})
if err != nil {
t.Fatal(err)
}
ev := mkEvent("user_im_message_receive_group", "1")
ev.SubscribeID = "sub-1"
h.Deliver(ev)
assertNoEvent(c, t)
}
func TestHub_RegisterRejectsBadFilterRegex(t *testing.T) {
h := NewHub(10)
_, err := h.Register(transport.Hello{Filter: `(unclosed`})
if err == nil {
t.Fatal("expected RegisterError for bad regex")
}
var re *RegisterError
if !errors.As(err, &re) {
t.Fatalf("err = %v, want *RegisterError", err)
}
}
func TestHub_DropOldestOnFullChannel(t *testing.T) {
h := NewHub(2) // small buffer
c, _ := h.Register(transport.Hello{})
// Push 5 events without draining → 2 stay, 3 dropped
for i := 0; i < 5; i++ {
h.Deliver(mkEvent("foo", string(rune('0'+i))))
}
if c.received.Load() != 2 {
t.Fatalf("received = %d, want 2", c.received.Load())
}
if c.dropped.Load() != 3 {
t.Fatalf("dropped = %d, want 3", c.dropped.Load())
}
// Bus counters reflect it too
snap := h.Counters().Snapshot()
if snap["foo"].Dropped != 3 {
t.Fatalf("hub dropped = %d, want 3", snap["foo"].Dropped)
}
}
func TestHub_UnregisterClosesChannel(t *testing.T) {
h := NewHub(10)
c, _ := h.Register(transport.Hello{})
h.Unregister(c.ID)
// Channel should be closed; receive returns zero value with ok=false
_, ok := <-c.SendCh
if ok {
t.Fatal("SendCh should be closed after Unregister")
}
// Further Deliver must not panic (closed flag prevents send)
h.Deliver(mkEvent("foo", "x"))
if h.Len() != 0 {
t.Fatalf("Len after Unregister = %d, want 0", h.Len())
}
}
func TestHub_UnregisterIdempotent(t *testing.T) {
h := NewHub(10)
c, _ := h.Register(transport.Hello{})
h.Unregister(c.ID)
h.Unregister(c.ID) // must not panic
h.Unregister(9999) // unknown ID
}
func TestHub_ConcurrentDeliverBroadcastUnregister(t *testing.T) {
for iteration := 0; iteration < 200; iteration++ {
h := NewHub(4)
c, err := h.Register(transport.Hello{ConsumerPID: iteration + 1})
if err != nil {
t.Fatal(err)
}
start := make(chan struct{})
var wg sync.WaitGroup
for producer := 0; producer < 4; producer++ {
wg.Add(1)
go func() {
defer wg.Done()
<-start
for i := 0; i < 25; i++ {
h.Deliver(mkEvent("foo", "event"))
h.Broadcast(transport.SourceState{Type: transport.FrameTypeSourceState, State: "connected"})
}
}()
}
wg.Add(1)
go func() {
defer wg.Done()
<-start
h.Unregister(c.ID)
}()
close(start)
wg.Wait()
deadline := time.After(time.Second)
drain:
for {
select {
case _, ok := <-c.SendCh:
if !ok {
break drain
}
case <-deadline:
t.Fatal("SendCh was not closed after concurrent unregister")
}
}
seq := c.seq.Load()
received := c.received.Load()
dropped := c.dropped.Load()
h.Deliver(mkEvent("foo", "after-close"))
h.Broadcast(transport.Bye{Type: transport.FrameTypeBye})
if c.seq.Load() != seq || c.received.Load() != received || c.dropped.Load() != dropped {
t.Fatal("closed consumer counters changed after unregister")
}
}
}
func TestHub_BroadcastReachesAllConsumers(t *testing.T) {
h := NewHub(10)
a, _ := h.Register(transport.Hello{ConsumerPID: 1})
b, _ := h.Register(transport.Hello{ConsumerPID: 2})
bye := transport.Bye{Type: transport.FrameTypeBye, Reason: "shutdown"}
if got := h.Broadcast(bye); got != 2 {
t.Fatalf("Broadcast delivered to %d, want 2", got)
}
for _, c := range []*Consumer{a, b} {
select {
case f := <-c.SendCh:
if _, ok := f.(transport.Bye); !ok {
t.Errorf("PID %d got %T, want Bye", c.PID, f)
}
case <-time.After(time.Second):
t.Errorf("PID %d did not receive broadcast", c.PID)
}
}
}
func TestHub_Snapshot_SortedByPID(t *testing.T) {
h := NewHub(10)
for _, pid := range []int{30, 10, 20} {
_, _ = h.Register(transport.Hello{ConsumerPID: pid, EventTypes: []string{"a"}})
}
snap := h.Snapshot()
if len(snap) != 3 || snap[0].PID != 10 || snap[1].PID != 20 || snap[2].PID != 30 {
t.Fatalf("snapshot not sorted: %+v", snap)
}
}
func TestHub_PerConsumerSeqRestartsAtOne(t *testing.T) {
h := NewHub(10)
a, _ := h.Register(transport.Hello{ConsumerPID: 1})
b, _ := h.Register(transport.Hello{ConsumerPID: 2})
h.Deliver(mkEvent("foo", "x"))
h.Deliver(mkEvent("foo", "y"))
ea := drain(a, 2, t)
eb := drain(b, 2, t)
for i, ev := range ea {
if ev.Seq != uint64(i+1) {
t.Errorf("consumer a seq[%d] = %d, want %d", i, ev.Seq, i+1)
}
}
for i, ev := range eb {
if ev.Seq != uint64(i+1) {
t.Errorf("consumer b seq[%d] = %d, want %d", i, ev.Seq, i+1)
}
}
}
func TestHub_NilEventNoOp(t *testing.T) {
h := NewHub(10)
c, _ := h.Register(transport.Hello{})
h.Deliver(nil)
if c.received.Load() != 0 {
t.Fatal("nil event should not increment")
}
}
+157
View File
@@ -0,0 +1,157 @@
// 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 bus
import (
"errors"
"fmt"
"io"
"os"
"strconv"
"strings"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/lock"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/process"
)
// LockFileName is the on-disk name of the bus single-instance lock. It lives
// inside the bus working directory
// (<ConfigDir>/events/<edition>/<source_kind>/<identity_hash>/).
const LockFileName = "bus.lock"
// ErrBusy is re-exported from lock for callers that only depend on bus.
var ErrBusy = lock.ErrBusy
// ErrStaleOwnerAlive indicates the PID stored in bus.lock points at a
// live process — there is already a bus running for this ClientID and we
// must not start another one. (This case is hit when the holder is still
// alive but its flock was somehow released; in practice flock + PID always
// agree, so this is mostly defensive.)
var ErrStaleOwnerAlive = errors.New("bus: lock file PID is alive but flock was released; assuming live owner")
// Lock represents a held bus.lock. Close releases the flock and removes the
// PID file, so a subsequent bus can acquire cleanly. A zero Lock is unusable.
type Lock struct {
inner *lock.File
}
// Acquire takes the bus lock at path and writes our PID into the file body.
//
// If the file already has a PID written by a previous run:
// 1. Try the flock first — if another process holds it, return ErrBusy
// (a live bus is running, abort).
// 2. flock acquired but file contains a PID → check if that PID is
// alive via process.Alive(). If alive → return ErrStaleOwnerAlive
// (defensive; release our flock first). If dead → take over (orphan
// cleanup) and overwrite PID with our own.
//
// On success the returned Lock owns an exclusive flock and a file body
// containing our PID. Concurrent competing processes will get ErrBusy.
func Acquire(path string) (*Lock, error) {
l, err := lock.TryAcquire(path)
if err != nil {
return nil, err // already wraps ErrBusy when busy
}
// flock acquired. Read existing PID (if any).
f := l.File()
if _, err := f.Seek(0, io.SeekStart); err != nil {
_ = l.Close()
return nil, fmt.Errorf("bus: seek lock: %w", err)
}
old, err := io.ReadAll(f)
if err != nil {
_ = l.Close()
return nil, fmt.Errorf("bus: read lock: %w", err)
}
if pid := parsePID(old); pid > 0 && process.Alive(pid) {
// Defensive: flock returned us the lock, but the stored PID is
// alive. This shouldn't normally happen (the live process holds
// the flock), but it's possible across odd kernel/FS edge cases
// (NFS, container restarts). Release and refuse to start.
_ = l.Close()
return nil, ErrStaleOwnerAlive
}
// Orphan or first-ever acquisition. Rewrite the file with our PID.
if err := truncateAndWritePID(f, os.Getpid()); err != nil {
_ = l.Close()
return nil, fmt.Errorf("bus: write PID: %w", err)
}
return &Lock{inner: l}, nil
}
// ReadHolderPID returns the PID stored in path, or 0 if the file is missing
// or unreadable. Does NOT attempt to acquire the lock — useful for `event
// status` to display the holder without contention.
func ReadHolderPID(path string) int {
b, err := os.ReadFile(path)
if err != nil {
return 0
}
return parsePID(b)
}
// Close releases the flock and best-effort blanks the PID body so a stale
// reader (e.g. `event status` racing our shutdown) does not see our
// long-dead PID and try to signal it. The lock file itself is NOT removed
// — keeping it on disk avoids a race where a competing bus could acquire
// inode-on-create faster than our truncate.
func (l *Lock) Close() error {
if l == nil || l.inner == nil {
return nil
}
// Blank the body before releasing the lock.
f := l.inner.File()
_ = truncateAndWritePID(f, 0)
err := l.inner.Close()
l.inner = nil
return err
}
// HoldsPath returns the path the lock is held on.
func (l *Lock) HoldsPath() string {
if l == nil || l.inner == nil {
return ""
}
return l.inner.Path()
}
func truncateAndWritePID(f *os.File, pid int) error {
if err := f.Truncate(0); err != nil {
return err
}
if _, err := f.Seek(0, io.SeekStart); err != nil {
return err
}
if pid > 0 {
if _, err := fmt.Fprintf(f, "%d\n", pid); err != nil {
return err
}
}
return f.Sync()
}
func parsePID(b []byte) int {
s := strings.TrimSpace(string(b))
if s == "" {
return 0
}
n, err := strconv.Atoi(s)
if err != nil || n <= 0 {
return 0
}
return n
}
+147
View File
@@ -0,0 +1,147 @@
// 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 bus
import (
"errors"
"fmt"
"os"
"path/filepath"
"testing"
)
func TestAcquire_WritesOurPID(t *testing.T) {
path := filepath.Join(t.TempDir(), LockFileName)
l, err := Acquire(path)
if err != nil {
t.Fatalf("Acquire: %v", err)
}
defer l.Close()
got := ReadHolderPID(path)
if got != os.Getpid() {
t.Fatalf("ReadHolderPID = %d, want %d", got, os.Getpid())
}
}
func TestAcquire_SecondCallerGetsBusy(t *testing.T) {
path := filepath.Join(t.TempDir(), LockFileName)
l1, err := Acquire(path)
if err != nil {
t.Fatalf("first Acquire: %v", err)
}
defer l1.Close()
_, err = Acquire(path)
if !errors.Is(err, ErrBusy) {
t.Fatalf("second Acquire = %v, want ErrBusy", err)
}
}
func TestAcquire_StaleOrphanIsReclaimed(t *testing.T) {
path := filepath.Join(t.TempDir(), LockFileName)
// Pre-populate file with a definitely-dead PID (max int32 is unlikely to be alive).
if err := os.WriteFile(path, []byte("2147483646\n"), 0o600); err != nil {
t.Fatalf("pre-populate: %v", err)
}
l, err := Acquire(path)
if err != nil {
t.Fatalf("Acquire on stale orphan: %v", err)
}
defer l.Close()
if got := ReadHolderPID(path); got != os.Getpid() {
t.Fatalf("after orphan reclaim, ReadHolderPID = %d, want %d", got, os.Getpid())
}
}
func TestAcquire_EmptyExistingFileWorks(t *testing.T) {
path := filepath.Join(t.TempDir(), LockFileName)
if err := os.WriteFile(path, []byte{}, 0o600); err != nil {
t.Fatalf("pre-create: %v", err)
}
l, err := Acquire(path)
if err != nil {
t.Fatalf("Acquire on empty file: %v", err)
}
defer l.Close()
if got := ReadHolderPID(path); got != os.Getpid() {
t.Fatalf("ReadHolderPID = %d, want %d", got, os.Getpid())
}
}
func TestClose_BlanksPID(t *testing.T) {
path := filepath.Join(t.TempDir(), LockFileName)
l, err := Acquire(path)
if err != nil {
t.Fatalf("Acquire: %v", err)
}
if err := l.Close(); err != nil {
t.Fatalf("Close: %v", err)
}
// Body should be empty (or at least not our PID anymore).
if got := ReadHolderPID(path); got != 0 {
t.Fatalf("after Close, ReadHolderPID = %d, want 0", got)
}
}
func TestReadHolderPID_MissingFileReturnsZero(t *testing.T) {
if got := ReadHolderPID(filepath.Join(t.TempDir(), "does-not-exist")); got != 0 {
t.Fatalf("ReadHolderPID(missing) = %d, want 0", got)
}
}
func TestReadHolderPID_MalformedReturnsZero(t *testing.T) {
path := filepath.Join(t.TempDir(), "junk.lock")
if err := os.WriteFile(path, []byte("not-a-pid"), 0o600); err != nil {
t.Fatalf("write: %v", err)
}
if got := ReadHolderPID(path); got != 0 {
t.Fatalf("ReadHolderPID(malformed) = %d, want 0", got)
}
}
func TestAcquire_AfterReleaseReclaimable(t *testing.T) {
path := filepath.Join(t.TempDir(), LockFileName)
l1, err := Acquire(path)
if err != nil {
t.Fatalf("first: %v", err)
}
if err := l1.Close(); err != nil {
t.Fatalf("close first: %v", err)
}
l2, err := Acquire(path)
if err != nil {
t.Fatalf("reclaim after Close: %v", err)
}
defer l2.Close()
}
// Sanity check that PID round-trips through truncateAndWritePID.
func TestTruncateAndWritePID_Roundtrip(t *testing.T) {
path := filepath.Join(t.TempDir(), "x")
f, err := os.OpenFile(path, os.O_RDWR|os.O_CREATE, 0o600)
if err != nil {
t.Fatal(err)
}
defer f.Close()
if err := truncateAndWritePID(f, 9999); err != nil {
t.Fatal(err)
}
b, _ := os.ReadFile(path)
want := fmt.Sprintf("%d\n", 9999)
if string(b) != want {
t.Fatalf("body = %q, want %q", b, want)
}
}
+103
View File
@@ -0,0 +1,103 @@
// 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 bus
import (
"encoding/json"
"fmt"
"os"
"path/filepath"
"time"
dwsevent "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
)
// MetaFileName is the on-disk name of the bus metadata file. It lives
// alongside bus.lock and bus.sock inside the bus working directory.
const MetaFileName = "bus.meta"
// Meta is the JSON document written once at bus startup. Its primary
// purpose is to let `dws event list/status --all` reverse-map directory
// names (clientIDHash hex) back to the human-readable ClientID. It also
// records bus identity for protocol-compatibility diagnostics (a future
// consume client built against bus_version="v2" can refuse to dial a
// bus_version="v1" bus, etc.).
//
// The file is overwritten on each bus startup (so a previous bus's stale
// meta does not persist past a fresh boot) and intentionally NOT deleted
// on Close — keeping it on disk helps `event status` diagnose an orphan
// (bus.lock empty + bus.meta present + PID dead = clean orphan).
type Meta struct {
ClientID string `json:"client_id"`
Edition string `json:"edition"`
SourceKind dwsevent.SourceKind `json:"source_kind,omitempty"`
IdentityHash string `json:"identity_hash,omitempty"`
SourceID string `json:"source_id,omitempty"`
StartedAt time.Time `json:"started_at"`
SDKVersion string `json:"sdk_version,omitempty"`
BusVersion string `json:"bus_version"`
BusPID int `json:"bus_pid"`
}
// CurrentBusVersion identifies the bus wire/storage compatibility level.
// Bumped only on breaking changes (IPC protocol, lockfile shape, meta
// schema). v1 is the initial value; the field is parsed defensively by
// readers (older readers tolerate unknown fields via encoding/json).
const CurrentBusVersion = "v1"
// WriteMeta atomically writes m to <dir>/bus.meta. Atomic via tmp-file +
// rename. Directory permissions are not changed; caller must mkdir the
// containing directory beforehand with pkg/config.DirPerm.
func WriteMeta(dir string, m Meta) error {
if m.BusVersion == "" {
m.BusVersion = CurrentBusVersion
}
if m.BusPID == 0 {
m.BusPID = os.Getpid()
}
if m.StartedAt.IsZero() {
m.StartedAt = time.Now().UTC()
}
b, err := json.MarshalIndent(m, "", " ")
if err != nil {
return fmt.Errorf("bus: marshal meta: %w", err)
}
final := filepath.Join(dir, MetaFileName)
tmp := final + ".tmp"
if err := os.WriteFile(tmp, b, config.FilePerm); err != nil {
return fmt.Errorf("bus: write tmp meta: %w", err)
}
if err := os.Rename(tmp, final); err != nil {
_ = os.Remove(tmp)
return fmt.Errorf("bus: rename meta: %w", err)
}
return nil
}
// ReadMeta loads and parses <dir>/bus.meta. Returns (nil, error) when the
// file is missing or malformed. Used by `event list/status --all` to
// resolve clientIDHash → original ClientID.
func ReadMeta(dir string) (*Meta, error) {
path := filepath.Join(dir, MetaFileName)
b, err := os.ReadFile(path)
if err != nil {
return nil, err
}
var m Meta
if err := json.Unmarshal(b, &m); err != nil {
return nil, fmt.Errorf("bus: parse meta %s: %w", path, err)
}
return &m, nil
}
+90
View File
@@ -0,0 +1,90 @@
// 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 bus
import (
"os"
"path/filepath"
"testing"
"time"
)
func TestWriteRead_Roundtrip(t *testing.T) {
dir := t.TempDir()
m := Meta{
ClientID: "ding_xyz",
Edition: "open",
StartedAt: time.Date(2026, 5, 28, 12, 0, 0, 0, time.UTC),
SDKVersion: "v0.9.1",
BusPID: 12345,
}
if err := WriteMeta(dir, m); err != nil {
t.Fatalf("WriteMeta: %v", err)
}
got, err := ReadMeta(dir)
if err != nil {
t.Fatalf("ReadMeta: %v", err)
}
if got.ClientID != m.ClientID || got.Edition != m.Edition || got.BusPID != m.BusPID {
t.Errorf("roundtrip mismatch: got %+v want %+v", got, m)
}
if got.BusVersion != CurrentBusVersion {
t.Errorf("BusVersion default not applied: %q", got.BusVersion)
}
}
func TestWriteMeta_DefaultsPID(t *testing.T) {
dir := t.TempDir()
if err := WriteMeta(dir, Meta{ClientID: "x", Edition: "open"}); err != nil {
t.Fatalf("WriteMeta: %v", err)
}
got, err := ReadMeta(dir)
if err != nil {
t.Fatalf("ReadMeta: %v", err)
}
if got.BusPID != os.Getpid() {
t.Errorf("BusPID default = %d, want %d", got.BusPID, os.Getpid())
}
if got.StartedAt.IsZero() {
t.Error("StartedAt default not applied")
}
}
func TestWriteMeta_AtomicNoTmpLeft(t *testing.T) {
dir := t.TempDir()
if err := WriteMeta(dir, Meta{ClientID: "x", Edition: "open"}); err != nil {
t.Fatal(err)
}
// .tmp file must not remain after successful rename
if _, err := os.Stat(filepath.Join(dir, MetaFileName+".tmp")); err == nil {
t.Fatal(".tmp file leaked after WriteMeta")
}
}
func TestReadMeta_MissingFileErrors(t *testing.T) {
if _, err := ReadMeta(t.TempDir()); err == nil {
t.Fatal("ReadMeta on missing file should error")
}
}
func TestReadMeta_MalformedErrors(t *testing.T) {
dir := t.TempDir()
if err := os.WriteFile(filepath.Join(dir, MetaFileName), []byte("{garbage"), 0o600); err != nil {
t.Fatal(err)
}
if _, err := ReadMeta(dir); err == nil {
t.Fatal("ReadMeta on malformed JSON should error")
}
}
+120
View File
@@ -0,0 +1,120 @@
// 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 bus
import (
"sort"
"sync"
"sync/atomic"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/transport"
)
// PerTypeCounters is bus-wide and tracks how many events of each event_type
// the bus delivered (or dropped due to backpressure) since startup. Surfaced
// via `dws event status`. Concurrent-safe.
//
// Implementation note: a map of *atomic uint64 pairs lets us avoid taking
// the mu in the hot Add() path; mu is only held when inserting a new
// event_type key.
type PerTypeCounters struct {
mu sync.RWMutex
m map[string]*typeRow
}
type typeRow struct {
received atomic.Uint64
dropped atomic.Uint64
}
// NewPerTypeCounters returns an empty counter set.
func NewPerTypeCounters() *PerTypeCounters {
return &PerTypeCounters{m: make(map[string]*typeRow)}
}
// AddReceived increments the received counter for eventType (allocating the
// row on first sight of a new type).
func (c *PerTypeCounters) AddReceived(eventType string) {
c.row(eventType).received.Add(1)
}
// AddDropped increments the dropped counter for eventType.
func (c *PerTypeCounters) AddDropped(eventType string) {
c.row(eventType).dropped.Add(1)
}
// Snapshot returns a deterministic point-in-time view, sorted by
// event_type. Used by status RPC encoding.
func (c *PerTypeCounters) Snapshot() map[string]transport.Counters {
c.mu.RLock()
defer c.mu.RUnlock()
out := make(map[string]transport.Counters, len(c.m))
for k, row := range c.m {
out[k] = transport.Counters{
Received: row.received.Load(),
Dropped: row.dropped.Load(),
}
}
return out
}
// SortedTypes returns the known event_types in deterministic order. Useful
// for human-readable formatting (status table).
func (c *PerTypeCounters) SortedTypes() []string {
c.mu.RLock()
defer c.mu.RUnlock()
out := make([]string, 0, len(c.m))
for k := range c.m {
out = append(out, k)
}
sort.Strings(out)
return out
}
// DropRatePercent returns rounded (dropped / (received + dropped)) * 100
// for the given event_type, or -1 if the type was never observed. Used by
// the drop-rate stderr warning when crossing a configurable threshold.
func (c *PerTypeCounters) DropRatePercent(eventType string) int {
c.mu.RLock()
row, ok := c.m[eventType]
c.mu.RUnlock()
if !ok {
return -1
}
r := row.received.Load()
d := row.dropped.Load()
total := r + d
if total == 0 {
return -1
}
return int((d * 100) / total)
}
func (c *PerTypeCounters) row(eventType string) *typeRow {
c.mu.RLock()
row, ok := c.m[eventType]
c.mu.RUnlock()
if ok {
return row
}
c.mu.Lock()
defer c.mu.Unlock()
// double-check after acquiring write lock
if row, ok = c.m[eventType]; ok {
return row
}
row = &typeRow{}
c.m[eventType] = row
return row
}
+110
View File
@@ -0,0 +1,110 @@
// 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 bus
import (
"sync"
"testing"
)
func TestPerTypeCounters_BasicAdd(t *testing.T) {
c := NewPerTypeCounters()
c.AddReceived("im.message.receive_v1")
c.AddReceived("im.message.receive_v1")
c.AddDropped("im.message.receive_v1")
c.AddReceived("approval.task")
snap := c.Snapshot()
if snap["im.message.receive_v1"].Received != 2 {
t.Errorf("im received = %d, want 2", snap["im.message.receive_v1"].Received)
}
if snap["im.message.receive_v1"].Dropped != 1 {
t.Errorf("im dropped = %d, want 1", snap["im.message.receive_v1"].Dropped)
}
if snap["approval.task"].Received != 1 {
t.Errorf("approval received = %d, want 1", snap["approval.task"].Received)
}
}
func TestPerTypeCounters_SortedTypes(t *testing.T) {
c := NewPerTypeCounters()
for _, k := range []string{"z", "a", "m", "b"} {
c.AddReceived(k)
}
got := c.SortedTypes()
want := []string{"a", "b", "m", "z"}
if len(got) != len(want) {
t.Fatalf("len = %d, want %d", len(got), len(want))
}
for i := range want {
if got[i] != want[i] {
t.Fatalf("[%d] %q, want %q", i, got[i], want[i])
}
}
}
func TestPerTypeCounters_DropRatePercent(t *testing.T) {
c := NewPerTypeCounters()
for i := 0; i < 95; i++ {
c.AddReceived("a")
}
for i := 0; i < 5; i++ {
c.AddDropped("a")
}
if got := c.DropRatePercent("a"); got != 5 {
t.Errorf("DropRatePercent = %d, want 5", got)
}
if got := c.DropRatePercent("never-seen"); got != -1 {
t.Errorf("unseen DropRatePercent = %d, want -1", got)
}
}
func TestPerTypeCounters_ConcurrentAddSameType(t *testing.T) {
c := NewPerTypeCounters()
const N = 1000
var wg sync.WaitGroup
wg.Add(N * 2)
for i := 0; i < N; i++ {
go func() { defer wg.Done(); c.AddReceived("hot") }()
go func() { defer wg.Done(); c.AddDropped("hot") }()
}
wg.Wait()
snap := c.Snapshot()
if snap["hot"].Received != N || snap["hot"].Dropped != N {
t.Fatalf("got %+v, want received=%d dropped=%d", snap["hot"], N, N)
}
}
func TestPerTypeCounters_ConcurrentNewTypes(t *testing.T) {
c := NewPerTypeCounters()
const N = 500
var wg sync.WaitGroup
wg.Add(N)
for i := 0; i < N; i++ {
i := i
go func() {
defer wg.Done()
c.AddReceived(string(rune('a'+(i%26))) + "_" + string(rune('a'+((i/26)%26))))
}()
}
wg.Wait()
// Each goroutine adds exactly 1; total received across all types must be N.
var total uint64
for _, v := range c.Snapshot() {
total += v.Received
}
if total != N {
t.Fatalf("total received = %d, want %d", total, N)
}
}
+81
View File
@@ -0,0 +1,81 @@
// 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 bus
import (
"os"
"strconv"
"time"
)
// Tunable env vars (plan §15 已决项 — surfaced for operators without
// requiring a new CLI flag for each knob). Each lookup is read-once at
// bus startup; runtime changes require a bus restart.
const (
EnvIdleTimeout = "DWS_EVENT_BUS_IDLE_TIMEOUT" // Go duration, e.g. "10m"
EnvConsumerBuffer = "DWS_EVENT_CONSUMER_BUFFER" // integer
EnvDedupLRU = "DWS_EVENT_DEDUP_LRU" // integer
EnvDropWarnPct = "DWS_EVENT_DROP_WARN_PCT" // integer 1-100
)
// ApplyEnvTuning fills in Config defaults from the env vars listed above
// for any fields the caller left at zero. The cobra layer calls this after
// constructing Config so explicit flag values still win.
//
// Defaults (when env is absent or invalid):
//
// IdleTimeout → 5m
// ConsumerBuffer → DefaultSendBuffer
// DedupCapacity → 0 (let dedup package's DefaultCapacity apply)
// DropWarnPercent → DefaultDropWarnPercent
//
// Invalid env values (non-parseable, out of range) are silently ignored
// and the default is used. We deliberately do NOT fail the bus on bad
// env input — operators shouldn't lose a daemon over a typo'd env var.
func ApplyEnvTuning(cfg *Config) {
if cfg.IdleTimeout == 0 {
if v := os.Getenv(EnvIdleTimeout); v != "" {
if d, err := time.ParseDuration(v); err == nil && d > 0 {
cfg.IdleTimeout = d
}
}
if cfg.IdleTimeout == 0 {
cfg.IdleTimeout = 5 * time.Minute
}
}
if cfg.ConsumerBuffer == 0 {
if v := os.Getenv(EnvConsumerBuffer); v != "" {
if n, err := strconv.Atoi(v); err == nil && n > 0 {
cfg.ConsumerBuffer = n
}
}
}
if cfg.DedupCapacity == 0 {
if v := os.Getenv(EnvDedupLRU); v != "" {
if n, err := strconv.Atoi(v); err == nil && n > 0 {
cfg.DedupCapacity = n
}
}
}
if cfg.DropWarnPercent == 0 {
if v := os.Getenv(EnvDropWarnPct); v != "" {
if n, err := strconv.Atoi(v); err == nil && n > 0 && n <= 100 {
cfg.DropWarnPercent = n
}
}
if cfg.DropWarnPercent == 0 {
cfg.DropWarnPercent = DefaultDropWarnPercent
}
}
}
+120
View File
@@ -0,0 +1,120 @@
// 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 bus
import (
"testing"
"time"
)
func TestApplyEnvTuning_DefaultsWhenAllUnset(t *testing.T) {
t.Setenv(EnvIdleTimeout, "")
t.Setenv(EnvConsumerBuffer, "")
t.Setenv(EnvDedupLRU, "")
t.Setenv(EnvDropWarnPct, "")
cfg := Config{}
ApplyEnvTuning(&cfg)
if cfg.IdleTimeout != 5*time.Minute {
t.Errorf("IdleTimeout default = %s, want 5m", cfg.IdleTimeout)
}
if cfg.DropWarnPercent != DefaultDropWarnPercent {
t.Errorf("DropWarnPercent default = %d, want %d",
cfg.DropWarnPercent, DefaultDropWarnPercent)
}
// ConsumerBuffer / DedupCapacity intentionally left zero — the Hub
// / dedup packages apply their own defaults from that signal.
if cfg.ConsumerBuffer != 0 {
t.Errorf("ConsumerBuffer should remain 0 (Hub picks default), got %d", cfg.ConsumerBuffer)
}
if cfg.DedupCapacity != 0 {
t.Errorf("DedupCapacity should remain 0 (dedup picks default), got %d", cfg.DedupCapacity)
}
}
func TestApplyEnvTuning_ReadsValidEnv(t *testing.T) {
t.Setenv(EnvIdleTimeout, "10m")
t.Setenv(EnvConsumerBuffer, "200")
t.Setenv(EnvDedupLRU, "16384")
t.Setenv(EnvDropWarnPct, "10")
cfg := Config{}
ApplyEnvTuning(&cfg)
if cfg.IdleTimeout != 10*time.Minute {
t.Errorf("IdleTimeout = %s, want 10m", cfg.IdleTimeout)
}
if cfg.ConsumerBuffer != 200 {
t.Errorf("ConsumerBuffer = %d, want 200", cfg.ConsumerBuffer)
}
if cfg.DedupCapacity != 16384 {
t.Errorf("DedupCapacity = %d, want 16384", cfg.DedupCapacity)
}
if cfg.DropWarnPercent != 10 {
t.Errorf("DropWarnPercent = %d, want 10", cfg.DropWarnPercent)
}
}
func TestApplyEnvTuning_ExplicitConfigWins(t *testing.T) {
t.Setenv(EnvIdleTimeout, "10m")
t.Setenv(EnvDropWarnPct, "20")
cfg := Config{
IdleTimeout: 3 * time.Minute,
DropWarnPercent: 7,
}
ApplyEnvTuning(&cfg)
if cfg.IdleTimeout != 3*time.Minute {
t.Errorf("explicit IdleTimeout overridden by env: %s", cfg.IdleTimeout)
}
if cfg.DropWarnPercent != 7 {
t.Errorf("explicit DropWarnPercent overridden by env: %d", cfg.DropWarnPercent)
}
}
func TestApplyEnvTuning_IgnoresInvalidValues(t *testing.T) {
t.Setenv(EnvIdleTimeout, "not-a-duration")
t.Setenv(EnvConsumerBuffer, "-50")
t.Setenv(EnvDedupLRU, "notanumber")
t.Setenv(EnvDropWarnPct, "150") // out of 1..100
cfg := Config{}
ApplyEnvTuning(&cfg)
if cfg.IdleTimeout != 5*time.Minute {
t.Errorf("invalid duration → default; got %s", cfg.IdleTimeout)
}
if cfg.ConsumerBuffer != 0 {
t.Errorf("negative buffer ignored → 0; got %d", cfg.ConsumerBuffer)
}
if cfg.DedupCapacity != 0 {
t.Errorf("non-numeric LRU ignored → 0; got %d", cfg.DedupCapacity)
}
if cfg.DropWarnPercent != DefaultDropWarnPercent {
t.Errorf("out-of-range pct → default; got %d", cfg.DropWarnPercent)
}
}
func TestApplyEnvTuning_IdleTimeoutZeroEnv(t *testing.T) {
t.Setenv(EnvIdleTimeout, "0")
cfg := Config{}
ApplyEnvTuning(&cfg)
// "0" is parseable but <= 0 → use default 5m, not 0 (which would mean
// "disabled" in the daemon — too dangerous as an env-driven default).
if cfg.IdleTimeout != 5*time.Minute {
t.Errorf("zero env should fall through to 5m default, got %s", cfg.IdleTimeout)
}
}
+150
View File
@@ -0,0 +1,150 @@
// 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 busctl
import (
"errors"
"fmt"
"net"
"os"
"path/filepath"
"time"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/transport"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
)
// SpawnFunc abstracts the fork-bus operation so tests can inject a fake
// instead of execing a real binary. Production callers pass busctl.Spawn.
type SpawnFunc func(SpawnConfig) (pid int, err error)
// DiscoverConfig describes one discover attempt. WorkDir holds bus.lock and
// usually (on Unix) bus.sock — see dwsevent.IPCEndpoint for the short-path
// fallback when WorkDir is too deep; the caller must mkdir it with
// pkg/config.DirPerm beforehand.
type DiscoverConfig struct {
WorkDir string
IPCEndpoint string
ClientID string
// Spawn is the fork-bus callback. Default busctl.Spawn.
Spawn SpawnFunc
// SpawnExtraArgs is forwarded to Spawn (for tests).
SpawnExtraArgs []string
// DialBackoff: initial sleep between retry dials when another process
// is spawning. Doubled each attempt up to DialMaxBackoff.
DialBackoff time.Duration
// DialMaxBackoff caps backoff.
DialMaxBackoff time.Duration
// DialDeadline caps total wall-clock time spent discovering.
DialDeadline time.Duration
}
// Default knobs. Conservative — dial is local and cheap, so retry is cheap.
const (
defaultDialBackoff = 25 * time.Millisecond
defaultDialMaxBackoff = 250 * time.Millisecond
defaultDialDeadline = 5 * time.Second
)
// Discover returns a connected net.Conn to the bus for cfg.ClientID. If the
// bus is not running, Discover forks a new one and waits for it to come up.
//
// Race-free three-step (plan §12 P3):
//
// 1. try dial IPC endpoint → success → done
// 2. failed → call Spawn (fork _bus); Spawn blocks until ready pipe says
// 'R' (or returns ErrSpawnFailed if another process won the race and
// our bus startup hit ErrBusy on the lock)
// 3. dial again — should succeed; retry with backoff up to DialDeadline
// in case Spawn succeeded but socket bind has tiny latency
//
// On concurrent Discover by N processes: only one Spawn succeeds (the
// others get ErrBusy via the bus daemon's own lock acquisition). Losers
// fall through to the retry-dial loop in step 3 and connect to the bus
// the winner brought up.
func Discover(cfg DiscoverConfig) (net.Conn, error) {
if cfg.WorkDir == "" {
return nil, errors.New("busctl: WorkDir is required")
}
if cfg.IPCEndpoint == "" {
return nil, errors.New("busctl: IPCEndpoint is required")
}
if cfg.ClientID == "" {
return nil, errors.New("busctl: ClientID is required")
}
if cfg.Spawn == nil {
cfg.Spawn = Spawn
}
if cfg.DialBackoff == 0 {
cfg.DialBackoff = defaultDialBackoff
}
if cfg.DialMaxBackoff == 0 {
cfg.DialMaxBackoff = defaultDialMaxBackoff
}
if cfg.DialDeadline == 0 {
cfg.DialDeadline = defaultDialDeadline
}
// Step 1: try dial.
if conn, err := transport.Dial(cfg.IPCEndpoint); err == nil {
return conn, nil
}
// Step 2: ensure WorkDir, then spawn _bus.
if err := os.MkdirAll(cfg.WorkDir, config.DirPerm); err != nil {
return nil, fmt.Errorf("busctl: mkdir workdir: %w", err)
}
_, spawnErr := cfg.Spawn(SpawnConfig{
ClientID: cfg.ClientID,
ExtraArgs: cfg.SpawnExtraArgs,
})
if spawnErr != nil && !errors.Is(spawnErr, ErrSpawnFailed) {
// Hard error (couldn't even exec the child). Stop here — no bus
// will come up.
return nil, fmt.Errorf("busctl: spawn bus: %w", spawnErr)
}
// spawnErr == ErrSpawnFailed → child reported it can't start, most
// commonly because another process already holds the lock. Either way
// we fall through to retry-dial — if someone else's bus IS up we'll
// connect to it.
// Step 3: retry dial until DialDeadline.
deadline := time.Now().Add(cfg.DialDeadline)
backoff := cfg.DialBackoff
var lastDialErr error
for time.Now().Before(deadline) {
conn, err := transport.Dial(cfg.IPCEndpoint)
if err == nil {
return conn, nil
}
lastDialErr = err
time.Sleep(backoff)
backoff *= 2
if backoff > cfg.DialMaxBackoff {
backoff = cfg.DialMaxBackoff
}
}
return nil, fmt.Errorf("busctl: discover deadline exceeded; last dial error: %w (spawn error: %v)", lastDialErr, spawnErr)
}
// LockPath returns the canonical bus.lock path for the given working dir.
// Centralised so consume / status / stop all agree on the location.
func LockPath(workDir string) string {
return filepath.Join(workDir, "bus.lock")
}
// MetaPath returns the canonical bus.meta path.
func MetaPath(workDir string) string {
return filepath.Join(workDir, "bus.meta")
}
+280
View File
@@ -0,0 +1,280 @@
// 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 busctl
import (
"errors"
"net"
"os"
"path/filepath"
"runtime"
"sync"
"sync/atomic"
"testing"
"time"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/transport"
)
func skipOnWindows(t *testing.T) {
t.Helper()
if runtime.GOOS == "windows" {
t.Skip("uses Unix socket; Windows transport covered by transport_windows_test.go")
}
}
func shortTempDir(t *testing.T) string {
t.Helper()
dir, err := os.MkdirTemp("/tmp", "dws-busctl-")
if err != nil {
t.Fatalf("mktemp: %v", err)
}
t.Cleanup(func() { _ = os.RemoveAll(dir) })
return dir
}
// startStubBus brings up a transport.Listener at sockPath and accepts
// connections in a goroutine, discarding the data. Returns a closer.
// Used as a stand-in for the real bus daemon in discover unit tests.
func startStubBus(t *testing.T, sockPath string) func() {
t.Helper()
l, err := transport.Listen(sockPath)
if err != nil {
t.Fatalf("startStubBus listen: %v", err)
}
done := make(chan struct{})
go func() {
for {
conn, err := l.Accept()
if err != nil {
close(done)
return
}
go func(c net.Conn) {
buf := make([]byte, 256)
for {
if _, err := c.Read(buf); err != nil {
_ = c.Close()
return
}
}
}(conn)
}
}()
return func() {
_ = l.Close()
<-done
}
}
func TestDiscover_BusAlreadyRunning_DirectDial(t *testing.T) {
skipOnWindows(t)
dir := shortTempDir(t)
sock := filepath.Join(dir, "bus.sock")
closer := startStubBus(t, sock)
defer closer()
var spawnCalled atomic.Bool
fakeSpawn := func(SpawnConfig) (int, error) {
spawnCalled.Store(true)
return 0, nil
}
conn, err := Discover(DiscoverConfig{
WorkDir: dir,
IPCEndpoint: sock,
ClientID: "ding_x",
Spawn: fakeSpawn,
})
if err != nil {
t.Fatalf("Discover: %v", err)
}
defer conn.Close()
if spawnCalled.Load() {
t.Fatal("Spawn must not be called when bus is already running")
}
}
func TestDiscover_NoBus_SpawnSucceeds(t *testing.T) {
skipOnWindows(t)
dir := shortTempDir(t)
sock := filepath.Join(dir, "bus.sock")
// fakeSpawn launches the stub bus *during* the spawn call to mirror
// the real flow (bus is up by the time spawn returns).
var closer func()
t.Cleanup(func() {
if closer != nil {
closer()
}
})
fakeSpawn := func(SpawnConfig) (int, error) {
closer = startStubBus(t, sock)
return 12345, nil
}
conn, err := Discover(DiscoverConfig{
WorkDir: dir,
IPCEndpoint: sock,
ClientID: "ding_x",
Spawn: fakeSpawn,
})
if err != nil {
t.Fatalf("Discover: %v", err)
}
defer conn.Close()
}
func TestDiscover_SpawnHardErrorFails(t *testing.T) {
skipOnWindows(t)
dir := shortTempDir(t)
sock := filepath.Join(dir, "bus.sock")
hardErr := errors.New("exec failed: not found")
fakeSpawn := func(SpawnConfig) (int, error) { return 0, hardErr }
_, err := Discover(DiscoverConfig{
WorkDir: dir,
IPCEndpoint: sock,
ClientID: "ding_x",
Spawn: fakeSpawn,
DialDeadline: 100 * time.Millisecond,
})
if err == nil {
t.Fatal("Discover should fail when Spawn returns non-ErrSpawnFailed error")
}
if !errors.Is(err, hardErr) {
t.Fatalf("err = %v, want wrap of %v", err, hardErr)
}
}
func TestDiscover_SpawnReportsFailButPeerBusComesUp(t *testing.T) {
skipOnWindows(t)
dir := shortTempDir(t)
sock := filepath.Join(dir, "bus.sock")
// Simulates the race: our Spawn loses (returns ErrSpawnFailed) but
// during retry-dial a peer bus shows up.
go func() {
time.Sleep(80 * time.Millisecond)
closer := startStubBus(t, sock)
t.Cleanup(closer)
}()
fakeSpawn := func(SpawnConfig) (int, error) { return 0, ErrSpawnFailed }
conn, err := Discover(DiscoverConfig{
WorkDir: dir,
IPCEndpoint: sock,
ClientID: "ding_x",
Spawn: fakeSpawn,
DialDeadline: 2 * time.Second,
DialBackoff: 20 * time.Millisecond,
})
if err != nil {
t.Fatalf("Discover should retry-dial after ErrSpawnFailed: %v", err)
}
defer conn.Close()
}
func TestDiscover_DeadlineExceeded(t *testing.T) {
skipOnWindows(t)
dir := shortTempDir(t)
sock := filepath.Join(dir, "bus.sock") // never created
fakeSpawn := func(SpawnConfig) (int, error) { return 0, ErrSpawnFailed }
start := time.Now()
_, err := Discover(DiscoverConfig{
WorkDir: dir,
IPCEndpoint: sock,
ClientID: "ding_x",
Spawn: fakeSpawn,
DialDeadline: 150 * time.Millisecond,
DialBackoff: 20 * time.Millisecond,
})
if err == nil {
t.Fatal("Discover should fail when bus never comes up within deadline")
}
if elapsed := time.Since(start); elapsed < 100*time.Millisecond || elapsed > 1*time.Second {
t.Errorf("deadline-driven exit took %s, expected ~150ms", elapsed)
}
}
func TestDiscover_ConcurrentCallersOnlyOneSpawn(t *testing.T) {
skipOnWindows(t)
dir := shortTempDir(t)
sock := filepath.Join(dir, "bus.sock")
// Use a mutex-guarded one-shot Spawn that actually starts the stub bus
// on first call. Concurrent callers may race: the first wins (returns
// success), the rest see ErrSpawnFailed but still retry-dial successfully.
var spawnMu sync.Mutex
var spawnCount atomic.Int32
var closer func()
t.Cleanup(func() {
if closer != nil {
closer()
}
})
fakeSpawn := func(SpawnConfig) (int, error) {
spawnMu.Lock()
defer spawnMu.Unlock()
spawnCount.Add(1)
if closer == nil {
closer = startStubBus(t, sock)
return 99, nil
}
return 0, ErrSpawnFailed
}
const N = 5
var wg sync.WaitGroup
errs := make(chan error, N)
conns := make([]net.Conn, 0, N)
connsMu := sync.Mutex{}
for i := 0; i < N; i++ {
wg.Add(1)
go func() {
defer wg.Done()
conn, err := Discover(DiscoverConfig{
WorkDir: dir,
IPCEndpoint: sock,
ClientID: "ding_x",
Spawn: fakeSpawn,
DialDeadline: 2 * time.Second,
})
if err != nil {
errs <- err
return
}
connsMu.Lock()
conns = append(conns, conn)
connsMu.Unlock()
}()
}
wg.Wait()
close(errs)
for err := range errs {
t.Errorf("Discover concurrent caller: %v", err)
}
for _, c := range conns {
_ = c.Close()
}
// Note: spawnCount can be 1..N because all goroutines fail dial first
// and call Spawn. The point is they all SUCCESSFULLY connected to the
// single bus that the first Spawn brought up.
if len(conns) != N {
t.Fatalf("only %d/%d callers got a conn", len(conns), N)
}
}
+25
View File
@@ -0,0 +1,25 @@
// 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 busctl glues the consume client to the bus daemon. It encapsulates
// the three operations a consumer needs at startup and shutdown:
//
// discover: find the running bus or fork a fresh one (race-free, plan §12
// P3 "try dial → try fork lock → retry dial")
// spawn: exec `dws event _bus --client-id <id>` as a detached
// background process (stdio detach, setsid, ready pipe handshake)
// stop: gracefully terminate the bus daemon (SIGTERM + IPC fallback)
//
// All three operations are short-lived helpers — they own no long-lived
// goroutines and return whole errors to their caller.
package busctl
+195
View File
@@ -0,0 +1,195 @@
// 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 busctl
import (
"errors"
"fmt"
"io"
"os"
"os/exec"
"strconv"
"time"
)
// ReadyFDEnv is the env var the spawned `event _bus` child inspects to find
// the ready-pipe write end. The parent passes the FD number; child opens it
// via os.NewFile(fd, "ready") and writes 'R' on success or 'E' on failure.
// 3 is the first FD slot beyond stdio in cmd.ExtraFiles.
const ReadyFDEnv = "DWS_EVENT_BUS_READY_FD"
// ReadyTimeout caps how long Spawn waits for the child to signal readiness.
// 10s is generous — bus startup is local-only work (file I/O + socket bind),
// so 1s would normally suffice; the extra headroom covers cold-start
// keychain prompts and slow CI machines.
const ReadyTimeout = 10 * time.Second
// ErrSpawnFailed is returned when the child reports startup failure via
// the ready pipe ('E' byte). The child's exit error / log file holds the
// actual cause; this sentinel just lets the caller distinguish "ready
// pipe said no" from "ready pipe timed out / closed early".
var ErrSpawnFailed = errors.New("busctl: bus child reported startup failure on ready pipe")
// ErrSpawnTimeout is returned when ReadyTimeout elapses without any signal.
var ErrSpawnTimeout = errors.New("busctl: bus child did not signal readiness within deadline")
// SpawnConfig describes one spawn attempt. ClientID is the only field
// inspected by the child; the rest govern process attributes the parent
// applies before exec.
type SpawnConfig struct {
// ExecPath is the dws binary to exec. Default os.Executable().
ExecPath string
// ClientID is passed as `--client-id` to `dws event _bus`. Required.
ClientID string
// ExtraArgs are appended after `--client-id`. Empty for normal use; tests
// pass `--extra-flag-for-test` etc.
ExtraArgs []string
// Env to pass to the child. Defaults to os.Environ(). The ReadyFDEnv
// entry is appended automatically.
Env []string
}
// Spawn forks a detached `dws event _bus --client-id <id>` child process and
// waits for it to signal readiness via the ready pipe. Returns the child's
// PID on success — the caller can then dial the bus IPC endpoint.
//
// stdio detach (plan invariant #7):
// - cmd.Stdout / cmd.Stderr set to nil so the child's own writes don't
// pollute the parent's NDJSON stream
// - Setsid on Unix so the child survives parent SIGHUP / parent exit
// - CREATE_NEW_PROCESS_GROUP on Windows (set in spawn_windows.go)
//
// Child startup (handled by the eventcmd._bus handler, P6):
// - Opens os.NewFile(<DWS_EVENT_BUS_READY_FD>, "ready")
// - On startup success → writes 'R' and closes
// - On startup failure → writes 'E' and closes (child exits)
//
// Parent (this function):
// - Holds the read end open until either 1 byte is read or ReadyTimeout
// - Returns ErrSpawnFailed for 'E', ErrSpawnTimeout otherwise
func Spawn(cfg SpawnConfig) (pid int, err error) {
if cfg.ClientID == "" {
return 0, errors.New("busctl: SpawnConfig.ClientID is required")
}
if cfg.ExecPath == "" {
execPath, err := os.Executable()
if err != nil {
return 0, fmt.Errorf("busctl: locate executable: %w", err)
}
cfg.ExecPath = execPath
}
if cfg.Env == nil {
cfg.Env = os.Environ()
}
pr, pw, err := os.Pipe()
if err != nil {
return 0, fmt.Errorf("busctl: pipe: %w", err)
}
defer pr.Close()
// pw is passed to the child; close in parent after Start so only the
// child holds the write end (so reads return EOF if child dies before
// signalling, helping us distinguish death from slow startup).
args := append([]string{"event", "_bus", "--client-id", cfg.ClientID}, cfg.ExtraArgs...)
cmd := exec.Command(cfg.ExecPath, args...)
cmd.Env = append(cfg.Env, ReadyFDEnv+"=3")
cmd.ExtraFiles = []*os.File{pw} // child sees fd 3 = pw
cmd.Stdin = nil
cmd.Stdout = nil
cmd.Stderr = nil
applyDetach(cmd) // platform-specific Setsid / new process group
if err := cmd.Start(); err != nil {
_ = pw.Close()
return 0, fmt.Errorf("busctl: start %s: %w", cfg.ExecPath, err)
}
pid = cmd.Process.Pid
// Close parent's copy of the write end immediately. Now only the child
// holds it; reading on pr will return EOF when the child exits without
// signalling, instead of blocking forever.
_ = pw.Close()
// The detached child owns its own process group. Wait in the background to
// release its process resources; the caller only retains the numeric PID.
go func() { _ = cmd.Wait() }()
// Wait for ready byte.
if err := waitReady(pr); err != nil {
return pid, err
}
return pid, nil
}
// waitReady reads exactly one byte from pr ('R' or 'E') within ReadyTimeout.
// pr is closed by the caller on return.
func waitReady(pr *os.File) error {
type result struct {
b byte
err error
}
done := make(chan result, 1)
go func() {
buf := make([]byte, 1)
n, err := io.ReadFull(pr, buf)
if err != nil {
done <- result{err: err}
return
}
if n != 1 {
done <- result{err: io.ErrUnexpectedEOF}
return
}
done <- result{b: buf[0]}
}()
select {
case res := <-done:
if res.err != nil {
if errors.Is(res.err, io.EOF) || errors.Is(res.err, io.ErrUnexpectedEOF) {
return ErrSpawnFailed // child closed pipe without writing
}
return fmt.Errorf("busctl: read ready pipe: %w", res.err)
}
switch res.b {
case 'R':
return nil
case 'E':
return ErrSpawnFailed
default:
return fmt.Errorf("busctl: unexpected ready byte %q", res.b)
}
case <-time.After(ReadyTimeout):
return ErrSpawnTimeout
}
}
// ReadyFDFromEnv returns the inherited ready pipe (or nil if not set). The
// `event _bus` command handler calls this at startup, passes the returned
// *os.File to bus.Run as Config.ReadyPipe, and the bus signals readiness
// through it.
func ReadyFDFromEnv() *os.File {
v := os.Getenv(ReadyFDEnv)
if v == "" {
return nil
}
fd, err := strconv.Atoi(v)
if err != nil || fd < 3 {
return nil
}
return os.NewFile(uintptr(fd), "dws-bus-ready")
}
+248
View File
@@ -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 busctl
import (
"errors"
"os"
"os/exec"
"runtime"
"testing"
"time"
)
// Test child mode: when DWS_BUSCTL_TEST_CHILD is set, this test binary
// acts as a fake "dws event _bus" child. It reads ReadyFDFromEnv, writes
// the byte specified by the env var, and either exits or sleeps based on
// the second env var. Used by the Spawn tests below to exercise the
// real fork path without requiring a separate test-helper binary.
//
// We piggy-back on the test binary because building a separate helper
// would require either build tags or a `TestMain` two-phase exec; the
// env-marker pattern is what Go's own os/exec tests use and stays
// confined to this file.
const (
childEnvMarker = "DWS_BUSCTL_TEST_CHILD"
// values:
// "ready" — write 'R' then sleep 30s (parent should see ready)
// "fail" — write 'E' then exit (parent should see ErrSpawnFailed)
// "silent" — exit without writing (parent should see ErrSpawnFailed via EOF)
// "stall" — sleep without writing (parent should see ErrSpawnTimeout)
// "write-stdout" — write 'R' to ready FD AND to stdout (parent verifies stdout was detached)
)
// TestMain detects the child mode marker and executes the corresponding
// behaviour before delegating to the normal test runner. Production
// invocations never have this env var set so the dispatch is a no-op.
func TestMain(m *testing.M) {
switch os.Getenv(childEnvMarker) {
case "ready":
writeReady('R')
time.Sleep(30 * time.Second)
os.Exit(0)
case "fail":
writeReady('E')
os.Exit(1)
case "silent":
// don't open the ready FD at all — let the pipe close on exec exit
os.Exit(2)
case "stall":
time.Sleep(30 * time.Second)
os.Exit(0)
case "write-stdout":
// Write to stdout BEFORE signalling ready. Parent's Spawn should
// have detached stdout to /dev/null, so the parent's stdout
// buffer (captured separately in the test) must NOT see this.
_, _ = os.Stdout.Write([]byte("POLLUTION-FROM-CHILD\n"))
writeReady('R')
time.Sleep(30 * time.Second)
os.Exit(0)
}
os.Exit(m.Run())
}
func writeReady(b byte) {
pipe := ReadyFDFromEnv()
if pipe == nil {
return
}
_, _ = pipe.Write([]byte{b})
_ = pipe.Close()
}
// spawnWithMarker invokes Spawn against the current test binary with a
// child-mode env var set. Returns the child's PID + the Spawn error.
func spawnWithMarker(t *testing.T, marker string, opts ...func(*SpawnConfig)) (int, error) {
t.Helper()
cfg := SpawnConfig{
ExecPath: os.Args[0],
ClientID: "ding_spawn_test",
Env: append(os.Environ(),
childEnvMarker+"="+marker,
),
}
for _, o := range opts {
o(&cfg)
}
return Spawn(cfg)
}
func TestSpawn_ReadySuccess(t *testing.T) {
if runtime.GOOS == "windows" {
t.Skip("Spawn uses Unix Setsid; Windows path covered separately")
}
pid, err := spawnWithMarker(t, "ready")
if err != nil {
t.Fatalf("Spawn ready: %v", err)
}
if pid <= 0 {
t.Errorf("Spawn returned non-positive pid %d", pid)
}
// Reap the child so it doesn't outlive the test.
if proc, err := os.FindProcess(pid); err == nil {
_ = proc.Kill()
_, _ = proc.Wait()
}
}
func TestSpawn_ReadyFailReportsErrSpawnFailed(t *testing.T) {
if runtime.GOOS == "windows" {
t.Skip("Spawn uses Unix Setsid")
}
pid, err := spawnWithMarker(t, "fail")
if !errors.Is(err, ErrSpawnFailed) {
t.Fatalf("err = %v, want ErrSpawnFailed", err)
}
if pid <= 0 {
t.Errorf("pid should still be reported even on fail, got %d", pid)
}
}
func TestSpawn_ChildSilentExitReportsErrSpawnFailed(t *testing.T) {
if runtime.GOOS == "windows" {
t.Skip("Spawn uses Unix Setsid")
}
_, err := spawnWithMarker(t, "silent")
if !errors.Is(err, ErrSpawnFailed) {
t.Fatalf("err = %v, want ErrSpawnFailed (EOF on pipe)", err)
}
}
func TestSpawn_StallReportsTimeout(t *testing.T) {
if runtime.GOOS == "windows" {
t.Skip("Spawn uses Unix Setsid")
}
// Temporarily shorten ReadyTimeout via a local Spawn variant. The
// production timeout is 10s — too long for a unit test. We exec
// manually with a tiny io-wait wrapper to verify the behaviour.
//
// We can't change the package-level const, so we re-implement the
// timeout part directly using the same primitives the production
// code uses.
cmd := exec.Command(os.Args[0])
cmd.Env = append(os.Environ(), childEnvMarker+"=stall", ReadyFDEnv+"=3")
pr, pw, err := os.Pipe()
if err != nil {
t.Fatal(err)
}
cmd.ExtraFiles = []*os.File{pw}
if err := cmd.Start(); err != nil {
t.Fatalf("start: %v", err)
}
defer func() {
_ = cmd.Process.Kill()
_, _ = cmd.Process.Wait()
}()
_ = pw.Close()
// Replicate waitReady with a tiny timeout.
done := make(chan error, 1)
go func() {
b := make([]byte, 1)
_, err := pr.Read(b)
done <- err
}()
select {
case <-done:
t.Fatal("child should have stalled; got data on ready pipe")
case <-time.After(200 * time.Millisecond):
// expected — child is stalling, no ready byte arrived
}
_ = pr.Close()
}
func TestSpawn_StdioDetached(t *testing.T) {
if runtime.GOOS == "windows" {
t.Skip("Spawn uses Unix Setsid; Windows stdio handling differs")
}
// Capture this process's stdout for the duration of the child run.
// If applyDetach is broken and cmd.Stdout would otherwise inherit,
// the child's "POLLUTION-FROM-CHILD" line would land in our pipe.
origStdout := os.Stdout
r, w, err := os.Pipe()
if err != nil {
t.Fatal(err)
}
os.Stdout = w
defer func() {
os.Stdout = origStdout
_ = w.Close()
_ = r.Close()
}()
pid, err := spawnWithMarker(t, "write-stdout")
if err != nil {
t.Fatalf("Spawn: %v", err)
}
defer func() {
if proc, err := os.FindProcess(pid); err == nil {
_ = proc.Kill()
_, _ = proc.Wait()
}
}()
// Close the write end on the parent side so reading r will EOF if
// nothing arrived. Give the child a beat to attempt the write.
time.Sleep(150 * time.Millisecond)
_ = w.Close()
buf := make([]byte, 256)
r.SetReadDeadline(time.Now().Add(500 * time.Millisecond))
n, _ := r.Read(buf)
if n > 0 {
t.Fatalf("child wrote %q to parent's stdout — stdio not detached!", buf[:n])
}
}
func TestReadyFDFromEnv_NoEnvReturnsNil(t *testing.T) {
t.Setenv(ReadyFDEnv, "")
if f := ReadyFDFromEnv(); f != nil {
t.Errorf("ReadyFDFromEnv with empty env should be nil, got %v", f)
}
}
func TestReadyFDFromEnv_InvalidIntReturnsNil(t *testing.T) {
t.Setenv(ReadyFDEnv, "not-an-int")
if f := ReadyFDFromEnv(); f != nil {
t.Errorf("ReadyFDFromEnv with bad value should be nil, got %v", f)
}
}
func TestReadyFDFromEnv_LowFDRejected(t *testing.T) {
// fd 0/1/2 are stdio — refusing them defends against accidental
// stdin/stdout/stderr corruption if someone misconfigures.
t.Setenv(ReadyFDEnv, "1")
if f := ReadyFDFromEnv(); f != nil {
t.Errorf("ReadyFDFromEnv should reject stdio fds, got %v", f)
}
}
+31
View File
@@ -0,0 +1,31 @@
// 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.
//go:build !windows
package busctl
import (
"os/exec"
"syscall"
)
// applyDetach configures the child to live past parent death and not share
// the parent's controlling terminal. Setsid puts the child in a new
// session, so SIGHUP on the parent's controlling tty (e.g. SSH disconnect)
// does not propagate. Setpgid is implied by Setsid.
func applyDetach(cmd *exec.Cmd) {
cmd.SysProcAttr = &syscall.SysProcAttr{
Setsid: true,
}
}
+32
View File
@@ -0,0 +1,32 @@
// 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.
//go:build windows
package busctl
import (
"os/exec"
"syscall"
)
// CREATE_NEW_PROCESS_GROUP (0x00000200) prevents the child from receiving
// the parent's Ctrl+C signal, similar in spirit to Setsid on Unix.
const createNewProcessGroup = 0x00000200
func applyDetach(cmd *exec.Cmd) {
cmd.SysProcAttr = &syscall.SysProcAttr{
CreationFlags: createNewProcessGroup,
HideWindow: true,
}
}
+290
View File
@@ -0,0 +1,290 @@
// 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 busctl
import (
"fmt"
"os"
"path/filepath"
"sort"
"time"
dwsevent "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/bus"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/process"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/transport"
)
// BusEntryState classifies a discovered bus directory's runtime state.
// Used by `dws event status/list` to render the table and by
// --fail-on-orphan to drive exit code.
type BusEntryState string
const (
// BusStateRunning: bus.lock holds an alive PID — the daemon is up.
BusStateRunning BusEntryState = "running"
// BusStateOrphan: bus.meta exists but bus.lock PID is dead. The user
// should `dws event stop --client-id <id>` (which detects the dead
// PID and unblocks fresh starts) or rm -rf the working directory.
BusStateOrphan BusEntryState = "orphan"
// BusStateNotRunning: directory exists (e.g. bus.meta retained for
// historic reasons) but bus.lock is missing or empty. Clean state.
BusStateNotRunning BusEntryState = "not_running"
)
// BusEntry is one bus working directory found on disk plus its detected
// lifecycle state. EnumerateBuses produces these; the cobra layer joins
// them with QueryStatus output to render the full status view.
type BusEntry struct {
WorkDir string `json:"workdir"`
Edition string `json:"edition"`
SourceKind dwsevent.SourceKind `json:"source_kind,omitempty"`
ClientIDHash string `json:"client_id_hash"`
IdentityHash string `json:"identity_hash,omitempty"`
HolderPID int `json:"holder_pid"`
State BusEntryState `json:"state"`
// Meta, if non-nil, lets list/status display the original ClientID
// (reverse-mapped from the hash) and the bus start time.
Meta *bus.Meta `json:"meta,omitempty"`
}
// IPCEndpoint returns the IPC endpoint for this entry. Delegates to
// dwsevent.IPCEndpoint so status/stop dial exactly where consume and the
// bus daemon bound (including the short-path fallback when WorkDir is too
// deep for sun_path).
func (e BusEntry) IPCEndpoint() string {
hash := e.ClientIDHash
if e.IdentityHash != "" {
hash = e.IdentityHash
}
return dwsevent.IPCEndpoint(e.WorkDir, e.Edition, e.SourceKind, hash)
}
// EnumerateBuses scans <configDir>/events/<editionFilter>/*/ for bus
// working directories. An empty editionFilter scans every edition
// directory found under events/.
//
// Returns a deterministic slice sorted by (edition, source_kind, identity_hash).
// Missing/inaccessible directories are skipped silently — list/status
// commands should still succeed when only some editions have ever run a
// bus.
func EnumerateBuses(configDir string, editionFilter string) ([]BusEntry, error) {
root := filepath.Join(configDir, "events")
editions, err := listSubdirs(root)
if err != nil {
// events/ might not exist if no bus ever ran — that's fine,
// surface an empty list rather than an error.
if os.IsNotExist(err) {
return nil, nil
}
return nil, fmt.Errorf("busctl: scan events dir: %w", err)
}
var out []BusEntry
for _, ed := range editions {
if editionFilter != "" && ed != editionFilter {
continue
}
editionDir := filepath.Join(root, ed)
hashDirs, err := listSubdirs(editionDir)
if err != nil {
continue
}
for _, h := range hashDirs {
candidate := filepath.Join(editionDir, h)
if isSourceKindDir(h) {
identityDirs, err := listSubdirs(candidate)
if err != nil {
continue
}
for _, ih := range identityDirs {
workDir := filepath.Join(candidate, ih)
out = append(out, inspectEntry(workDir, ed, h, ih))
}
continue
}
// Legacy v1 app-stream layout: events/<edition>/<client_hash>.
out = append(out, inspectEntry(candidate, ed, "", h))
}
}
sort.Slice(out, func(i, j int) bool {
if out[i].Edition != out[j].Edition {
return out[i].Edition < out[j].Edition
}
if out[i].SourceKind != out[j].SourceKind {
return out[i].SourceKind < out[j].SourceKind
}
return out[i].IdentityHash < out[j].IdentityHash
})
return out, nil
}
// FindBusByClientID is the "current ClientID" lookup used by `event status`
// (no --all). Returns the entry for the given (edition, clientIDHash) pair
// or nil if no bus has ever started for it. The caller derives the hash
// using event.ClientIDHash.
func FindBusByClientID(configDir, editionName, clientIDHash string) *BusEntry {
workDir := filepath.Join(configDir, "events", editionName, string(dwsevent.SourceKindAppStream), clientIDHash)
if _, err := os.Stat(workDir); err != nil {
legacy := filepath.Join(configDir, "events", editionName, clientIDHash)
if _, legacyErr := os.Stat(legacy); legacyErr != nil {
return nil
}
workDir = legacy
}
e := inspectEntry(workDir, editionName, string(dwsevent.SourceKindAppStream), clientIDHash)
return &e
}
// FindBusByIdentity looks up a bus in the source-kind-aware layout.
func FindBusByIdentity(configDir, editionName string, sourceKind dwsevent.SourceKind, identityHash string) *BusEntry {
if sourceKind == "" {
sourceKind = dwsevent.SourceKindAppStream
}
workDir := filepath.Join(configDir, "events", editionName, string(sourceKind), identityHash)
if _, err := os.Stat(workDir); err != nil {
return nil
}
e := inspectEntry(workDir, editionName, string(sourceKind), identityHash)
return &e
}
// listSubdirs returns immediate subdirectories of path, by basename. Ignores
// regular files. Returns the err from ReadDir unchanged (callers handle
// os.IsNotExist).
func listSubdirs(path string) ([]string, error) {
ents, err := os.ReadDir(path)
if err != nil {
return nil, err
}
out := make([]string, 0, len(ents))
for _, e := range ents {
if e.IsDir() {
out = append(out, e.Name())
}
}
return out, nil
}
// inspectEntry reads bus.meta + bus.lock and derives the lifecycle state.
// Never returns an error: any read failure folds into BusStateNotRunning.
func inspectEntry(workDir, editionName, sourceKindRaw, identityHash string) BusEntry {
sourceKind := dwsevent.SourceKind(sourceKindRaw)
if sourceKind == "" {
sourceKind = dwsevent.SourceKindAppStream
}
e := BusEntry{
WorkDir: workDir,
Edition: editionName,
SourceKind: sourceKind,
ClientIDHash: identityHash,
IdentityHash: identityHash,
State: BusStateNotRunning,
}
if m, err := bus.ReadMeta(workDir); err == nil {
e.Meta = m
if m.SourceKind != "" {
e.SourceKind = m.SourceKind
}
if m.IdentityHash != "" {
e.IdentityHash = m.IdentityHash
e.ClientIDHash = m.IdentityHash
}
}
pid := bus.ReadHolderPID(filepath.Join(workDir, bus.LockFileName))
e.HolderPID = pid
switch {
case pid > 0 && process.Alive(pid):
e.State = BusStateRunning
case pid > 0 && !process.Alive(pid):
e.State = BusStateOrphan
case pid == 0 && e.Meta != nil:
// meta retained but lock cleared — bus exited cleanly. Render as
// not_running (with the historical meta visible if the user asked
// for --format json).
e.State = BusStateNotRunning
}
return e
}
func isSourceKindDir(name string) bool {
return name == string(dwsevent.SourceKindAppStream) || name == string(dwsevent.SourceKindPersonalStream)
}
// DefaultStatusRPCTimeout caps how long QueryStatus waits for the bus to
// reply. 2s is generous — the bus's status_resp is a synchronous
// in-memory snapshot, sub-millisecond in practice; the timeout exists
// only to bound pathological cases (bus stuck in shutdown).
const DefaultStatusRPCTimeout = 2 * time.Second
// QueryStatus dials the bus IPC, sends Hello with Role=status, sends a
// StatusReq, reads exactly one StatusResp, and closes. Returns the
// decoded response or an error if any step fails.
//
// Used by `dws event status` and `dws event list` to fetch live
// per-consumer / per-event-type counters. The bus's handleStatusRPC
// path (see internal/event/bus/daemon.go) handles this connection
// without registering with the Hub — ad-hoc tooling does not count as
// a consumer in `status.active_consumers`.
func QueryStatus(endpoint string) (*transport.StatusResp, error) {
conn, err := transport.Dial(endpoint)
if err != nil {
return nil, fmt.Errorf("busctl: dial bus for status: %w", err)
}
defer conn.Close()
_ = conn.SetDeadline(time.Now().Add(DefaultStatusRPCTimeout))
w := transport.NewWriter(conn)
r := transport.NewReader(conn)
if err := w.WriteJSON(transport.Hello{
Type: transport.FrameTypeHello,
ConsumerPID: os.Getpid(),
Role: transport.HelloRoleStatus,
}); err != nil {
return nil, fmt.Errorf("busctl: write hello: %w", err)
}
if err := w.WriteJSON(transport.StatusReq{Type: transport.FrameTypeStatusReq}); err != nil {
return nil, fmt.Errorf("busctl: write status_req: %w", err)
}
var resp transport.StatusResp
if err := r.ReadJSON(&resp); err != nil {
return nil, fmt.Errorf("busctl: read status_resp: %w", err)
}
return &resp, nil
}
// EntryStatus combines static FS info (BusEntry) with the live RPC
// snapshot (StatusResp). For not_running / orphan entries Live is nil.
type EntryStatus struct {
Entry BusEntry `json:"entry"`
Live *transport.StatusResp `json:"live,omitempty"`
}
// QueryEntry fetches the live status for one BusEntry. Returns the entry
// wrapped with a nil Live when state != running (or when the dial fails).
// Errors from QueryStatus are folded into Live=nil so the caller's table
// rendering does not need to surface per-bus dial failures (they are
// already conveyed by State).
func QueryEntry(entry BusEntry) EntryStatus {
out := EntryStatus{Entry: entry}
if entry.State != BusStateRunning {
return out
}
live, err := QueryStatus(entry.IPCEndpoint())
if err == nil {
out.Live = live
}
return out
}
+203
View File
@@ -0,0 +1,203 @@
// 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 busctl
import (
"context"
"os"
"path/filepath"
"strconv"
"testing"
"time"
dwsevent "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/bus"
)
func TestEnumerateBuses_EmptyConfigDir(t *testing.T) {
dir := shortTempDir(t)
got, err := EnumerateBuses(dir, "")
if err != nil {
t.Fatalf("err: %v", err)
}
if got != nil {
t.Fatalf("expected nil for empty configDir, got %v", got)
}
}
// makeBusDir creates events/<edition>/<hash>/ with optional meta + lock
// content. Returns the path.
func makeBusDir(t *testing.T, configDir, ed, hash string, withMeta bool, lockPID int) string {
t.Helper()
workDir := filepath.Join(configDir, "events", ed, hash)
if err := os.MkdirAll(workDir, 0o700); err != nil {
t.Fatal(err)
}
if withMeta {
shortHash := hash
if len(shortHash) > 8 {
shortHash = shortHash[:8]
}
if err := bus.WriteMeta(workDir, bus.Meta{
ClientID: "ding_" + shortHash,
Edition: ed,
StartedAt: time.Now().UTC(),
}); err != nil {
t.Fatal(err)
}
}
if lockPID != 0 {
if err := os.WriteFile(filepath.Join(workDir, bus.LockFileName),
[]byte(strconv.Itoa(lockPID)+"\n"), 0o600); err != nil {
t.Fatal(err)
}
}
return workDir
}
func TestEnumerateBuses_DetectsRunningOrphanNotRunning(t *testing.T) {
dir := shortTempDir(t)
makeBusDir(t, dir, "open", "aaaa1111", true, os.Getpid()) // running (self pid)
makeBusDir(t, dir, "open", "bbbb2222", true, 2147483646) // orphan (dead pid)
makeBusDir(t, dir, "open", "cccc3333", true, 0) // not_running (meta only, no lock content)
makeBusDir(t, dir, "wukong", "dddd4444", true, os.Getpid()) // running, different edition
all, err := EnumerateBuses(dir, "")
if err != nil {
t.Fatalf("EnumerateBuses all: %v", err)
}
if len(all) != 4 {
t.Fatalf("expected 4 entries, got %d: %+v", len(all), all)
}
// Filter by edition.
openOnly, err := EnumerateBuses(dir, "open")
if err != nil {
t.Fatal(err)
}
if len(openOnly) != 3 {
t.Fatalf("expected 3 open-edition entries, got %d", len(openOnly))
}
for _, e := range openOnly {
if e.Edition != "open" {
t.Errorf("editionFilter leak: %+v", e)
}
}
// State classification (in the all-editions slice the entries are
// sorted by edition,hash so we can index deterministically).
byHash := map[string]BusEntry{}
for _, e := range all {
byHash[e.ClientIDHash] = e
}
if got := byHash["aaaa1111"].State; got != BusStateRunning {
t.Errorf("aaaa1111 state = %s, want running", got)
}
if got := byHash["bbbb2222"].State; got != BusStateOrphan {
t.Errorf("bbbb2222 state = %s, want orphan", got)
}
if got := byHash["cccc3333"].State; got != BusStateNotRunning {
t.Errorf("cccc3333 state = %s, want not_running", got)
}
if got := byHash["dddd4444"].State; got != BusStateRunning {
t.Errorf("dddd4444 state = %s, want running", got)
}
}
func TestEnumerateBuses_SortedDeterministic(t *testing.T) {
dir := shortTempDir(t)
makeBusDir(t, dir, "open", "zzzz", true, 0)
makeBusDir(t, dir, "open", "aaaa", true, 0)
makeBusDir(t, dir, "wukong", "bbbb", true, 0)
got, _ := EnumerateBuses(dir, "")
if len(got) != 3 {
t.Fatalf("got %d entries", len(got))
}
// expected order: open/aaaa, open/zzzz, wukong/bbbb
if got[0].ClientIDHash != "aaaa" || got[1].ClientIDHash != "zzzz" || got[2].ClientIDHash != "bbbb" {
t.Fatalf("sort order wrong:\n %+v\n %+v\n %+v", got[0], got[1], got[2])
}
}
func TestFindBusByClientID(t *testing.T) {
dir := shortTempDir(t)
makeBusDir(t, dir, "open", "aaaa", true, os.Getpid())
if e := FindBusByClientID(dir, "open", "aaaa"); e == nil {
t.Fatal("FindBusByClientID returned nil for existing entry")
} else if e.State != BusStateRunning {
t.Errorf("State = %s, want running", e.State)
}
if e := FindBusByClientID(dir, "open", "missing"); e != nil {
t.Errorf("missing entry should return nil, got %+v", e)
}
}
func TestQueryStatus_RealBusE2E(t *testing.T) {
skipOnWindows(t)
// Bring up a real bus daemon, then query it.
workDir := shortTempDir(t)
sock := filepath.Join(workDir, "bus.sock")
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
runDone := make(chan error, 1)
go func() {
runDone <- bus.Run(ctx, bus.Config{
WorkDir: workDir,
IPCEndpoint: sock,
ClientID: "ding_test_query",
Edition: "open",
Source: &fakeSrc{},
})
}()
defer func() { cancel(); <-runDone }()
waitForSocket(t, sock, 2*time.Second)
resp, err := QueryStatus(sock)
if err != nil {
t.Fatalf("QueryStatus: %v", err)
}
if resp.Bus.ClientID != "ding_test_query" {
t.Errorf("ClientID round-trip = %q", resp.Bus.ClientID)
}
if resp.Bus.Edition != "open" {
t.Errorf("Edition = %q", resp.Bus.Edition)
}
if resp.Bus.PID != os.Getpid() {
t.Errorf("Bus.PID = %d, want %d", resp.Bus.PID, os.Getpid())
}
}
// fakeSrc is a no-op SourceAdapter used by QueryStatus E2E. It just
// blocks on ctx so the bus daemon stays up long enough for the test to
// dial it.
type fakeSrc struct{}
func (fakeSrc) Start(ctx context.Context, _ dwsevent.EmitFn) error {
<-ctx.Done()
return ctx.Err()
}
// waitForSocket polls for the unix socket file. Reused by tests that
// need to dial a freshly-spawned bus.
func waitForSocket(t *testing.T, path string, timeout time.Duration) {
t.Helper()
deadline := time.Now().Add(timeout)
for time.Now().Before(deadline) {
if _, err := os.Stat(path); err == nil {
return
}
time.Sleep(10 * time.Millisecond)
}
t.Fatalf("socket %q did not appear within %s", path, timeout)
}
+95
View File
@@ -0,0 +1,95 @@
// 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 busctl
import (
"errors"
"fmt"
"os"
"time"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/bus"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/process"
)
// DefaultStopTimeout is the wall-clock budget Stop waits for the bus to
// exit after the signal is sent. 5s covers the bus's own graceful tear-down
// (broadcast Bye → consumer goroutines drain → cleanup) with margin.
const DefaultStopTimeout = 5 * time.Second
// ErrNotRunning indicates bus.lock either does not exist or its recorded
// PID is not alive. Stop returns this as a sentinel so the caller can
// distinguish "nothing to stop" from "failed to stop".
var ErrNotRunning = errors.New("busctl: bus is not running")
// StopConfig identifies the target bus and tunes timing.
type StopConfig struct {
// WorkDir holds bus.lock; Stop reads the PID from there.
WorkDir string
// Timeout is the total wall-clock budget for graceful exit. After this,
// Stop returns an error; it does NOT escalate to SIGKILL — leave that
// to the operator.
Timeout time.Duration
}
// Stop signals the bus daemon for cfg.WorkDir to exit gracefully and waits
// for the process to actually die. Returns ErrNotRunning if no bus is
// running for that work dir.
//
// Implementation note: on Unix we send SIGTERM. The bus daemon's Run loop
// watches its parent ctx for cancellation; the cobra `event _bus` command
// wires signal.NotifyContext so SIGTERM triggers ctx.Done() → graceful
// shutdown path. On Windows we use os.Process.Signal(os.Interrupt) which
// the Go runtime maps to TerminateProcess for processes outside our
// console group; for v1 that's acceptable (Windows graceful shutdown is
// future work — plan §16 v2).
func Stop(cfg StopConfig) error {
if cfg.WorkDir == "" {
return errors.New("busctl: StopConfig.WorkDir is required")
}
if cfg.Timeout == 0 {
cfg.Timeout = DefaultStopTimeout
}
pid := bus.ReadHolderPID(LockPath(cfg.WorkDir))
if pid <= 0 {
return ErrNotRunning
}
if !process.Alive(pid) {
return ErrNotRunning
}
proc, err := os.FindProcess(pid)
if err != nil {
return fmt.Errorf("busctl: find process %d: %w", pid, err)
}
if err := proc.Signal(stopSignal()); err != nil {
// On many Unix platforms Signal returns "process already finished"
// when the bus has just exited on its own — treat that as success.
if errors.Is(err, os.ErrProcessDone) {
return nil
}
return fmt.Errorf("busctl: signal bus pid=%d: %w", pid, err)
}
// Poll for actual exit.
deadline := time.Now().Add(cfg.Timeout)
for time.Now().Before(deadline) {
if !process.Alive(pid) {
return nil
}
time.Sleep(50 * time.Millisecond)
}
return fmt.Errorf("busctl: bus pid=%d did not exit within %s", pid, cfg.Timeout)
}
+140
View File
@@ -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 busctl
import (
"context"
"errors"
"os"
"os/exec"
"path/filepath"
"strconv"
"testing"
"time"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/bus"
)
func TestStop_NotRunningWhenLockMissing(t *testing.T) {
dir := shortTempDir(t)
err := Stop(StopConfig{WorkDir: dir})
if !errors.Is(err, ErrNotRunning) {
t.Fatalf("Stop on missing lock = %v, want ErrNotRunning", err)
}
}
func TestStop_NotRunningWhenPIDDead(t *testing.T) {
dir := shortTempDir(t)
// Write a definitely-dead PID into bus.lock.
if err := os.WriteFile(LockPath(dir), []byte("2147483646\n"), 0o600); err != nil {
t.Fatal(err)
}
err := Stop(StopConfig{WorkDir: dir})
if !errors.Is(err, ErrNotRunning) {
t.Fatalf("Stop on dead PID = %v, want ErrNotRunning", err)
}
}
func TestStop_SignalsLiveProcess(t *testing.T) {
skipOnWindows(t)
dir := shortTempDir(t)
// Spawn `sleep 10` to act as the "bus daemon".
cmd := exec.CommandContext(context.Background(), "sh", "-c", "sleep 10")
if err := cmd.Start(); err != nil {
t.Fatalf("start sleep child: %v", err)
}
defer func() {
// Best-effort cleanup if test fails.
if cmd.ProcessState == nil {
_ = cmd.Process.Kill()
_, _ = cmd.Process.Wait()
}
}()
pid := cmd.Process.Pid
// Reap the child in background so Wait doesn't leave a zombie.
waited := make(chan error, 1)
go func() { waited <- cmd.Wait() }()
// Write PID into bus.lock.
if err := os.WriteFile(LockPath(dir), []byte(strconv.Itoa(pid)+"\n"), 0o600); err != nil {
t.Fatal(err)
}
// Stop should signal SIGTERM and observe the process exit.
start := time.Now()
if err := Stop(StopConfig{WorkDir: dir, Timeout: 3 * time.Second}); err != nil {
t.Fatalf("Stop: %v", err)
}
elapsed := time.Since(start)
// sleep should react to SIGTERM almost immediately.
if elapsed > 2*time.Second {
t.Errorf("Stop took %s, expected <2s for SIGTERM-responsive child", elapsed)
}
// Confirm the child actually exited.
select {
case err := <-waited:
// sh -c "sleep 10" exits with non-zero on signal; either is fine.
_ = err
case <-time.After(2 * time.Second):
t.Fatal("child did not exit after Stop")
}
}
// TestStop_TimeoutWhenChildIgnoresSignal ensures Stop honours its deadline
// and returns a useful error instead of hanging forever.
func TestStop_TimeoutWhenChildIgnoresSignal(t *testing.T) {
skipOnWindows(t)
dir := shortTempDir(t)
// shell that traps SIGTERM and ignores it for a long time
cmd := exec.Command("sh", "-c", "trap '' TERM; sleep 30")
if err := cmd.Start(); err != nil {
t.Fatalf("start trap child: %v", err)
}
defer func() {
_ = cmd.Process.Kill()
_, _ = cmd.Process.Wait()
}()
pid := cmd.Process.Pid
if err := os.WriteFile(LockPath(dir), []byte(strconv.Itoa(pid)+"\n"), 0o600); err != nil {
t.Fatal(err)
}
err := Stop(StopConfig{WorkDir: dir, Timeout: 250 * time.Millisecond})
if err == nil {
t.Fatal("Stop should error when child ignores SIGTERM")
}
}
// TestStop_RealBusGracefulShutdown is the integration sanity check: bring
// up a real bus.Run instance, set bus.lock content to its PID, call Stop,
// and verify Run returned cleanly (via ctx done propagation in the test).
//
// NOTE: bus.Run installs its own ctx handler from the caller's ctx; here
// we don't have signal.NotifyContext (we're running in-process), so Stop's
// SIGTERM won't reach bus.Run unless we install a signal handler. Instead,
// we test the underlying primitives: PID-read, signal-send, alive-poll.
func TestStop_BusLockPathHelper(t *testing.T) {
dir := shortTempDir(t)
got := LockPath(dir)
want := filepath.Join(dir, bus.LockFileName)
if got != want {
t.Fatalf("LockPath = %q, want %q", got, want)
}
}
+26
View File
@@ -0,0 +1,26 @@
// 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.
//go:build !windows
package busctl
import (
"os"
"syscall"
)
// stopSignal returns the graceful-shutdown signal for Unix (SIGTERM). The
// `event _bus` command wires signal.NotifyContext on SIGINT/SIGTERM so
// the daemon's Run sees ctx.Done() and runs its shutdown sequence.
func stopSignal() os.Signal { return syscall.SIGTERM }
+24
View File
@@ -0,0 +1,24 @@
// 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.
//go:build windows
package busctl
import "os"
// stopSignal returns the graceful-shutdown signal for Windows. The Go
// runtime maps os.Interrupt to TerminateProcess for non-console-group
// processes — not truly graceful, but acceptable for v1 (Windows graceful
// shutdown via Ctrl+Break is in the v2 backlog, plan §16).
func stopSignal() os.Signal { return os.Interrupt }
+18
View File
@@ -0,0 +1,18 @@
// 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 consume implements the consumer-side process of `dws event
// consume`: dial the bus, send Hello, read Event frames, format them, and
// write them out (stdout / file / dir). v1 (P3) implements the minimal
// path — NDJSON to stdout. P4 adds filter/format/route/compact pipeline.
package consume
@@ -0,0 +1,123 @@
// 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 consume
import (
"context"
"io"
"path/filepath"
"testing"
"time"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/bus"
)
// TestRun_DurationExitsCleanly verifies that --duration triggers a
// graceful exit (nil return, no error surfaced) and does so within a
// small multiple of the requested duration. The contract: --duration
// is a wall-clock budget, not an "abort" — Run wraps the caller's ctx
// with WithTimeout and returns nil rather than DeadlineExceeded so the
// exit code stays 0.
func TestRun_DurationExitsCleanly(t *testing.T) {
skipOnWindows(t)
dir, sock, cancel, runDone, trigger := bringUpBus(t, nil)
defer func() { cancel(); <-runDone }()
close(trigger) // no events to fire — Run will exit on duration alone
duration := 200 * time.Millisecond
start := time.Now()
err := Run(context.Background(), Config{
WorkDir: dir,
IPCEndpoint: sock,
ClientID: "ding_test",
Stdout: io.Discard,
Stderr: io.Discard,
Duration: duration,
})
if err != nil {
t.Fatalf("Run with --duration should exit nil, got %v", err)
}
elapsed := time.Since(start)
if elapsed < duration {
t.Errorf("Run returned before duration elapsed: %s < %s", elapsed, duration)
}
if elapsed > duration+2*time.Second {
t.Errorf("Run took %s, much longer than duration %s", elapsed, duration)
}
}
// TestRun_DurationZeroMeansUnlimited verifies the documented "0 = no
// limit" semantic of --duration. We start with a small parent-ctx
// timeout to bound the test runtime; Run should respect that ctx
// instead of having its own (zero) duration trigger.
func TestRun_DurationZeroMeansUnlimited(t *testing.T) {
skipOnWindows(t)
dir, sock, cancel, runDone, trigger := bringUpBus(t, nil)
defer func() { cancel(); <-runDone }()
close(trigger)
ctx, ctxCancel := context.WithTimeout(context.Background(), 150*time.Millisecond)
defer ctxCancel()
start := time.Now()
err := Run(ctx, Config{
WorkDir: dir,
IPCEndpoint: sock,
ClientID: "ding_test",
Stdout: io.Discard,
Stderr: io.Discard,
Duration: 0, // explicitly unlimited
})
if err != nil {
t.Fatalf("Run with --duration=0 should exit nil (ctx cancel), got %v", err)
}
// Run should respect the parent ctx — not Duration. So it returns
// at roughly the ctx deadline.
if elapsed := time.Since(start); elapsed < 100*time.Millisecond {
t.Errorf("Run returned too quickly (%s); --duration=0 should defer to parent ctx", elapsed)
}
}
// TestRun_DryRunPrintsConfigAndExits verifies --dry-run is end-to-end
// observable: Run never dials the bus (so even with a bogus IPC endpoint
// it returns nil) and writes the config block to Stderr.
func TestRun_DryRunDoesNotDial(t *testing.T) {
// We deliberately give a non-existent endpoint to prove Run does
// not try to dial.
dir := shortTempDir(t)
bogusSock := filepath.Join(dir, "no-such.sock")
err := Run(context.Background(), Config{
WorkDir: dir,
IPCEndpoint: bogusSock,
ClientID: "ding_test",
Stdout: io.Discard,
Stderr: io.Discard,
DryRun: true,
})
if err != nil {
t.Fatalf("DryRun should bypass dial and return nil, got %v", err)
}
}
// Sanity check that bus.ApplyEnvTuning is wired the same way Run reads
// from Config.Duration — both should be additive and not interfere.
func TestApplyEnvTuning_DoesNotTouchDuration(t *testing.T) {
// Duration is a consume.Config field, not bus.Config — but we still
// want a smoke test that ApplyEnvTuning doesn't accidentally reach
// into the consume layer.
cfg := bus.Config{}
bus.ApplyEnvTuning(&cfg)
// (no Duration field on bus.Config; this test compiles only if the
// invariant holds — caught by reviewer if someone adds one)
_ = cfg
}
+157
View File
@@ -0,0 +1,157 @@
// 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 consume
import (
"encoding/json"
"fmt"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/registry"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/transport"
)
// Format identifies the wire shape `dws event consume` writes per event.
// Values mirror dws's global -f/--format flag vocabulary (defined in
// internal/output) with the subset that makes sense for streaming.
type Format string
const (
// FormatNDJSON is the default: one compact JSON object per line.
// Pipe-friendly; one event per `read` line. Recommended for agents.
FormatNDJSON Format = "ndjson"
// FormatJSON pretty-prints each event as multi-line JSON. NOT a
// JSON array — still NDJSON-style (one document per output unit),
// just with indentation. See plan §3.1 输出约束 note about why we
// do not emit a JSON array for an unbounded stream.
FormatJSON Format = "json"
// FormatPretty is the same as FormatJSON in v1; reserved for future
// human-friendly colorisation. Kept distinct so we never silently
// degrade `--format pretty` to compact-ndjson.
FormatPretty Format = "pretty"
// FormatRaw writes only the SDK's original Data string (one per
// event, newline-terminated). Useful when piping into jq / a tool
// that wants the cloud payload verbatim without our envelope.
FormatRaw Format = "raw"
// FormatCompact runs the per-event-type compact processor (see
// registry.LookupProcessor) and emits one flattened JSON line.
FormatCompact Format = "compact"
)
// NormalizeFormat maps a raw flag value to a supported Format. Values
// outside the event command's supported set fall back to NDJSON with the
// fallback flag set true — callers SHOULD warn on stderr when fallback is
// true and the original value was non-empty (e.g. user passed
// --format table which has no meaning for an event stream).
//
// Empty input maps to NDJSON without a fallback warning.
func NormalizeFormat(raw string) (f Format, fellback bool) {
switch raw {
case "":
return FormatNDJSON, false
case string(FormatNDJSON):
return FormatNDJSON, false
case string(FormatJSON):
return FormatJSON, false
case string(FormatPretty):
return FormatPretty, false
case string(FormatRaw):
return FormatRaw, false
case string(FormatCompact):
return FormatCompact, false
default:
// Includes table/csv from the global -f vocabulary, plus any
// typo. Fall back to ndjson (the safe stream default) and let
// the caller stderr-WARN.
return FormatNDJSON, true
}
}
// Formatter renders a transport.Event into the byte stream the sink writes
// out. Implementations append their own line terminator when appropriate
// (NDJSON / Raw add '\n'; Pretty/JSON embed newlines in the JSON itself).
type Formatter interface {
Render(ev transport.Event) ([]byte, error)
}
// NewFormatter returns a Formatter for the given Format. Compact wraps
// registry.LookupProcessor so adding a new specialised compactor is just
// a registry-side change. Returns an error only if format is an internally
// unsupported value (defensive — NormalizeFormat guarantees the input is
// one of the constants).
func NewFormatter(format Format) (Formatter, error) {
switch format {
case FormatNDJSON:
return &ndjsonFormatter{}, nil
case FormatJSON, FormatPretty:
return &prettyFormatter{}, nil
case FormatRaw:
return &rawFormatter{}, nil
case FormatCompact:
return &compactFormatter{}, nil
default:
return nil, fmt.Errorf("consume: unsupported format %q", format)
}
}
// ndjsonFormatter encodes each Event as one compact JSON line + '\n'.
type ndjsonFormatter struct{}
func (ndjsonFormatter) Render(ev transport.Event) ([]byte, error) {
b, err := json.Marshal(ev)
if err != nil {
return nil, err
}
return append(b, '\n'), nil
}
// prettyFormatter encodes each Event as multi-line indented JSON + '\n'.
// json.MarshalIndent does not append a trailing newline; we add one so
// successive events are visually separated in the output.
type prettyFormatter struct{}
func (prettyFormatter) Render(ev transport.Event) ([]byte, error) {
b, err := json.MarshalIndent(ev, "", " ")
if err != nil {
return nil, err
}
return append(b, '\n'), nil
}
// rawFormatter writes ev.Data verbatim. If Data is already JSON it stays
// JSON; if it's some other string format it stays that. A trailing newline
// is appended so successive raw events are separable.
type rawFormatter struct{}
func (rawFormatter) Render(ev transport.Event) ([]byte, error) {
out := make([]byte, 0, len(ev.Data)+1)
out = append(out, ev.Data...)
if len(ev.Data) == 0 || ev.Data[len(ev.Data)-1] != '\n' {
out = append(out, '\n')
}
return out, nil
}
// compactFormatter dispatches to the registry per event_type and writes
// the flattened map as one compact JSON line + '\n'.
type compactFormatter struct{}
func (compactFormatter) Render(ev transport.Event) ([]byte, error) {
p := registry.LookupProcessor(ev.EventType)
v := p(ev)
b, err := json.Marshal(v)
if err != nil {
return nil, err
}
return append(b, '\n'), nil
}
+156
View File
@@ -0,0 +1,156 @@
// 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 consume
import (
"encoding/json"
"strings"
"testing"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/transport"
)
func sampleEvent() transport.Event {
return transport.Event{
Type: transport.FrameTypeEvent,
Seq: 42,
EventID: "ev_abc",
EventBornTime: 1700000000123,
EventType: "im.message.receive_v1",
EventCorpID: "corp_x",
EventUnifiedAppID: "app_y",
Data: `{"message":{"message_id":"om_x","chat_id":"oc_y","content":"hi"},"sender":{"sender_id":{"open_id":"ou_z"}}}`,
ReceivedAtUnixMS: 1700000000999,
}
}
func TestNormalizeFormat(t *testing.T) {
cases := []struct {
in string
want Format
fellback bool
}{
{"", FormatNDJSON, false},
{"ndjson", FormatNDJSON, false},
{"json", FormatJSON, false},
{"pretty", FormatPretty, false},
{"raw", FormatRaw, false},
{"compact", FormatCompact, false},
{"table", FormatNDJSON, true},
{"csv", FormatNDJSON, true},
{"yaml", FormatNDJSON, true}, // typo / unsupported → fallback
}
for _, c := range cases {
got, fb := NormalizeFormat(c.in)
if got != c.want || fb != c.fellback {
t.Errorf("NormalizeFormat(%q) = (%s, %v), want (%s, %v)", c.in, got, fb, c.want, c.fellback)
}
}
}
func TestNDJSONFormatter_OneLinePerEvent(t *testing.T) {
f, _ := NewFormatter(FormatNDJSON)
out, err := f.Render(sampleEvent())
if err != nil {
t.Fatal(err)
}
if !strings.HasSuffix(string(out), "\n") {
t.Fatal("ndjson output must end with \\n")
}
lines := strings.Split(strings.TrimRight(string(out), "\n"), "\n")
if len(lines) != 1 {
t.Fatalf("ndjson must be exactly one line, got %d: %s", len(lines), out)
}
// Must be valid JSON
var ev transport.Event
if err := json.Unmarshal([]byte(lines[0]), &ev); err != nil {
t.Fatalf("not valid JSON: %v", err)
}
if ev.EventID != "ev_abc" {
t.Errorf("round-trip lost EventID: %q", ev.EventID)
}
}
func TestPrettyFormatter_MultilineIndented(t *testing.T) {
f, _ := NewFormatter(FormatPretty)
out, err := f.Render(sampleEvent())
if err != nil {
t.Fatal(err)
}
s := string(out)
if !strings.Contains(s, "\n ") {
t.Fatal("pretty output should have 2-space indentation")
}
if !strings.HasSuffix(s, "\n") {
t.Fatal("pretty output should end with \\n")
}
// Strip trailing newline and ensure round-trip works (still valid JSON).
var ev transport.Event
if err := json.Unmarshal([]byte(strings.TrimRight(s, "\n")), &ev); err != nil {
t.Fatalf("pretty not valid JSON: %v", err)
}
}
func TestRawFormatter_WritesDataVerbatim(t *testing.T) {
f, _ := NewFormatter(FormatRaw)
out, err := f.Render(transport.Event{Data: `{"foo":"bar"}`})
if err != nil {
t.Fatal(err)
}
if string(out) != "{\"foo\":\"bar\"}\n" {
t.Errorf("raw output = %q", out)
}
}
func TestRawFormatter_PreservesExistingTrailingNewline(t *testing.T) {
f, _ := NewFormatter(FormatRaw)
out, _ := f.Render(transport.Event{Data: "already-ends-newline\n"})
if string(out) != "already-ends-newline\n" {
t.Errorf("raw should not double the trailing \\n, got %q", out)
}
}
func TestRawFormatter_EmptyDataYieldsBareNewline(t *testing.T) {
f, _ := NewFormatter(FormatRaw)
out, _ := f.Render(transport.Event{})
if string(out) != "\n" {
t.Errorf("empty Data → %q, want bare \\n", out)
}
}
func TestCompactFormatter_DispatchesPerEventType(t *testing.T) {
f, _ := NewFormatter(FormatCompact)
out, err := f.Render(sampleEvent())
if err != nil {
t.Fatal(err)
}
var got map[string]any
if err := json.Unmarshal([]byte(strings.TrimRight(string(out), "\n")), &got); err != nil {
t.Fatalf("compact not valid JSON: %v", err)
}
// IM message processor should have lifted message_id/chat_id/etc.
if got["message_id"] != "om_x" || got["chat_id"] != "oc_y" {
t.Fatalf("compact output missing lifted fields: %+v", got)
}
// Header field `type` must equal event_type.
if got["type"] != "im.message.receive_v1" {
t.Errorf("type = %v", got["type"])
}
}
func TestNewFormatter_RejectsUnknown(t *testing.T) {
if _, err := NewFormatter(Format("nope")); err == nil {
t.Fatal("expected error for unknown format")
}
}
+316
View File
@@ -0,0 +1,316 @@
// 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 consume
import (
"bytes"
"context"
"encoding/json"
"errors"
"io"
"os"
"path/filepath"
"strings"
"sync"
"testing"
"time"
dwsevent "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/bus"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/transport"
)
// TestIntegration_HelloPushdownFiltersAtBus verifies the Hello-time
// event_types pushdown contract (plan §4 unsung superpower): a consumer
// subscribing to "im.*" must NOT receive "approval.*" events even when
// bus and source are flowing both. The filter happens at the Hub layer,
// not at the consumer pipeline — saves IPC bytes for narrow consumers.
func TestIntegration_HelloPushdownFiltersAtBus(t *testing.T) {
skipOnWindows(t)
events := []dwsevent.RawEvent{
{EventID: "1", EventType: "im.message.receive_v1", Data: `{}`},
{EventID: "2", EventType: "approval.task", Data: `{}`},
{EventID: "3", EventType: "im.message.at_v1", Data: `{}`},
{EventID: "4", EventType: "approval.instance.status_changed", Data: `{}`},
{EventID: "5", EventType: "im.chat.member.user.added_v1", Data: `{}`},
}
dir, sock, cancel, runDone, trigger := bringUpBus(t, events)
defer func() { cancel(); <-runDone }()
var imBuf, approvalBuf bytes.Buffer
var wg sync.WaitGroup
wg.Add(2)
// Consumer A: im.* only
go func() {
defer wg.Done()
_ = Run(context.Background(), Config{
WorkDir: dir,
IPCEndpoint: sock,
ClientID: "ding_test",
Stdout: &imBuf,
Stderr: io.Discard,
EventTypes: []string{"im.*"},
MaxEvents: 3, // 3 im events expected
})
}()
// Consumer B: approval.* only
go func() {
defer wg.Done()
_ = Run(context.Background(), Config{
WorkDir: dir,
IPCEndpoint: sock,
ClientID: "ding_test",
Stdout: &approvalBuf,
Stderr: io.Discard,
EventTypes: []string{"approval.*"},
MaxEvents: 2, // 2 approval events expected
})
}()
// Give both consumers time to Hello + register.
time.Sleep(200 * time.Millisecond)
close(trigger)
wg.Wait()
// Verify consumer A got exactly the 3 im.* events.
imLines := nonEmptyLines(imBuf.String())
if len(imLines) != 3 {
t.Fatalf("im consumer got %d events, want 3:\n%s", len(imLines), imBuf.String())
}
for i, line := range imLines {
var ev transport.Event
if err := json.Unmarshal([]byte(line), &ev); err != nil {
t.Fatalf("im[%d] not valid JSON: %v", i, err)
}
if !strings.HasPrefix(ev.EventType, "im.") {
t.Errorf("im consumer got non-im event: %s", ev.EventType)
}
}
// Verify consumer B got exactly the 2 approval.* events.
apprLines := nonEmptyLines(approvalBuf.String())
if len(apprLines) != 2 {
t.Fatalf("approval consumer got %d events, want 2:\n%s", len(apprLines), approvalBuf.String())
}
for i, line := range apprLines {
var ev transport.Event
if err := json.Unmarshal([]byte(line), &ev); err != nil {
t.Fatalf("approval[%d] not valid JSON: %v", i, err)
}
if !strings.HasPrefix(ev.EventType, "approval.") {
t.Errorf("approval consumer got non-approval event: %s", ev.EventType)
}
}
}
// TestIntegration_FilterRegexInAdditionToEventTypes verifies the
// regex --filter is applied on top of --event-types (logical AND).
// Both narrow the stream; the test confirms only events matching BOTH
// surface to the consumer.
func TestIntegration_FilterRegexNarrowsFurther(t *testing.T) {
skipOnWindows(t)
events := []dwsevent.RawEvent{
{EventID: "1", EventType: "im.message.receive_v1", Data: `{}`},
{EventID: "2", EventType: "im.message.at_v1", Data: `{}`},
{EventID: "3", EventType: "im.chat.member.user.added_v1", Data: `{}`},
}
dir, sock, cancel, runDone, trigger := bringUpBus(t, events)
defer func() { cancel(); <-runDone }()
var buf bytes.Buffer
consumeDone := make(chan struct{})
go func() {
defer close(consumeDone)
_ = Run(context.Background(), Config{
WorkDir: dir,
IPCEndpoint: sock,
ClientID: "ding_test",
Stdout: &buf,
Stderr: io.Discard,
EventTypes: []string{"im.*"},
Filter: `\.at_v1$`, // only at_v1 events
MaxEvents: 1,
})
}()
time.Sleep(150 * time.Millisecond)
close(trigger)
<-consumeDone
lines := nonEmptyLines(buf.String())
if len(lines) != 1 {
t.Fatalf("expected 1 event after im.* + .at_v1 regex, got %d:\n%s", len(lines), buf.String())
}
var ev transport.Event
_ = json.Unmarshal([]byte(lines[0]), &ev)
if ev.EventType != "im.message.at_v1" {
t.Errorf("got %q, want im.message.at_v1", ev.EventType)
}
}
// TestIntegration_BusRestartConsumerReconnects verifies bus death + restart
// scenario: a consumer dialing after the first bus died and was replaced
// should connect to the fresh bus and receive new events. This proves
// stale lock cleanup + fresh-bus startup flow work correctly.
func TestIntegration_BusRestartConsumerReconnects(t *testing.T) {
skipOnWindows(t)
dir := shortTempDir(t)
sock := filepath.Join(dir, "bus.sock")
ctx1, cancel1 := context.WithCancel(context.Background())
run1Done := make(chan error, 1)
go func() {
run1Done <- bus.Run(ctx1, bus.Config{
WorkDir: dir,
IPCEndpoint: sock,
ClientID: "ding_test",
Edition: "open",
Source: &fakeSource{},
})
}()
waitForSock(t, sock, 2*time.Second)
// Verify first bus is up by dialing it briefly.
conn, err := transport.Dial(sock)
if err != nil {
t.Fatalf("dial first bus: %v", err)
}
conn.Close()
// Kill the first bus.
cancel1()
if err := <-run1Done; err != nil && !errors.Is(err, context.Canceled) {
t.Fatalf("first bus exited with: %v", err)
}
// Start second bus on same workdir (stale lock should be reclaimed).
ctx2, cancel2 := context.WithCancel(context.Background())
defer cancel2()
run2Done := make(chan error, 1)
events := []dwsevent.RawEvent{
{EventID: "post_restart", EventType: "im.message.receive_v1", Data: `{}`},
}
trigger := make(chan struct{})
src := &fakeSource{events: events, trigger: trigger}
go func() {
run2Done <- bus.Run(ctx2, bus.Config{
WorkDir: dir,
IPCEndpoint: sock,
ClientID: "ding_test",
Edition: "open",
Source: src,
})
}()
waitForSock(t, sock, 2*time.Second)
// Consumer dials the second bus and receives the new event.
var buf bytes.Buffer
consumeDone := make(chan error, 1)
go func() {
consumeDone <- Run(context.Background(), Config{
WorkDir: dir,
IPCEndpoint: sock,
ClientID: "ding_test",
Stdout: &buf,
Stderr: io.Discard,
MaxEvents: 1,
})
}()
time.Sleep(150 * time.Millisecond)
close(trigger)
if err := <-consumeDone; err != nil {
t.Fatalf("consume on restarted bus: %v", err)
}
lines := nonEmptyLines(buf.String())
if len(lines) != 1 {
t.Fatalf("expected 1 event after restart, got %d", len(lines))
}
var ev transport.Event
_ = json.Unmarshal([]byte(lines[0]), &ev)
if ev.EventID != "post_restart" {
t.Errorf("EventID = %q, want post_restart", ev.EventID)
}
cancel2()
<-run2Done
}
// TestIntegration_PipelineWithRouteAndOutputDir is the end-to-end version
// of the unit pipeline tests: a real bus + a real consume.Run process
// configured with --route and --output-dir. Verifies that matched
// events land in route dirs and unmatched events in the fallback dir.
func TestIntegration_PipelineWithRouteAndOutputDir(t *testing.T) {
skipOnWindows(t)
events := []dwsevent.RawEvent{
{EventID: "im1", EventType: "im.message.receive_v1", Data: `{}`},
{EventID: "ap1", EventType: "approval.task", Data: `{}`},
{EventID: "im2", EventType: "im.message.at_v1", Data: `{}`},
}
dir, sock, cancel, runDone, trigger := bringUpBus(t, events)
defer func() { cancel(); <-runDone }()
outputRoot := shortTempDir(t)
imDir := filepath.Join(outputRoot, "im")
defaultDir := filepath.Join(outputRoot, "default")
routes, _ := ParseRoutes([]string{`^im\.=dir:` + imDir})
consumeDone := make(chan struct{})
go func() {
defer close(consumeDone)
_ = Run(context.Background(), Config{
WorkDir: dir,
IPCEndpoint: sock,
ClientID: "ding_test",
Stdout: io.Discard,
Stderr: io.Discard,
OutputDir: defaultDir,
Routes: routes,
MaxEvents: 3,
})
}()
time.Sleep(150 * time.Millisecond)
close(trigger)
<-consumeDone
if entries, _ := os.ReadDir(imDir); len(entries) != 2 {
t.Errorf("imDir got %d files, want 2 (im1+im2)", len(entries))
}
if entries, _ := os.ReadDir(defaultDir); len(entries) != 1 {
t.Errorf("defaultDir got %d files, want 1 (ap1)", len(entries))
}
}
func nonEmptyLines(s string) []string {
parts := strings.Split(strings.TrimRight(s, "\n"), "\n")
out := make([]string, 0, len(parts))
for _, p := range parts {
if p != "" {
out = append(out, p)
}
}
return out
}
// waitForSock polls for the unix socket file. Used by bus-restart tests
// that bring up bus.Run twice in the same workdir.
func waitForSock(t *testing.T, path string, timeout time.Duration) {
t.Helper()
deadline := time.Now().Add(timeout)
for time.Now().Before(deadline) {
if _, err := os.Stat(path); err == nil {
return
}
time.Sleep(10 * time.Millisecond)
}
t.Fatalf("socket %q did not appear within %s", path, timeout)
}
+86
View File
@@ -0,0 +1,86 @@
// 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 consume
import (
"io"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/transport"
)
// Pipeline is the consumer-side delivery chain: format → route → sink.
// Each delivered event goes through formatting once; the routed sink then
// dispatches the formatted bytes to either a route-specific directory or
// the fallback (stdout/file).
//
// A Pipeline is bound to a single Config snapshot. Reconfiguring (changing
// format / routes mid-stream) is out of scope for v1.
type Pipeline struct {
formatter Formatter
sink Sink
}
// NewPipeline builds a Pipeline for the given formatter and sink.
func NewPipeline(formatter Formatter, sink Sink) *Pipeline {
return &Pipeline{formatter: formatter, sink: sink}
}
// Deliver renders ev with the configured formatter and hands the result
// to the sink. Returns ErrPipeClosed (re-raised from the sink) when the
// downstream stdout pipe closed; otherwise returns whatever formatting or
// sink error surfaced.
func (p *Pipeline) Deliver(ev transport.Event) error {
body, err := p.formatter.Render(ev)
if err != nil {
return err
}
return p.sink.Write(ev, body)
}
// Close releases sink resources. Safe to call multiple times because
// underlying Sink Close methods are idempotent.
func (p *Pipeline) Close() error { return p.sink.Close() }
// BuildPipeline constructs a Pipeline from the cobra-side flag bundle. The
// cobra command first parses --format / --output-dir / --route into the
// derived inputs here so this function stays free of cobra dependencies.
//
// Sink selection rules (plan §3.1 输出约束):
// - --route present → routed sink with per-rule dirs;
// fallback is --output-dir if set, else stdout
// - --output-dir only → file-per-event sink at the dir
// - neither → stdout sink with stdoutW
//
// stdoutW is injected for tests (os.Stdout in production). When nil it
// defaults to io.Discard so a misconfigured pipeline never writes to
// the host process's actual stdout.
func BuildPipeline(format Format, outputDir string, routes []Route, stdoutW io.Writer) (*Pipeline, error) {
fmter, err := NewFormatter(format)
if err != nil {
return nil, err
}
if stdoutW == nil {
stdoutW = io.Discard
}
var fallback Sink
if outputDir != "" {
fallback = NewFileDirSink(outputDir)
} else {
fallback = NewStdoutSink(stdoutW)
}
if len(routes) > 0 {
return NewPipeline(fmter, NewRoutedSink(NewRouter(routes), fallback)), nil
}
return NewPipeline(fmter, fallback), nil
}
+124
View File
@@ -0,0 +1,124 @@
// 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 consume
import (
"bytes"
"encoding/json"
"os"
"path/filepath"
"strings"
"testing"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/transport"
)
func TestPipeline_FormatNDJSONToStdout(t *testing.T) {
var buf bytes.Buffer
p, err := BuildPipeline(FormatNDJSON, "", nil, &buf)
if err != nil {
t.Fatal(err)
}
defer p.Close()
for i := 0; i < 3; i++ {
_ = p.Deliver(transport.Event{Type: transport.FrameTypeEvent, EventID: "x", EventType: "y", Data: "{}"})
}
lines := strings.Split(strings.TrimRight(buf.String(), "\n"), "\n")
if len(lines) != 3 {
t.Fatalf("expected 3 NDJSON lines, got %d", len(lines))
}
for _, line := range lines {
var ev transport.Event
if err := json.Unmarshal([]byte(line), &ev); err != nil {
t.Errorf("not valid JSON: %v\n%s", err, line)
}
}
}
func TestPipeline_OutputDirFallback(t *testing.T) {
dir := t.TempDir()
p, err := BuildPipeline(FormatNDJSON, dir, nil, nil)
if err != nil {
t.Fatal(err)
}
defer p.Close()
_ = p.Deliver(transport.Event{EventType: "im.message.receive_v1", EventID: "ev1", ReceivedAtUnixMS: 100})
_ = p.Deliver(transport.Event{EventType: "im.message.receive_v1", EventID: "ev2", ReceivedAtUnixMS: 200})
entries, _ := os.ReadDir(dir)
if len(entries) != 2 {
t.Fatalf("expected 2 files in %s, got %d", dir, len(entries))
}
}
func TestPipeline_RouteWithStdoutFallback(t *testing.T) {
root := t.TempDir()
imDir := filepath.Join(root, "im")
var stdoutBuf bytes.Buffer
routes, err := ParseRoutes([]string{`^im\.=dir:` + imDir})
if err != nil {
t.Fatal(err)
}
p, err := BuildPipeline(FormatNDJSON, "", routes, &stdoutBuf)
if err != nil {
t.Fatal(err)
}
defer p.Close()
_ = p.Deliver(transport.Event{EventType: "im.message.receive_v1", EventID: "ev1", ReceivedAtUnixMS: 100})
_ = p.Deliver(transport.Event{EventType: "approval.task", EventID: "ev2", ReceivedAtUnixMS: 200})
// im event routed to imDir
if entries, _ := os.ReadDir(imDir); len(entries) != 1 {
t.Errorf("expected 1 file in im dir, got %d", len(entries))
}
// approval event fell through to stdout
if !strings.Contains(stdoutBuf.String(), `"event_id":"ev2"`) {
t.Errorf("approval event missing from stdout:\n%s", stdoutBuf.String())
}
}
func TestPipeline_RouteWithOutputDirFallback(t *testing.T) {
root := t.TempDir()
imDir := filepath.Join(root, "im")
defaultDir := filepath.Join(root, "default")
routes, _ := ParseRoutes([]string{`^im\.=dir:` + imDir})
p, _ := BuildPipeline(FormatNDJSON, defaultDir, routes, nil)
defer p.Close()
_ = p.Deliver(transport.Event{EventType: "im.message.receive_v1", EventID: "ev1", ReceivedAtUnixMS: 100})
_ = p.Deliver(transport.Event{EventType: "approval.task", EventID: "ev2", ReceivedAtUnixMS: 200})
if entries, _ := os.ReadDir(imDir); len(entries) != 1 {
t.Errorf("im events should go to imDir, got %d files", len(entries))
}
if entries, _ := os.ReadDir(defaultDir); len(entries) != 1 {
t.Errorf("unmatched events should go to default dir, got %d files", len(entries))
}
}
func TestPipeline_CompactFormat(t *testing.T) {
var buf bytes.Buffer
p, _ := BuildPipeline(FormatCompact, "", nil, &buf)
defer p.Close()
_ = p.Deliver(transport.Event{EventType: "im.message.receive_v1", EventID: "ev1", Data: `{"message":{"chat_id":"oc_x","message_id":"om_y","content":"hi"}}`})
line := strings.TrimRight(buf.String(), "\n")
var out map[string]any
if err := json.Unmarshal([]byte(line), &out); err != nil {
t.Fatalf("compact output not valid JSON: %v", err)
}
if out["chat_id"] != "oc_x" {
t.Errorf("compact missed chat_id: %+v", out)
}
}
+121
View File
@@ -0,0 +1,121 @@
// 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 consume
import (
"fmt"
"regexp"
"strings"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/transport"
)
// Route describes one --route rule. The CLI accepts the wire form
// `<regex>=dir:<path>`; ParseRoute compiles regex once at startup so the
// hot path is just a Match.
//
// Pattern matches against event.EventType (NOT the whole event JSON).
// The first matching rule in CLI order wins; unmatched events fall through
// to the default sink (stdout or --output-dir).
type Route struct {
Pattern *regexp.Regexp
Dir string
// Raw is the original CLI spec; preserved for status / debug output.
Raw string
}
// ParseRoute parses one `<regex>=dir:<path>` spec. Returns a typed error
// for bad inputs so the CLI can render a clear "did you mean" message.
//
// Wire grammar:
//
// spec = regex "=dir:" path
// regex = any chars except literal "=" (use \= to escape) (v1: no escape)
// path = any string (no validation here; sink validates at write time)
//
// Examples:
//
// "^im\\.message=dir:./im/"
// "^approval\\.=dir:./approval/"
func ParseRoute(spec string) (Route, error) {
if spec == "" {
return Route{}, fmt.Errorf("consume: empty route spec")
}
// v1 grammar is intentionally rigid: split on the first "=dir:".
// Earlier proposals supported other sink kinds (=file: / =mcp:), but
// the cobra layer rejects those — keep parsing tight here too.
const sep = "=dir:"
idx := strings.Index(spec, sep)
if idx <= 0 || idx == len(spec)-len(sep) {
return Route{}, fmt.Errorf("consume: route spec must be '<regex>=dir:<path>', got %q", spec)
}
pattern := spec[:idx]
path := spec[idx+len(sep):]
re, err := regexp.Compile(pattern)
if err != nil {
return Route{}, fmt.Errorf("consume: route regex %q: %w", pattern, err)
}
if path == "" {
return Route{}, fmt.Errorf("consume: route path is empty in %q", spec)
}
return Route{Pattern: re, Dir: path, Raw: spec}, nil
}
// ParseRoutes parses many specs in CLI order. On any parse failure returns
// the partial parse so far and the error — the caller decides whether to
// continue. (The cobra layer treats any parse error as fatal validation.)
func ParseRoutes(specs []string) ([]Route, error) {
out := make([]Route, 0, len(specs))
for _, s := range specs {
r, err := ParseRoute(s)
if err != nil {
return out, err
}
out = append(out, r)
}
return out, nil
}
// Router decides which sink an event goes to. Match returns the directory
// of the first matching Route, or empty string when no rule matches (fall
// through to default sink).
type Router struct {
rules []Route
}
// NewRouter constructs a router from pre-parsed rules.
func NewRouter(rules []Route) *Router { return &Router{rules: rules} }
// Match returns the destination directory for the event's type, or "" if
// no rule matched. Iterates rules in CLI order (first match wins).
func (r *Router) Match(ev transport.Event) string {
if r == nil {
return ""
}
for _, rule := range r.rules {
if rule.Pattern.MatchString(ev.EventType) {
return rule.Dir
}
}
return ""
}
// Rules returns the parsed routes for status / debug output. Caller MUST
// NOT mutate the returned slice.
func (r *Router) Rules() []Route {
if r == nil {
return nil
}
return r.rules
}
+104
View File
@@ -0,0 +1,104 @@
// 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 consume
import (
"testing"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/transport"
)
func TestParseRoute_Valid(t *testing.T) {
r, err := ParseRoute(`^im\.message=dir:./im/`)
if err != nil {
t.Fatalf("ParseRoute: %v", err)
}
if !r.Pattern.MatchString("im.message.receive_v1") {
t.Error("regex did not match expected event type")
}
if r.Dir != "./im/" {
t.Errorf("Dir = %q", r.Dir)
}
if r.Raw != `^im\.message=dir:./im/` {
t.Errorf("Raw = %q", r.Raw)
}
}
func TestParseRoute_Invalid(t *testing.T) {
cases := []string{
"", // empty
"no-separator", // no =dir:
"=dir:./x/", // empty regex
"^im=dir:", // empty path
"(unclosed=dir:./x/", // invalid regex
"=dir:", // both empty
}
for _, in := range cases {
if _, err := ParseRoute(in); err == nil {
t.Errorf("ParseRoute(%q) should error", in)
}
}
}
func TestParseRoutes_StopsOnFirstError(t *testing.T) {
good := `^im=dir:./im/`
bad := `(unclosed=dir:./x/`
out, err := ParseRoutes([]string{good, bad, good})
if err == nil {
t.Fatal("expected parse error")
}
if len(out) != 1 {
t.Errorf("partial parse should have 1 entry, got %d", len(out))
}
}
func TestRouter_FirstMatchWins(t *testing.T) {
routes, err := ParseRoutes([]string{
`^im\.message=dir:./im/`,
`^im\.=dir:./other-im/`,
})
if err != nil {
t.Fatal(err)
}
r := NewRouter(routes)
// First rule should win for im.message.* events.
got := r.Match(transport.Event{EventType: "im.message.receive_v1"})
if got != "./im/" {
t.Errorf("Match im.message = %q, want ./im/", got)
}
// Second rule covers im.chat.*
got = r.Match(transport.Event{EventType: "im.chat.member.bot.added_v1"})
if got != "./other-im/" {
t.Errorf("Match im.chat = %q, want ./other-im/", got)
}
}
func TestRouter_NoMatchReturnsEmpty(t *testing.T) {
routes, _ := ParseRoutes([]string{`^im\.=dir:./im/`})
r := NewRouter(routes)
if got := r.Match(transport.Event{EventType: "approval.task"}); got != "" {
t.Fatalf("no-match should return empty, got %q", got)
}
}
func TestRouter_NilSafe(t *testing.T) {
var r *Router
if got := r.Match(transport.Event{EventType: "x"}); got != "" {
t.Fatalf("nil Router.Match should return empty, got %q", got)
}
if rules := r.Rules(); rules != nil {
t.Fatalf("nil Router.Rules should return nil, got %v", rules)
}
}
+271
View File
@@ -0,0 +1,271 @@
// 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 consume
import (
"context"
"encoding/json"
"errors"
"fmt"
"io"
"net"
"os"
"time"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/busctl"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/transport"
)
// Config holds everything Run needs. Built by the cobra command handler
// (P5) from flag values + strict resolver output.
type Config struct {
// WorkDir is the bus working directory:
// <ConfigDir>/events/<edition>/<source_kind>/<identity_hash>/
WorkDir string
// IPCEndpoint is the Unix socket path / Windows pipe name. Caller
// computes from WorkDir on Unix, from edition+hash on Windows.
IPCEndpoint string
// ClientID is forwarded to busctl.Spawn so it can pass --client-id
// when forking _bus.
ClientID string
// SpawnExtraArgs are forwarded to the hidden _bus process when consume.Run
// needs to start a daemon. Used for source-mode options that must be
// reproduced in the child process, including portal ticket mode and
// personal_stream.
SpawnExtraArgs []string
// EventTypes / Filter / Compact are forwarded to the bus via Hello
// for server-side pushdown filtering.
EventTypes []string
Filter string
SubscribeID string
Compact bool
// MaxEvents: stop after receiving this many events. 0 = no limit.
MaxEvents int
// Duration: wall-clock budget for the consume run. After this elapses,
// Run returns nil (clean exit, exit code 0). Zero = no limit.
//
// Note: this is event-consume specific and intentionally NOT named
// "Timeout" — global dws --timeout is HTTP request timeout (int
// seconds) which would collide if reused. See plan §1 决策
// "事件运行时长 flag 不复用全局 --timeout".
Duration time.Duration
// DryRun, when true, prints the resolved configuration to Stderr and
// returns nil without dialing the bus. Used by the cobra layer to
// preview configuration with `--dry-run` (plan §3.1).
DryRun bool
// Foreground hint, passed through to status output but otherwise has
// no behavioural effect inside consume.Run — the cobra layer decides
// whether to call this Run or to bus.Run directly when --foreground
// is set.
Foreground bool
// Force, like Foreground, is informational at this layer. The cobra
// layer enforces the "--force requires --foreground" rule before
// calling Run.
Force bool
// --- Output / Sink config (P4) ---
// Format controls the per-event output shape (ndjson/json/pretty/raw/
// compact). The cobra layer maps --format string → Format via
// NormalizeFormat; an empty Format here defaults to NDJSON inside
// BuildPipeline.
Format Format
// OutputDir, if non-empty, switches the fallback sink from stdout to
// "file per event" under this directory.
OutputDir string
// Routes are pre-parsed --route specs. Empty = no routing.
Routes []Route
// Stdout sink; nil → os.Stdout. Injected for tests.
Stdout io.Writer
// Stderr sink for status lines (HelloAck info, bye reason); nil → os.Stderr.
// Set to io.Discard when --quiet is in effect.
Stderr io.Writer
// Quiet suppresses stderr status writes (the HelloAck / bye banners).
Quiet bool
}
// Run dials the bus (forking one if necessary), sends Hello, and writes
// each received Event frame as one NDJSON line to stdout. Blocks until
// ctx is cancelled, MaxEvents is reached, the bus sends Bye, or the
// stream is interrupted.
//
// Returns nil on graceful exits (ctx done, max-events reached, bye
// received, stdout pipe closed). Returns a non-nil error only for
// connection / protocol failures.
func Run(ctx context.Context, cfg Config) error {
if cfg.WorkDir == "" || cfg.IPCEndpoint == "" || cfg.ClientID == "" {
return errors.New("consume: WorkDir, IPCEndpoint, and ClientID are required")
}
if cfg.Stdout == nil {
cfg.Stdout = os.Stdout
}
if cfg.Stderr == nil {
cfg.Stderr = os.Stderr
}
if cfg.Quiet {
cfg.Stderr = io.Discard
}
if cfg.Format == "" {
cfg.Format = FormatNDJSON
}
// --dry-run: print resolved config, return without dialing.
if cfg.DryRun {
PrintDryRun(cfg.Stderr, cfg)
return nil
}
// --duration: layer a deadline on top of caller-provided ctx. Run
// returns nil on deadline (clean exit) rather than surfacing the
// context.DeadlineExceeded as an error to the user.
if cfg.Duration > 0 {
var cancel context.CancelFunc
ctx, cancel = context.WithTimeout(ctx, cfg.Duration)
defer cancel()
}
pipeline, err := BuildPipeline(cfg.Format, cfg.OutputDir, cfg.Routes, cfg.Stdout)
if err != nil {
return fmt.Errorf("consume: build pipeline: %w", err)
}
defer pipeline.Close()
conn, err := busctl.Discover(busctl.DiscoverConfig{
WorkDir: cfg.WorkDir,
IPCEndpoint: cfg.IPCEndpoint,
ClientID: cfg.ClientID,
SpawnExtraArgs: cfg.SpawnExtraArgs,
})
if err != nil {
return fmt.Errorf("consume: discover bus: %w", err)
}
defer conn.Close()
// Ensure the conn closes when ctx cancels so blocked Read returns.
closeOnContext(ctx, conn)
w := transport.NewWriter(conn)
r := transport.NewReader(conn)
hello := transport.Hello{
Type: transport.FrameTypeHello,
ConsumerPID: os.Getpid(),
EventTypes: cfg.EventTypes,
Filter: cfg.Filter,
SubscribeID: cfg.SubscribeID,
Compact: cfg.Compact,
}
if err := w.WriteJSON(hello); err != nil {
return fmt.Errorf("consume: write hello: %w", err)
}
var ack transport.HelloAck
if err := r.ReadJSON(&ack); err != nil {
return fmt.Errorf("consume: read hello_ack: %w", err)
}
if ack.Type != transport.FrameTypeHelloAck {
return fmt.Errorf("consume: unexpected first frame type %q", ack.Type)
}
if !cfg.Quiet {
fmt.Fprintf(cfg.Stderr,
"connected bus pid=%d source=%s state=%s idle_timeout=%ds\n",
ack.BusPID, ack.StateSource, ack.SourceState, ack.IdleTimeoutSecs)
}
received := 0
for {
raw, err := r.Read()
if err != nil {
if errors.Is(err, io.EOF) {
return nil // peer closed cleanly
}
if isCtxCancelled(ctx) {
return nil
}
return fmt.Errorf("consume: read frame: %w", err)
}
typ, err := transport.PeekType(raw)
if err != nil {
// Malformed frame; skip and continue.
continue
}
switch typ {
case transport.FrameTypeEvent:
var ev transport.Event
if err := json.Unmarshal(raw, &ev); err != nil {
continue
}
if err := pipeline.Deliver(ev); err != nil {
if errors.Is(err, ErrPipeClosed) {
// Downstream stdout consumer closed; exit cleanly.
_ = w.WriteJSON(transport.Bye{
Type: transport.FrameTypeBye,
Reason: "client_done",
})
return nil
}
return fmt.Errorf("consume: deliver event: %w", err)
}
received++
if cfg.MaxEvents > 0 && received >= cfg.MaxEvents {
_ = w.WriteJSON(transport.Bye{
Type: transport.FrameTypeBye,
Reason: "client_done",
})
return nil
}
case transport.FrameTypeBye:
var bye transport.Bye
_ = json.Unmarshal(raw, &bye)
if !cfg.Quiet {
fmt.Fprintf(cfg.Stderr, "bus closing: %s\n", bye.Reason)
}
return nil
case transport.FrameTypeSourceState:
if !cfg.Quiet {
var s transport.SourceState
_ = json.Unmarshal(raw, &s)
fmt.Fprintf(cfg.Stderr, "source state: %s (source=%s, attempt=%d)\n", s.State, s.StateSource, s.Attempt)
}
case transport.FrameTypeHeartbeat:
// silent
default:
// future frame types: ignored for forward compat
}
}
}
// closeOnContext spawns a goroutine that closes conn when ctx is done.
// This unblocks any pending Read on conn so the main loop can return.
func closeOnContext(ctx context.Context, conn net.Conn) {
go func() {
<-ctx.Done()
_ = conn.Close()
}()
}
func isCtxCancelled(ctx context.Context) bool {
select {
case <-ctx.Done():
return true
default:
return false
}
}
+311
View File
@@ -0,0 +1,311 @@
// 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 consume
import (
"bytes"
"context"
"encoding/json"
"errors"
"io"
"os"
"path/filepath"
"runtime"
"strings"
"sync"
"testing"
"time"
dwsevent "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/bus"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/transport"
)
func skipOnWindows(t *testing.T) {
t.Helper()
if runtime.GOOS == "windows" {
t.Skip("uses Unix socket")
}
}
func shortTempDir(t *testing.T) string {
t.Helper()
dir, err := os.MkdirTemp("/tmp", "dws-consume-")
if err != nil {
t.Fatalf("mktemp: %v", err)
}
t.Cleanup(func() { _ = os.RemoveAll(dir) })
return dir
}
// fakeSource mirrors the one in bus tests; reproduced here to keep the
// integration test self-contained.
type fakeSource struct {
events []dwsevent.RawEvent
trigger <-chan struct{}
}
func (f *fakeSource) Start(ctx context.Context, emit dwsevent.EmitFn) error {
if f.trigger != nil {
select {
case <-f.trigger:
case <-ctx.Done():
return ctx.Err()
}
}
for i := range f.events {
select {
case <-ctx.Done():
return ctx.Err()
default:
}
ev := f.events[i]
ev.ReceivedAt = time.Now().UTC()
emit(&ev)
time.Sleep(5 * time.Millisecond)
}
<-ctx.Done()
return ctx.Err()
}
// bringUpBus starts a bus.Run in a goroutine and waits for its socket.
// Returns (workDir, sockPath, cancelFunc, runDone, fakeSource trigger).
func bringUpBus(t *testing.T, events []dwsevent.RawEvent) (string, string, context.CancelFunc, <-chan error, chan struct{}) {
t.Helper()
dir := shortTempDir(t)
sock := filepath.Join(dir, "bus.sock")
ctx, cancel := context.WithCancel(context.Background())
trigger := make(chan struct{})
src := &fakeSource{events: events, trigger: trigger}
done := make(chan error, 1)
go func() {
done <- bus.Run(ctx, bus.Config{
WorkDir: dir,
IPCEndpoint: sock,
ClientID: "ding_test",
Edition: "open",
Source: src,
})
}()
// Wait for socket.
deadline := time.Now().Add(2 * time.Second)
for time.Now().Before(deadline) {
if _, err := os.Stat(sock); err == nil {
return dir, sock, cancel, done, trigger
}
time.Sleep(10 * time.Millisecond)
}
cancel()
t.Fatalf("bus socket did not appear")
return "", "", nil, nil, nil
}
// dialOnlyDiscover is a Discover-impl-bypass: tests don't want consume.Run
// to exec a real dws binary, so we sidestep by pre-bringing-up the bus and
// letting Discover succeed on its first dial attempt. No Spawn is required.
func TestRun_StdoutNDJSON(t *testing.T) {
skipOnWindows(t)
dir, sock, cancel, runDone, trigger := bringUpBus(t, []dwsevent.RawEvent{
{EventID: "1", EventType: "im.message.receive_v1", Data: `{"text":"hi"}`},
{EventID: "2", EventType: "im.message.at_v1", Data: `{"at":1}`},
})
defer func() { cancel(); <-runDone }()
// Trigger source emission after we've started consume (otherwise events
// race ahead of consumer registration).
go func() {
time.Sleep(150 * time.Millisecond)
close(trigger)
}()
var stdout bytes.Buffer
var stderr bytes.Buffer
ctx, consumeCancel := context.WithTimeout(context.Background(), 3*time.Second)
defer consumeCancel()
err := Run(ctx, Config{
WorkDir: dir,
IPCEndpoint: sock,
ClientID: "ding_test",
Stdout: &stdout,
Stderr: &stderr,
EventTypes: []string{"im.*"},
MaxEvents: 2,
})
if err != nil {
t.Fatalf("Run: %v", err)
}
// Verify NDJSON: each non-empty line is a valid Event JSON.
lines := strings.Split(strings.TrimRight(stdout.String(), "\n"), "\n")
if len(lines) != 2 {
t.Fatalf("expected 2 NDJSON lines, got %d:\n%s", len(lines), stdout.String())
}
for i, line := range lines {
var ev transport.Event
if err := json.Unmarshal([]byte(line), &ev); err != nil {
t.Errorf("line %d not valid JSON: %v\n%s", i, err, line)
}
if ev.Type != transport.FrameTypeEvent {
t.Errorf("line %d type = %s, want event", i, ev.Type)
}
}
// Stderr should have the connected-to-bus banner.
if !strings.Contains(stderr.String(), "connected bus pid=") {
t.Errorf("stderr missing connected banner:\n%s", stderr.String())
}
}
func TestRun_QuietSuppressesStderr(t *testing.T) {
skipOnWindows(t)
dir, sock, cancel, runDone, trigger := bringUpBus(t, []dwsevent.RawEvent{
{EventID: "1", EventType: "im.message.receive_v1", Data: `{}`},
})
defer func() { cancel(); <-runDone }()
go func() { time.Sleep(150 * time.Millisecond); close(trigger) }()
var stdout, stderr bytes.Buffer
ctx, consumeCancel := context.WithTimeout(context.Background(), 3*time.Second)
defer consumeCancel()
err := Run(ctx, Config{
WorkDir: dir,
IPCEndpoint: sock,
ClientID: "ding_test",
Stdout: &stdout,
Stderr: &stderr,
Quiet: true,
MaxEvents: 1,
})
if err != nil {
t.Fatalf("Run: %v", err)
}
if stderr.Len() != 0 {
t.Errorf("--quiet should suppress all stderr; got: %s", stderr.String())
}
if stdout.Len() == 0 {
t.Error("stdout should still contain the NDJSON event")
}
}
func TestRun_CtxCancelReturnsCleanly(t *testing.T) {
skipOnWindows(t)
dir, sock, cancel, runDone, _ := bringUpBus(t, nil)
defer func() { cancel(); <-runDone }()
ctx, consumeCancel := context.WithCancel(context.Background())
consumeDone := make(chan error, 1)
go func() {
consumeDone <- Run(ctx, Config{
WorkDir: dir,
IPCEndpoint: sock,
ClientID: "ding_test",
Stdout: io.Discard,
Stderr: io.Discard,
})
}()
// Let consume connect.
time.Sleep(100 * time.Millisecond)
consumeCancel()
select {
case err := <-consumeDone:
if err != nil && !errors.Is(err, context.Canceled) {
t.Fatalf("Run returned %v, want nil or canceled", err)
}
case <-time.After(2 * time.Second):
t.Fatal("Run did not return after ctx cancel")
}
}
func TestRun_MaxEventsZeroIsUnlimited(t *testing.T) {
skipOnWindows(t)
dir, sock, cancel, runDone, trigger := bringUpBus(t, []dwsevent.RawEvent{
{EventID: "1", EventType: "x", Data: `{}`},
{EventID: "2", EventType: "x", Data: `{}`},
{EventID: "3", EventType: "x", Data: `{}`},
})
defer func() { cancel(); <-runDone }()
var stdout bytes.Buffer
ctx, consumeCancel := context.WithCancel(context.Background())
consumeDone := make(chan error, 1)
go func() {
consumeDone <- Run(ctx, Config{
WorkDir: dir,
IPCEndpoint: sock,
ClientID: "ding_test",
Stdout: &stdout,
Stderr: io.Discard,
MaxEvents: 0, // unlimited
})
}()
// Wait for consume to dial + Hello (otherwise events fire before
// consumer registers and the Hub drops them silently — no consumer
// to deliver to).
time.Sleep(150 * time.Millisecond)
close(trigger)
// Let all 3 events flow.
time.Sleep(200 * time.Millisecond)
consumeCancel()
<-consumeDone
lines := strings.Split(strings.TrimRight(stdout.String(), "\n"), "\n")
if len(lines) != 3 {
t.Fatalf("expected 3 events with MaxEvents=0, got %d:\n%s", len(lines), stdout.String())
}
}
// TestRun_MultipleConsumersOneBus exercises the daemon's multi-consumer
// fan-out via the real consume.Run path. Both consumers should receive
// every matching event independently.
func TestRun_MultipleConsumersOneBus(t *testing.T) {
skipOnWindows(t)
dir, sock, cancel, runDone, trigger := bringUpBus(t, []dwsevent.RawEvent{
{EventID: "1", EventType: "im.message.receive_v1", Data: `{}`},
{EventID: "2", EventType: "im.message.receive_v1", Data: `{}`},
})
defer func() { cancel(); <-runDone }()
var wg sync.WaitGroup
bufs := make([]*bytes.Buffer, 2)
for i := 0; i < 2; i++ {
i := i
bufs[i] = &bytes.Buffer{}
wg.Add(1)
go func() {
defer wg.Done()
_ = Run(context.Background(), Config{
WorkDir: dir,
IPCEndpoint: sock,
ClientID: "ding_test",
Stdout: bufs[i],
Stderr: io.Discard,
MaxEvents: 2,
})
}()
}
// Both consumers should be Hello'd before trigger.
time.Sleep(200 * time.Millisecond)
close(trigger)
wg.Wait()
for i, buf := range bufs {
lines := strings.Split(strings.TrimRight(buf.String(), "\n"), "\n")
if len(lines) != 2 {
t.Errorf("consumer %d: got %d lines, want 2:\n%s", i, len(lines), buf.String())
}
}
}
+30
View File
@@ -0,0 +1,30 @@
// 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.
//go:build !windows
package consume
import (
"errors"
"syscall"
)
// isBrokenPipe reports whether err originates from a closed downstream
// pipe (typical: `dws event consume | head -1`). On Unix this surfaces as
// EPIPE; the Go runtime by default also raises SIGPIPE which would kill
// the process, but Go programs ignore SIGPIPE on stdio writes (since
// Go 1.x). We just need to detect EPIPE and exit cleanly.
func isBrokenPipe(err error) bool {
return errors.Is(err, syscall.EPIPE)
}
+29
View File
@@ -0,0 +1,29 @@
// 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.
//go:build windows
package consume
import (
"errors"
"golang.org/x/sys/windows"
)
// isBrokenPipe reports whether err is the Windows equivalent of EPIPE
// (ERROR_BROKEN_PIPE / ERROR_NO_DATA) surfaced when a downstream pipe
// consumer closes its read end.
func isBrokenPipe(err error) bool {
return errors.Is(err, windows.ERROR_BROKEN_PIPE) || errors.Is(err, windows.ERROR_NO_DATA)
}
+170
View File
@@ -0,0 +1,170 @@
// 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 consume
import (
"errors"
"fmt"
"io"
"os"
"path/filepath"
"strings"
"time"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/transport"
)
// Sink is what a Pipeline writes formatted event bytes to. Implementations
// take both the event (for filename derivation in file sinks) and the
// already-formatted bytes (so the same event can be rendered with
// different formats per call without re-running the formatter inside the
// sink).
type Sink interface {
// Write places one event on the sink. Returns an error for IO failures;
// returns ErrPipeClosed when the downstream consumer closed the pipe
// (typical pattern: `dws event consume | head -1`).
Write(ev transport.Event, formatted []byte) error
// Close releases sink-owned resources. Idempotent.
Close() error
}
// ErrPipeClosed is returned by stdout-style sinks when the downstream
// reader closed its end (SIGPIPE / EPIPE on Unix, ERROR_BROKEN_PIPE on
// Windows). The pipeline catches this sentinel and exits cleanly without
// surfacing it as a fatal error.
var ErrPipeClosed = errors.New("sink: downstream pipe closed")
// NewStdoutSink wraps the given writer (typically os.Stdout) in a Sink
// that writes formatted bytes verbatim. Detects broken-pipe on the host
// platform and returns ErrPipeClosed so the caller can exit code 0.
func NewStdoutSink(w io.Writer) Sink { return &stdoutSink{w: w} }
type stdoutSink struct{ w io.Writer }
func (s *stdoutSink) Write(_ transport.Event, formatted []byte) error {
_, err := s.w.Write(formatted)
if err != nil && isBrokenPipe(err) {
return ErrPipeClosed
}
return err
}
func (s *stdoutSink) Close() error { return nil }
// NewFileDirSink returns a sink that writes each event to its own file
// under dir, naming files `{type}_{id}_{ts}.json`. The directory is
// mkdir'd on first write so callers don't have to ensure it themselves.
//
// Filename pieces are sanitised: characters that would escape the
// directory (path separators) or break shell globbing are replaced with
// '_'. `ts` is the ReceivedAtUnixMS (or current time if zero) so two
// events with the same id (re-delivery, dedup-defeated edge cases) don't
// collide.
func NewFileDirSink(dir string) Sink { return &fileDirSink{dir: dir} }
type fileDirSink struct{ dir string }
func (s *fileDirSink) Write(ev transport.Event, formatted []byte) error {
if err := os.MkdirAll(s.dir, 0o700); err != nil {
return fmt.Errorf("sink: mkdir %s: %w", s.dir, err)
}
name := buildFilename(ev)
full := filepath.Join(s.dir, name)
return atomicWrite(full, formatted)
}
func (s *fileDirSink) Close() error { return nil }
// NewRoutedSink composes a Router with per-route dir sinks plus a fallback.
// On each Write, Router.Match decides the target dir; if non-empty, the
// event is written there; otherwise the fallback sink handles it. The
// fallback is typically NewStdoutSink (default) or NewFileDirSink
// (--output-dir mode).
func NewRoutedSink(router *Router, fallback Sink) Sink {
return &routedSink{router: router, fallback: fallback}
}
type routedSink struct {
router *Router
fallback Sink
}
func (s *routedSink) Write(ev transport.Event, formatted []byte) error {
if dir := s.router.Match(ev); dir != "" {
return NewFileDirSink(dir).Write(ev, formatted)
}
return s.fallback.Write(ev, formatted)
}
func (s *routedSink) Close() error { return s.fallback.Close() }
// buildFilename produces `{type}_{id}_{ts}.json`. All three pieces are
// sanitised to be safe filesystem path segments — see safePart.
func buildFilename(ev transport.Event) string {
typ := safePart(ev.EventType)
if typ == "" {
typ = "unknown"
}
id := safePart(ev.EventID)
if id == "" {
id = "no-id"
}
ts := ev.ReceivedAtUnixMS
if ts == 0 {
ts = time.Now().UTC().UnixMilli()
}
return fmt.Sprintf("%s_%s_%d.json", typ, id, ts)
}
// safePart strips path separators, NULs, and leading/trailing whitespace
// from a filename piece. Replaces unsafe chars with '_' instead of
// dropping them so different inputs don't collide.
//
// We DO allow dots and dashes (common in event types like "im.message.at_v1");
// we just reject path separators and parent-directory traversal.
func safePart(s string) string {
s = strings.TrimSpace(s)
if s == "" {
return ""
}
var b strings.Builder
b.Grow(len(s))
for _, r := range s {
switch r {
case '/', '\\', 0, ':':
b.WriteRune('_')
default:
b.WriteRune(r)
}
}
return b.String()
// We intentionally do NOT collapse ".." sequences. After path
// separator replacement (above), a bare ".." in the middle of a
// filename cannot perform parent-directory traversal because there
// is no separator to anchor it against. The final filename is
// joined into a known-safe directory with filepath.Join which itself
// will clean any traversal that does sneak through.
}
// atomicWrite writes content to path via tmp-file + rename, so a concurrent
// reader either sees the previous version or the new version — never a
// half-written file.
func atomicWrite(path string, content []byte) error {
tmp := path + ".tmp"
if err := os.WriteFile(tmp, content, 0o600); err != nil {
return fmt.Errorf("sink: write tmp %s: %w", tmp, err)
}
if err := os.Rename(tmp, path); err != nil {
_ = os.Remove(tmp)
return fmt.Errorf("sink: rename to %s: %w", path, err)
}
return nil
}
+138
View File
@@ -0,0 +1,138 @@
// 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 consume
import (
"bytes"
"os"
"path/filepath"
"strings"
"testing"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/transport"
)
func TestStdoutSink_WritesVerbatim(t *testing.T) {
var buf bytes.Buffer
s := NewStdoutSink(&buf)
if err := s.Write(transport.Event{}, []byte("line one\n")); err != nil {
t.Fatal(err)
}
if err := s.Write(transport.Event{}, []byte("line two\n")); err != nil {
t.Fatal(err)
}
if buf.String() != "line one\nline two\n" {
t.Fatalf("got: %q", buf.String())
}
}
func TestFileDirSink_WritesPerEvent(t *testing.T) {
dir := t.TempDir()
s := NewFileDirSink(dir)
ev := transport.Event{
EventType: "im.message.receive_v1",
EventID: "ev_abc",
ReceivedAtUnixMS: 1700000000123,
}
body := []byte(`{"hello":"world"}`)
if err := s.Write(ev, body); err != nil {
t.Fatalf("Write: %v", err)
}
entries, _ := os.ReadDir(dir)
if len(entries) != 1 {
t.Fatalf("expected 1 file, got %d", len(entries))
}
got, _ := os.ReadFile(filepath.Join(dir, entries[0].Name()))
if !bytes.Equal(got, body) {
t.Fatalf("body mismatch: %q", got)
}
if entries[0].Name() != "im.message.receive_v1_ev_abc_1700000000123.json" {
t.Errorf("filename = %q", entries[0].Name())
}
}
func TestFileDirSink_MkdirAutomatic(t *testing.T) {
dir := filepath.Join(t.TempDir(), "nested", "events")
s := NewFileDirSink(dir)
err := s.Write(transport.Event{EventType: "x", EventID: "1", ReceivedAtUnixMS: 1}, []byte("ok"))
if err != nil {
t.Fatalf("Write should auto-mkdir: %v", err)
}
if _, err := os.Stat(dir); err != nil {
t.Fatalf("dir not created: %v", err)
}
}
func TestFileDirSink_AtomicWriteNoTmpLeft(t *testing.T) {
dir := t.TempDir()
s := NewFileDirSink(dir)
_ = s.Write(transport.Event{EventType: "x", EventID: "1", ReceivedAtUnixMS: 1}, []byte("ok"))
entries, _ := os.ReadDir(dir)
for _, e := range entries {
if strings.HasSuffix(e.Name(), ".tmp") {
t.Fatalf("tmp file leaked: %s", e.Name())
}
}
}
func TestSafePart_StripsPathSeparators(t *testing.T) {
cases := []struct{ in, want string }{
{"normal", "normal"},
{"im.message.receive_v1", "im.message.receive_v1"}, // dots OK
{"a/b", "a_b"},
{"a\\b", "a_b"},
{"a:b", "a_b"},
{"../../etc/passwd", ".._.._etc_passwd"}, // slashes → _; dots preserved (safe — no separator anchor)
{" spaces ", "spaces"},
{"", ""},
}
for _, c := range cases {
if got := safePart(c.in); got != c.want {
t.Errorf("safePart(%q) = %q, want %q", c.in, got, c.want)
}
}
}
func TestBuildFilename_DefaultsForMissingFields(t *testing.T) {
got := buildFilename(transport.Event{})
// Should contain unknown_no-id_<ts>.json
if !strings.HasPrefix(got, "unknown_no-id_") || !strings.HasSuffix(got, ".json") {
t.Fatalf("default filename shape unexpected: %q", got)
}
}
func TestRoutedSink_MatchedGoesToDir(t *testing.T) {
tmp := t.TempDir()
imDir := filepath.Join(tmp, "im")
var fallback bytes.Buffer
routes, _ := ParseRoutes([]string{`^im\.=dir:` + imDir})
rs := NewRoutedSink(NewRouter(routes), NewStdoutSink(&fallback))
// IM event → file in imDir
if err := rs.Write(transport.Event{EventType: "im.message.receive_v1", EventID: "x", ReceivedAtUnixMS: 1}, []byte("body")); err != nil {
t.Fatal(err)
}
if entries, _ := os.ReadDir(imDir); len(entries) != 1 {
t.Errorf("expected 1 file in %s, got %d", imDir, len(entries))
}
if fallback.Len() != 0 {
t.Errorf("fallback should be empty for matched route")
}
// Non-IM event → fallback (stdout)
_ = rs.Write(transport.Event{EventType: "approval.task", EventID: "y", ReceivedAtUnixMS: 2}, []byte("body2\n"))
if fallback.String() != "body2\n" {
t.Errorf("fallback got: %q", fallback.String())
}
}
+151
View File
@@ -0,0 +1,151 @@
// 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 consume
import (
"errors"
"fmt"
"io"
"strings"
)
// ValidationError represents a flag-level user error. It wraps a clear,
// actionable message — the cobra command layer surfaces it to the user
// with exit code 2 (validation error).
type ValidationError struct{ Msg string }
func (e *ValidationError) Error() string { return e.Msg }
// validation sentinels (cobra layer uses errors.Is to set the exit code).
var (
// ErrForceRequiresForeground is the plan §3.1 contract: --force only
// makes sense in foreground mode where the bus runs in the current
// process (then --force skips the single-instance lock so a second
// foreground bus can co-exist for a brief debug window). Outside of
// --foreground, --force would silently produce two daemons writing
// to the same socket — refuse upfront.
ErrForceRequiresForeground = &ValidationError{
Msg: "--force is only meaningful with --foreground (in daemon mode it would produce multiple bus instances; cloud events would be randomly split across connections). To restart the bus: dws event stop && dws event consume",
}
// ErrJSONFormatRequiresBounded is the plan §3.1 contract: --format
// json renders each event as a multi-line JSON object suitable for
// human inspection. With an unbounded stream the output mixes events
// without delimiters. Force the user to bound the run.
ErrJSONFormatRequiresBounded = &ValidationError{
Msg: "--format json requires --max-events or --duration (an unbounded JSON stream is not parseable). Use --format ndjson for unbounded streams.",
}
)
// ValidateConfig performs all pre-flight validation that does not require
// disk / network I/O. Returns a *ValidationError for any rule violation;
// returns nil if the cfg is launchable. The cobra layer calls this BEFORE
// calling Run so the user gets clear errors at parse time.
//
// Rules implemented:
// 1. WorkDir / IPCEndpoint / ClientID non-empty
// 2. --force requires --foreground (plan §3.1)
// 3. --format json requires --max-events OR --duration (bounded)
// 4. Routes already pre-parsed (any parse error is reported by ParseRoutes)
// 5. --output-dir conflict with global --output (caller-supplied flag —
// we expose ValidateNoOutputConflict separately because global -o is
// a cobra-layer concern)
//
// Rules NOT enforced here (deferred to caller / Run):
// - Credentials presence (auth.ResolveAppCredentialsStrict already
// reports a typed error)
// - bus availability (busctl.Discover handles)
func ValidateConfig(cfg Config) error {
if strings.TrimSpace(cfg.WorkDir) == "" {
return &ValidationError{Msg: "consume: WorkDir is required"}
}
if strings.TrimSpace(cfg.IPCEndpoint) == "" {
return &ValidationError{Msg: "consume: IPCEndpoint is required"}
}
if strings.TrimSpace(cfg.ClientID) == "" {
return &ValidationError{Msg: "consume: ClientID is required"}
}
if cfg.Force && !cfg.Foreground {
return ErrForceRequiresForeground
}
if cfg.Format == FormatJSON && cfg.MaxEvents <= 0 && cfg.Duration <= 0 {
return ErrJSONFormatRequiresBounded
}
return nil
}
// ValidateNoOutputConflict ensures --output-dir / --route (event-stream
// sinks) are not combined with the dws global hidden -o/--output flag
// (request-output to file). The cobra layer reads the global output flag
// from inherited flags and passes its value here; an empty globalOutput
// means the flag was unset.
func ValidateNoOutputConflict(cfg Config, globalOutput string) error {
if globalOutput == "" {
return nil
}
if cfg.OutputDir != "" || len(cfg.Routes) > 0 {
return &ValidationError{
Msg: fmt.Sprintf("--output-dir/--route cannot be combined with global -o/--output=%q (event stream sinks are mutually exclusive with single-file output capture)", globalOutput),
}
}
return nil
}
// IsValidationError reports whether err is a flag-level user error.
// Cobra command handlers use this to map validation errors to exit code 2.
func IsValidationError(err error) bool {
var v *ValidationError
return errors.As(err, &v)
}
// PrintDryRun writes the resolved configuration to w in a single
// human-readable block. Called by Run when cfg.DryRun is true. Format
// avoids JSON so users can `dws event consume --dry-run | head` cleanly.
//
// Secret-bearing fields are never present in Config (credentials never
// reach this layer), so no redaction is required here.
func PrintDryRun(w io.Writer, cfg Config) {
if w == nil {
return
}
fmt.Fprintln(w, "dws event consume — dry run (no bus connection will be made)")
fmt.Fprintf(w, " client_id : %s\n", cfg.ClientID)
fmt.Fprintf(w, " workdir : %s\n", cfg.WorkDir)
fmt.Fprintf(w, " ipc_endpoint : %s\n", cfg.IPCEndpoint)
if len(cfg.EventTypes) > 0 {
fmt.Fprintf(w, " event_types : %s\n", strings.Join(cfg.EventTypes, ","))
} else {
fmt.Fprintln(w, " event_types : (catch-all)")
}
if cfg.Filter != "" {
fmt.Fprintf(w, " filter : %s\n", cfg.Filter)
}
fmt.Fprintf(w, " format : %s\n", cfg.Format)
if cfg.OutputDir != "" {
fmt.Fprintf(w, " output_dir : %s\n", cfg.OutputDir)
}
for i, r := range cfg.Routes {
fmt.Fprintf(w, " route[%d] : %s\n", i, r.Raw)
}
if cfg.MaxEvents > 0 {
fmt.Fprintf(w, " max_events : %d\n", cfg.MaxEvents)
}
if cfg.Duration > 0 {
fmt.Fprintf(w, " duration : %s\n", cfg.Duration)
}
fmt.Fprintf(w, " compact : %v\n", cfg.Compact)
fmt.Fprintf(w, " quiet : %v\n", cfg.Quiet)
fmt.Fprintf(w, " foreground : %v\n", cfg.Foreground)
fmt.Fprintf(w, " force : %v\n", cfg.Force)
}
+166
View File
@@ -0,0 +1,166 @@
// 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 consume
import (
"bytes"
"errors"
"strings"
"testing"
"time"
)
func base() Config {
return Config{
WorkDir: "/tmp/x",
IPCEndpoint: "/tmp/x/bus.sock",
ClientID: "ding_abc",
}
}
func TestValidate_Happy(t *testing.T) {
if err := ValidateConfig(base()); err != nil {
t.Fatalf("baseline should be valid, got %v", err)
}
}
func TestValidate_RequiredFields(t *testing.T) {
cases := []struct {
name string
mut func(*Config)
}{
{"empty WorkDir", func(c *Config) { c.WorkDir = "" }},
{"whitespace WorkDir", func(c *Config) { c.WorkDir = " " }},
{"empty IPCEndpoint", func(c *Config) { c.IPCEndpoint = "" }},
{"empty ClientID", func(c *Config) { c.ClientID = "" }},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
c := base()
tc.mut(&c)
err := ValidateConfig(c)
if !IsValidationError(err) {
t.Fatalf("expected ValidationError, got %v", err)
}
})
}
}
func TestValidate_ForceRequiresForeground(t *testing.T) {
c := base()
c.Force = true
c.Foreground = false
err := ValidateConfig(c)
if !errors.Is(err, ErrForceRequiresForeground) {
t.Fatalf("err = %v, want ErrForceRequiresForeground", err)
}
if !strings.Contains(err.Error(), "event stop && dws event consume") {
t.Errorf("error message must include the recovery hint, got: %s", err.Error())
}
// With --foreground it's fine.
c.Foreground = true
if err := ValidateConfig(c); err != nil {
t.Fatalf("--force + --foreground should be valid, got %v", err)
}
}
func TestValidate_FormatJSONRequiresBounded(t *testing.T) {
c := base()
c.Format = FormatJSON
c.MaxEvents = 0
c.Duration = 0
err := ValidateConfig(c)
if !errors.Is(err, ErrJSONFormatRequiresBounded) {
t.Fatalf("err = %v, want ErrJSONFormatRequiresBounded", err)
}
// With --max-events it passes.
c.MaxEvents = 10
if err := ValidateConfig(c); err != nil {
t.Fatalf("--format json + --max-events should be valid: %v", err)
}
c.MaxEvents = 0
c.Duration = 30 * time.Second
if err := ValidateConfig(c); err != nil {
t.Fatalf("--format json + --duration should be valid: %v", err)
}
// NDJSON has no such requirement.
c.Format = FormatNDJSON
c.MaxEvents = 0
c.Duration = 0
if err := ValidateConfig(c); err != nil {
t.Fatalf("ndjson unbounded should be valid: %v", err)
}
}
func TestValidateNoOutputConflict(t *testing.T) {
c := base()
c.OutputDir = "/tmp/events"
if err := ValidateNoOutputConflict(c, ""); err != nil {
t.Fatalf("no global -o → ok, got %v", err)
}
if err := ValidateNoOutputConflict(c, "/tmp/out.json"); !IsValidationError(err) {
t.Fatalf("--output-dir + global -o should be ValidationError, got %v", err)
}
c2 := base()
c2.Routes, _ = ParseRoutes([]string{`^im=dir:./im/`})
if err := ValidateNoOutputConflict(c2, "/tmp/out.json"); !IsValidationError(err) {
t.Fatalf("--route + global -o should be ValidationError, got %v", err)
}
}
func TestPrintDryRun_NilWriterSafe(t *testing.T) {
// Must not panic.
PrintDryRun(nil, base())
}
func TestPrintDryRun_RendersAllSetFields(t *testing.T) {
var buf bytes.Buffer
c := base()
c.EventTypes = []string{"im.*", "approval.*"}
c.Filter = "^im\\."
c.Format = FormatCompact
c.OutputDir = "/tmp/events"
c.Routes, _ = ParseRoutes([]string{`^im\.=dir:/tmp/im/`})
c.MaxEvents = 5
c.Duration = 30 * time.Second
c.Compact = true
c.Quiet = true
c.Foreground = true
c.Force = true
PrintDryRun(&buf, c)
out := buf.String()
wants := []string{
"client_id", "workdir", "ipc_endpoint", "im.*,approval.*",
"^im\\.", "compact", "/tmp/events", "route[0]", "max_events : 5",
"duration", "true",
}
for _, w := range wants {
if !strings.Contains(out, w) {
t.Errorf("dry-run missing %q in output:\n%s", w, out)
}
}
}
func TestPrintDryRun_CatchAllWhenEventTypesEmpty(t *testing.T) {
var buf bytes.Buffer
PrintDryRun(&buf, base())
if !strings.Contains(buf.String(), "(catch-all)") {
t.Errorf("expected '(catch-all)' for empty event_types:\n%s", buf.String())
}
}
+94
View File
@@ -0,0 +1,94 @@
// 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 dedup implements a fixed-capacity LRU set used by the bus to
// suppress duplicate events that the DingTalk Stream SDK redelivers on
// reconnect (see plan invariant #2). The set stores only keys, never values,
// and is safe for concurrent use by multiple goroutines.
package dedup
import (
"container/list"
"sync"
)
// DefaultCapacity is the default LRU size when LRU is constructed without an
// explicit capacity. 8192 is sized to absorb a typical reconnect-storm window
// (~5 min × 30 events/s) while staying memory-cheap (~256 KB at 32 bytes/key).
const DefaultCapacity = 8192
// LRU is a fixed-capacity LRU set of string keys. Zero value is not usable;
// call New or NewWithCapacity.
type LRU struct {
mu sync.Mutex
cap int
keys map[string]*list.Element
eviction *list.List // back = newest, front = oldest
}
// New returns an LRU with DefaultCapacity.
func New() *LRU { return NewWithCapacity(DefaultCapacity) }
// NewWithCapacity returns an LRU sized to hold up to cap keys. cap must be > 0.
func NewWithCapacity(cap int) *LRU {
if cap <= 0 {
cap = DefaultCapacity
}
return &LRU{
cap: cap,
keys: make(map[string]*list.Element, cap),
eviction: list.New(),
}
}
// Seen reports whether key was already present and inserts it if not. The
// return value is true when the caller should treat the event as a duplicate
// (drop it) and false when this is the first occurrence.
//
// Empty keys are never considered duplicates and are not stored — callers
// without a stable identifier should use RawEvent.DedupKey() which falls back
// to a content hash.
func (l *LRU) Seen(key string) bool {
if key == "" {
return false
}
l.mu.Lock()
defer l.mu.Unlock()
if el, ok := l.keys[key]; ok {
l.eviction.MoveToBack(el)
return true
}
if len(l.keys) >= l.cap {
oldest := l.eviction.Front()
if oldest != nil {
delete(l.keys, oldest.Value.(string))
l.eviction.Remove(oldest)
}
}
el := l.eviction.PushBack(key)
l.keys[key] = el
return false
}
// Len returns the current number of stored keys.
func (l *LRU) Len() int {
l.mu.Lock()
defer l.mu.Unlock()
return len(l.keys)
}
// Cap returns the configured capacity.
func (l *LRU) Cap() int { return l.cap }
+111
View File
@@ -0,0 +1,111 @@
// 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 dedup
import (
"strconv"
"sync"
"testing"
)
func TestLRU_FirstSeenIsFalse(t *testing.T) {
l := New()
if l.Seen("a") {
t.Fatal("first occurrence should not be reported as seen")
}
if !l.Seen("a") {
t.Fatal("second occurrence should be reported as seen")
}
}
func TestLRU_EmptyKeyNeverSeen(t *testing.T) {
l := New()
if l.Seen("") {
t.Fatal("empty key must never be reported as seen")
}
if l.Seen("") {
t.Fatal("empty key must never be stored / reported as seen on second call either")
}
if l.Len() != 0 {
t.Fatalf("empty key must not be stored, got Len=%d", l.Len())
}
}
func TestLRU_EvictsOldestAtCapacity(t *testing.T) {
l := NewWithCapacity(3)
for _, k := range []string{"a", "b", "c"} {
if l.Seen(k) {
t.Fatalf("unexpected seen for %s", k)
}
}
// d inserts → a should be evicted
if l.Seen("d") {
t.Fatal("d is new, should not be seen")
}
if l.Seen("a") {
t.Fatal("a should have been evicted; second insert returns not-seen")
}
// Now b should be the oldest. After re-querying "c" (refreshes c),
// inserting "e" should evict b not c.
if !l.Seen("c") {
t.Fatal("c is still in set, should be seen")
}
if l.Seen("e") {
t.Fatal("e is new")
}
if l.Seen("b") {
t.Fatal("b should have been evicted by e (c was just refreshed)")
}
}
func TestLRU_LenAndCap(t *testing.T) {
l := NewWithCapacity(5)
if l.Cap() != 5 {
t.Fatalf("Cap = %d, want 5", l.Cap())
}
if l.Len() != 0 {
t.Fatalf("initial Len = %d, want 0", l.Len())
}
_ = l.Seen("x")
_ = l.Seen("y")
if l.Len() != 2 {
t.Fatalf("after 2 inserts Len = %d, want 2", l.Len())
}
}
func TestLRU_ZeroCapacityUsesDefault(t *testing.T) {
l := NewWithCapacity(0)
if l.Cap() != DefaultCapacity {
t.Fatalf("zero cap should use DefaultCapacity, got %d", l.Cap())
}
}
func TestLRU_ConcurrentSafety(t *testing.T) {
l := NewWithCapacity(1000)
const N = 200
var wg sync.WaitGroup
wg.Add(N)
for i := 0; i < N; i++ {
i := i
go func() {
defer wg.Done()
_ = l.Seen(strconv.Itoa(i))
_ = l.Seen(strconv.Itoa(i)) // duplicate
}()
}
wg.Wait()
if l.Len() != N {
t.Fatalf("Len after concurrent insert = %d, want %d", l.Len(), N)
}
}
+36
View File
@@ -0,0 +1,36 @@
// 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 event implements the DingTalk Stream event subscription pipeline for
// dws. The architecture is a single-cloud-connection bus daemon (one per
// ClientID) plus N local consumer processes communicating over Unix socket /
// Windows Named Pipe. The bus keeps one cloud connection per identity while
// exposing observable connection state, per-event-type metrics, and Hello-time
// filter pushdown to local consumers.
//
// Package layout:
//
// event/ // top-level types (RawEvent, EmitFn, hash helpers)
// event/dedup/ // event_id LRU dedup
// event/registry/ // catch-all event types + compact processor registry
// event/source/ // wrap dingtalk-stream-sdk-go + connection state machine
// event/bus/ // daemon loop, hub, metrics, lockfile, meta
// event/transport/ // UDS/Pipe abstraction, frame protocol
// event/busctl/ // discover, spawn, stop helpers
// event/consume/ // consumer-side pipeline, formatter, router, sink
// event/lock/ // cross-platform flock primitive (Unix flock / Windows LockFileEx)
// event/process/ // cross-platform process-alive check (Unix signal 0 / Windows OpenProcess)
//
// See plans/2026-05-28_event_capability_v1.plan.md for the full design,
// invariants, and protocol.
package event
+60
View File
@@ -0,0 +1,60 @@
// 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 event
import (
"os"
"path/filepath"
"runtime"
)
// MaxUnixSocketPath returns the longest Unix socket path accepted by
// bind/connect on this OS (Go rejects longer names with EINVAL before
// the syscall). sockaddr_un.sun_path is 104 bytes on darwin and the
// BSDs and 108 on Linux; the usable budget is one less.
func MaxUnixSocketPath() int {
if runtime.GOOS == "linux" {
return 107
}
return 103
}
// IPCEndpoint returns the bus IPC endpoint for one identity: a Named Pipe
// name on Windows, otherwise bus.sock inside workDir.
//
// The canonical Unix location is <workDir>/bus.sock, but workDir derives
// from the config dir, which can be arbitrarily deep (e.g. dwssb sandboxes
// use ~/.dwssb/sandboxes/<name>/config/...). When the canonical path would
// exceed the OS sun_path limit, the socket falls back to a short
// deterministic path under os.TempDir keyed by a hash of workDir, so every
// process (consume parent, forked _bus child, status/stop tooling) that
// derives the endpoint from the same workDir agrees on the location.
// bus.lock / bus.meta / bus.log always stay in workDir — only the socket
// moves.
//
// This is the single source of truth for endpoint derivation; the cobra
// layer and busctl must not re-implement the shape.
func IPCEndpoint(workDir, editionName string, sourceKind SourceKind, identityHash string) string {
if sourceKind == "" {
sourceKind = SourceKindAppStream
}
if runtime.GOOS == "windows" {
return `\\.\pipe\dws-event-` + editionName + "-" + string(sourceKind) + "-" + identityHash
}
sock := filepath.Join(workDir, "bus.sock")
if len(sock) <= MaxUnixSocketPath() {
return sock
}
return filepath.Join(os.TempDir(), "dws-evt-"+IdentityHash(workDir)+".sock")
}
+61
View File
@@ -0,0 +1,61 @@
// 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.
//go:build !windows
package event
import (
"os"
"path/filepath"
"strings"
"testing"
)
func TestIPCEndpointShortWorkDirUsesCanonicalPath(t *testing.T) {
workDir := "/tmp/dws/events/open/app_stream/aabbccdd00112233"
got := IPCEndpoint(workDir, "open", SourceKindAppStream, "aabbccdd00112233")
want := filepath.Join(workDir, "bus.sock")
if got != want {
t.Fatalf("IPCEndpoint = %q, want %q", got, want)
}
}
func TestIPCEndpointLongWorkDirFallsBackUnderTempDir(t *testing.T) {
// Mirrors the dwssb sandbox layout that produced a 111-byte socket
// path — over macOS's 103-byte usable sun_path budget.
workDir := "/Users/zhengyubai/.dwssb/sandboxes/event-subscribe/config/events/open/personal_stream/3928ce0fb4860a52"
got := IPCEndpoint(workDir, "open", SourceKindPersonalStream, "3928ce0fb4860a52")
if strings.HasPrefix(got, workDir) {
t.Fatalf("IPCEndpoint = %q, want fallback outside workDir", got)
}
if !strings.HasPrefix(got, os.TempDir()) {
t.Fatalf("IPCEndpoint = %q, want fallback under os.TempDir %q", got, os.TempDir())
}
if len(got) > MaxUnixSocketPath() {
t.Fatalf("fallback path still too long: %d > %d (%q)", len(got), MaxUnixSocketPath(), got)
}
}
func TestIPCEndpointFallbackIsDeterministicPerWorkDir(t *testing.T) {
long := strings.Repeat("x", 120)
a := IPCEndpoint("/base/"+long+"/one", "open", SourceKindPersonalStream, "hash")
b := IPCEndpoint("/base/"+long+"/one", "open", SourceKindPersonalStream, "hash")
c := IPCEndpoint("/base/"+long+"/two", "open", SourceKindPersonalStream, "hash")
if a != b {
t.Fatalf("same workDir produced different endpoints: %q vs %q", a, b)
}
if a == c {
t.Fatalf("different workDirs collided on endpoint %q", a)
}
}
+24
View File
@@ -0,0 +1,24 @@
// 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 lock implements a cross-platform exclusive file lock primitive used
// by the bus daemon to enforce the "single bus per ClientID" invariant
// (plan invariant #3). The primitive is intentionally tiny — it acquires and
// releases a non-blocking exclusive lock on an opened file handle, with no
// knowledge of PID files or business semantics. Higher layers (bus/lockfile.go)
// combine this primitive with PID content read/write to provide the full
// single-file bus.lock design.
//
// Unix: syscall.Flock(LOCK_EX|LOCK_NB).
// Windows: windows.LockFileEx with LOCKFILE_EXCLUSIVE_LOCK | LOCKFILE_FAIL_IMMEDIATELY.
package lock
+79
View File
@@ -0,0 +1,79 @@
// 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 lock
import (
"errors"
"fmt"
"os"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
)
// ErrBusy indicates the lock is currently held by another process. Callers
// distinguish "actually busy" (another live bus) from "lock file
// inaccessible" (FS error) by checking for this sentinel.
var ErrBusy = errors.New("lock: file is held by another process")
// File represents a held exclusive file lock. Close releases the lock and
// closes the underlying file handle. A zero File is not usable.
type File struct {
f *os.File
}
// TryAcquire opens path (creating it if absent, mode 0600) and attempts to
// take an exclusive non-blocking lock. Returns ErrBusy when the lock is held
// by another process; any other error wraps the underlying I/O failure.
//
// The directory containing path must already exist; callers should mkdir
// with pkg/config.DirPerm beforehand.
func TryAcquire(path string) (*File, error) {
f, err := os.OpenFile(path, os.O_RDWR|os.O_CREATE, config.FilePerm)
if err != nil {
return nil, fmt.Errorf("lock: open %s: %w", path, err)
}
if err := lockFile(f); err != nil {
_ = f.Close()
if isBusy(err) {
return nil, ErrBusy
}
return nil, fmt.Errorf("lock: flock %s: %w", path, err)
}
return &File{f: f}, nil
}
// File returns the underlying *os.File so callers can Read/Write content
// while holding the lock. The handle MUST NOT be closed by the caller —
// use Close on the lock File instead.
func (l *File) File() *os.File { return l.f }
// Path returns the file path the lock is held on.
func (l *File) Path() string {
if l == nil || l.f == nil {
return ""
}
return l.f.Name()
}
// Close releases the lock and closes the file handle. Safe to call on a nil
// receiver. Subsequent calls are no-ops.
func (l *File) Close() error {
if l == nil || l.f == nil {
return nil
}
unlockFile(l.f)
err := l.f.Close()
l.f = nil
return err
}
+101
View File
@@ -0,0 +1,101 @@
// 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 lock
import (
"errors"
"io"
"path/filepath"
"testing"
)
func TestTryAcquire_FirstCallerWins(t *testing.T) {
path := filepath.Join(t.TempDir(), "bus.lock")
l, err := TryAcquire(path)
if err != nil {
t.Fatalf("first TryAcquire: %v", err)
}
defer l.Close()
if l.Path() != path {
t.Fatalf("Path() = %q, want %q", l.Path(), path)
}
}
func TestTryAcquire_SecondCallerGetsBusy(t *testing.T) {
path := filepath.Join(t.TempDir(), "bus.lock")
l1, err := TryAcquire(path)
if err != nil {
t.Fatalf("first TryAcquire: %v", err)
}
defer l1.Close()
l2, err := TryAcquire(path)
if !errors.Is(err, ErrBusy) {
t.Fatalf("second TryAcquire: err = %v, want ErrBusy", err)
}
if l2 != nil {
t.Fatal("on ErrBusy the returned lock must be nil")
}
}
func TestTryAcquire_ReleasedLockIsReacquirable(t *testing.T) {
path := filepath.Join(t.TempDir(), "bus.lock")
l1, err := TryAcquire(path)
if err != nil {
t.Fatalf("first TryAcquire: %v", err)
}
if err := l1.Close(); err != nil {
t.Fatalf("close: %v", err)
}
l2, err := TryAcquire(path)
if err != nil {
t.Fatalf("re-acquire after close: %v", err)
}
defer l2.Close()
}
func TestTryAcquire_ContentReadWriteWhileHeld(t *testing.T) {
path := filepath.Join(t.TempDir(), "bus.lock")
l, err := TryAcquire(path)
if err != nil {
t.Fatalf("acquire: %v", err)
}
defer l.Close()
// Write PID-like content through the underlying handle
const pid = "12345\n"
if _, err := l.File().WriteString(pid); err != nil {
t.Fatalf("write: %v", err)
}
// Rewind and read back
if _, err := l.File().Seek(0, io.SeekStart); err != nil {
t.Fatalf("seek: %v", err)
}
buf := make([]byte, len(pid))
if _, err := io.ReadFull(l.File(), buf); err != nil {
t.Fatalf("read: %v", err)
}
if string(buf) != pid {
t.Fatalf("read back = %q, want %q", buf, pid)
}
}
func TestClose_NilSafe(t *testing.T) {
var l *File
if err := l.Close(); err != nil {
t.Fatalf("nil Close should be no-op, got %v", err)
}
}
+35
View File
@@ -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.
//go:build !windows
package lock
import (
"errors"
"os"
"syscall"
)
func lockFile(f *os.File) error {
return syscall.Flock(int(f.Fd()), syscall.LOCK_EX|syscall.LOCK_NB)
}
func unlockFile(f *os.File) {
_ = syscall.Flock(int(f.Fd()), syscall.LOCK_UN)
}
func isBusy(err error) bool {
// Linux returns EWOULDBLOCK (==EAGAIN); macOS returns EWOULDBLOCK too.
return errors.Is(err, syscall.EWOULDBLOCK) || errors.Is(err, syscall.EAGAIN)
}
+60
View File
@@ -0,0 +1,60 @@
// 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.
//go:build windows
package lock
import (
"errors"
"os"
"unsafe"
"golang.org/x/sys/windows"
)
const (
lockfileExclusiveLock = 0x00000002
lockfileFailImmediately = 0x00000001
)
func lockFile(f *os.File) error {
handle := windows.Handle(f.Fd())
ol := new(windows.Overlapped)
return windows.LockFileEx(
handle,
lockfileExclusiveLock|lockfileFailImmediately,
0,
1,
0,
(*windows.Overlapped)(unsafe.Pointer(ol)),
)
}
func unlockFile(f *os.File) {
handle := windows.Handle(f.Fd())
ol := new(windows.Overlapped)
_ = windows.UnlockFileEx(
handle,
0,
1,
0,
(*windows.Overlapped)(unsafe.Pointer(ol)),
)
}
func isBusy(err error) bool {
// LockFileEx with LOCKFILE_FAIL_IMMEDIATELY returns ERROR_LOCK_VIOLATION
// when the region is already locked.
return errors.Is(err, windows.ERROR_LOCK_VIOLATION) || errors.Is(err, windows.ERROR_IO_PENDING)
}
+681
View File
@@ -0,0 +1,681 @@
// 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 personal
import (
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"io"
"log/slog"
"net/http"
"net/url"
"strings"
"time"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
)
const DefaultBasePath = "/dws"
const (
controlLogPayloadLimit = 8192
subscriptionListPageSize = 100
subscriptionListMaxPageGuard = 10000
)
type Identity struct {
AccessToken string `json:"-"`
LocalSubject string `json:"-"`
CorpID string `json:"corp_id"`
UserID string `json:"user_id"`
ClientID string `json:"client_id"`
SourceID string `json:"source_id"`
}
func (i Identity) Key() string {
corpID := strings.TrimSpace(i.CorpID)
userID := strings.TrimSpace(i.UserID)
clientID := strings.TrimSpace(i.ClientID)
sourceID := strings.TrimSpace(i.SourceID)
if corpID != "" && userID != "" {
return strings.Join([]string{"corp_user", corpID, userID, clientID, sourceID}, "\x00")
}
if localSubject := strings.TrimSpace(i.LocalSubject); localSubject != "" {
return strings.Join([]string{"local_subject", localSubject, clientID, sourceID}, "\x00")
}
return strings.Join([]string{"unknown", corpID, userID, clientID, sourceID}, "\x00")
}
type Client struct {
BaseURL string
HTTPClient *http.Client
Identity Identity
}
type CreateSubscriptionRequest struct {
EventKey string `json:"event_key"`
RuleType string `json:"rule_type"`
Name string `json:"name,omitempty"`
RuleParam map[string]any `json:"rule_param"`
Filter any `json:"filter,omitempty"`
Delivery map[string]any `json:"delivery"`
TTLSeconds int64 `json:"ttl_seconds,omitempty"`
IdempotencyKey string `json:"idempotency_key,omitempty"`
}
type Subscription struct {
SubscribeID string `json:"subscribe_id"`
EventKey string `json:"event_key,omitempty"`
RuleType string `json:"rule_type,omitempty"`
Status string `json:"status,omitempty"`
SourceID string `json:"source_id,omitempty"`
CreatedAt string `json:"created_at,omitempty"`
ExpiresAt *time.Time `json:"expires_at,omitempty"`
}
type ListOptions struct {
Status string
EventKey string
SubscribeID string
}
type dwsCreateSubscriptionRequest struct {
ClientID string `json:"clientId"`
SourceID string `json:"sourceId,omitempty"`
EventKey string `json:"eventKey"`
FilterRule string `json:"filterRule,omitempty"`
DeliveryPref string `json:"deliveryPref,omitempty"`
ExpiresAt string `json:"expiresAt,omitempty"`
Ext map[string]any `json:"ext,omitempty"`
}
type dwsSubListResult struct {
Total int `json:"total,omitempty"`
PageNo int `json:"pageNo,omitempty"`
PageSize int `json:"pageSize,omitempty"`
Items []dwsSubscription `json:"items"`
}
type dwsSubscription struct {
SubID string `json:"subId"`
SubscribeID string `json:"subscribe_id"`
EventKey string `json:"eventKey"`
EventKeySnake string `json:"event_key"`
RuleType string `json:"ruleType,omitempty"`
RuleTypeSnake string `json:"rule_type,omitempty"`
ClientID string `json:"clientId,omitempty"`
SourceID string `json:"sourceId"`
SourceIDSnake string `json:"source_id"`
DeliveryPref string `json:"deliveryPref,omitempty"`
Status json.RawMessage `json:"status,omitempty"`
GmtCreate string `json:"gmtCreate,omitempty"`
CreatedAt string `json:"created_at,omitempty"`
}
func (s dwsSubscription) toSubscription() Subscription {
return Subscription{
SubscribeID: firstNonEmpty(s.SubID, s.SubscribeID),
EventKey: firstNonEmpty(s.EventKey, s.EventKeySnake),
RuleType: firstNonEmpty(s.RuleType, s.RuleTypeSnake),
Status: dwsStatusString(s.Status),
SourceID: firstNonEmpty(s.SourceID, s.SourceIDSnake),
CreatedAt: firstNonEmpty(s.GmtCreate, s.CreatedAt),
}
}
type APIError struct {
Code string `json:"code"`
Message string `json:"message"`
Retryable bool `json:"retryable,omitempty"`
Details map[string]any `json:"details,omitempty"`
}
func (e *APIError) Error() string {
if e == nil {
return ""
}
if e.Code != "" && e.Message != "" {
return e.Code + ": " + e.Message
}
if e.Code != "" {
return e.Code
}
return e.Message
}
func NewClient(baseURL string, identity Identity) *Client {
baseURL = strings.TrimRight(strings.TrimSpace(baseURL), "/")
if baseURL == "" {
baseURL = strings.TrimRight(config.GetMCPBaseURL(), "/") + DefaultBasePath
}
return &Client{
BaseURL: baseURL,
HTTPClient: &http.Client{Timeout: 30 * time.Second},
Identity: identity,
}
}
func (c *Client) CreateSubscription(ctx context.Context, req CreateSubscriptionRequest) (*Subscription, error) {
if req.EventKey == "" || req.RuleType == "" {
return nil, errors.New("personal event: event_key and rule_type are required")
}
var sub Subscription
if err := c.do(ctx, http.MethodPost, "/subscription/user", nil, c.buildCreateRequest(req), &sub); err != nil {
var apiErr *APIError
if errors.As(err, &apiErr) {
if subID, ok := apiErr.Details["subscribe_id"].(string); ok && subID != "" {
return &Subscription{
SubscribeID: subID,
EventKey: req.EventKey,
RuleType: req.RuleType,
Status: "active",
SourceID: c.Identity.SourceID,
}, nil
}
}
return nil, err
}
if sub.EventKey == "" {
sub.EventKey = req.EventKey
}
if sub.RuleType == "" {
sub.RuleType = req.RuleType
}
if sub.Status == "" {
sub.Status = "active"
}
if sub.SourceID == "" {
sub.SourceID = c.Identity.SourceID
}
return &sub, nil
}
func (c *Client) GetSubscription(ctx context.Context, subscribeID string) (*Subscription, error) {
subscribeID = strings.TrimSpace(subscribeID)
if subscribeID == "" {
return nil, errors.New("personal event: subscribe_id is required")
}
subs, err := c.ListSubscriptions(ctx, ListOptions{SubscribeID: subscribeID})
if err != nil {
return nil, err
}
if len(subs) == 0 {
return nil, &APIError{Code: "PERSONAL_EVENT_NOT_FOUND", Message: "subscription not found"}
}
return &subs[0], nil
}
func (c *Client) ListSubscriptions(ctx context.Context, opts ListOptions) ([]Subscription, error) {
q := make(url.Values)
if clientID := strings.TrimSpace(c.Identity.ClientID); clientID != "" {
q.Set("clientId", clientID)
}
if sourceID := strings.TrimSpace(c.Identity.SourceID); sourceID != "" {
q.Set("sourceId", sourceID)
}
q.Set("pageSize", fmt.Sprintf("%d", subscriptionListPageSize))
all := make([]Subscription, 0, subscriptionListPageSize)
seen := make(map[string]struct{}, subscriptionListPageSize)
for pageNo := 1; pageNo <= subscriptionListMaxPageGuard; pageNo++ {
q.Set("pageNo", fmt.Sprintf("%d", pageNo))
var result dwsSubListResult
if err := c.do(ctx, http.MethodGet, "/event/sublist", q, nil, &result); err != nil {
return nil, err
}
if len(result.Items) == 0 {
break
}
effectivePageSize := subscriptionListPageSize
if result.PageSize > 0 {
effectivePageSize = result.PageSize
}
added := 0
for _, item := range result.Items {
sub := item.toSubscription()
if sub.SubscribeID != "" {
if _, ok := seen[sub.SubscribeID]; ok {
continue
}
seen[sub.SubscribeID] = struct{}{}
}
all = append(all, sub)
added++
}
if added == 0 && (result.Total > len(all) || len(result.Items) >= effectivePageSize) {
return nil, fmt.Errorf("personal event: subscription pagination made no progress at page %d", pageNo)
}
if result.Total > 0 && len(all) >= result.Total {
break
}
if len(result.Items) < effectivePageSize {
break
}
if pageNo == subscriptionListMaxPageGuard {
return nil, fmt.Errorf("personal event: subscription pagination exceeded %d pages", subscriptionListMaxPageGuard)
}
}
items := make([]Subscription, 0, len(all))
for _, sub := range all {
if opts.Status != "" && opts.Status != "all" && sub.Status != opts.Status {
continue
}
if opts.EventKey != "" && sub.EventKey != opts.EventKey {
continue
}
if opts.SubscribeID != "" && sub.SubscribeID != opts.SubscribeID {
continue
}
items = append(items, sub)
}
return items, nil
}
func (c *Client) DeleteSubscription(ctx context.Context, subscribeID string) error {
subscribeID = strings.TrimSpace(subscribeID)
if subscribeID == "" {
return errors.New("personal event: subscribe_id is required")
}
err := c.do(ctx, http.MethodPost, "/subscription/cancel", nil, map[string]string{"subId": subscribeID}, nil)
if isNotFound(err) {
return nil
}
return err
}
func (c *Client) buildCreateRequest(req CreateSubscriptionRequest) dwsCreateSubscriptionRequest {
filterRule := ""
if req.RuleParam != nil {
if b, err := json.Marshal(req.RuleParam); err == nil {
filterRule = string(b)
}
}
ext := map[string]any{
"ruleType": req.RuleType,
}
if req.Name != "" {
ext["name"] = req.Name
}
if req.Filter != nil {
ext["filter"] = req.Filter
}
if req.IdempotencyKey != "" {
ext["idempotencyKey"] = req.IdempotencyKey
}
out := dwsCreateSubscriptionRequest{
ClientID: c.Identity.ClientID,
SourceID: c.Identity.SourceID,
EventKey: req.EventKey,
FilterRule: filterRule,
DeliveryPref: "realtime",
Ext: ext,
}
if req.TTLSeconds > 0 {
out.ExpiresAt = time.Now().UTC().Add(time.Duration(req.TTLSeconds) * time.Second).Format(time.RFC3339)
}
return out
}
func (c *Client) do(ctx context.Context, method, path string, q url.Values, body any, out any) error {
if c == nil {
return errors.New("personal event: nil client")
}
if c.Identity.AccessToken == "" {
return errors.New("personal event: access token is required")
}
u := strings.TrimRight(c.BaseURL, "/") + path
if len(q) > 0 {
u += "?" + q.Encode()
}
var r io.Reader
requestLog := ""
if body != nil {
b, err := json.Marshal(body)
if err != nil {
return fmt.Errorf("personal event: encode request: %w", err)
}
requestLog = sanitizeLogPayload(b)
r = bytes.NewReader(b)
}
req, err := http.NewRequestWithContext(ctx, method, u, r)
if err != nil {
return fmt.Errorf("personal event: create request: %w", err)
}
c.decorate(req)
if body != nil {
req.Header.Set("Content-Type", "application/json")
}
hc := c.HTTPClient
if hc == nil {
hc = http.DefaultClient
}
resp, err := hc.Do(req)
if err != nil {
return fmt.Errorf("personal event: send request: %w", err)
}
defer resp.Body.Close()
data, err := io.ReadAll(io.LimitReader(resp.Body, config.MaxResponseBodySize))
if err != nil {
return fmt.Errorf("personal event: read response: %w", err)
}
responseLog := sanitizeLogPayload(data)
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
if apiErr := decodeAPIError(data); apiErr != nil {
apiErr = withRequestDetails(apiErr, method, path, resp.StatusCode, responseRequestID(data))
logControlRequest("personal event control request failed", method, path, q, resp.StatusCode, requestLog, responseLog, responseRequestID(data), apiErr)
return apiErr
}
logControlRequest("personal event control request failed", method, path, q, resp.StatusCode, requestLog, responseLog, responseRequestID(data), nil)
return fmt.Errorf("personal event: HTTP %d", resp.StatusCode)
}
if len(bytes.TrimSpace(data)) == 0 {
logControlRequest("personal event control request", method, path, q, resp.StatusCode, requestLog, responseLog, "", nil)
return nil
}
var env responseEnvelope
if err := json.Unmarshal(data, &env); err == nil && (env.Success != nil || env.Error != nil || env.Result != nil || env.ErrorCode != "" || env.ErrorMsg != "") {
if env.Success == nil {
if apiErr := env.apiError(); apiErr != nil {
apiErr = withRequestDetails(apiErr, method, path, resp.StatusCode, env.requestID())
logControlRequest("personal event control request failed", method, path, q, resp.StatusCode, requestLog, responseLog, env.requestID(), apiErr)
return apiErr
}
}
if env.Success != nil && !*env.Success {
if apiErr := env.apiError(); apiErr != nil {
apiErr = withRequestDetails(apiErr, method, path, resp.StatusCode, env.requestID())
logControlRequest("personal event control request failed", method, path, q, resp.StatusCode, requestLog, responseLog, env.requestID(), apiErr)
return apiErr
}
logControlRequest("personal event control request failed", method, path, q, resp.StatusCode, requestLog, responseLog, env.requestID(), nil)
return errors.New("personal event: request failed")
}
logControlRequest("personal event control request", method, path, q, resp.StatusCode, requestLog, responseLog, env.requestID(), nil)
if env.Result == nil {
return nil
}
if out == nil {
return nil
}
return decodeResult(env.Result, out)
}
if out == nil {
logControlRequest("personal event control request", method, path, q, resp.StatusCode, requestLog, responseLog, responseRequestID(data), nil)
return nil
}
logControlRequest("personal event control request", method, path, q, resp.StatusCode, requestLog, responseLog, responseRequestID(data), nil)
return json.Unmarshal(data, out)
}
func (c *Client) decorate(req *http.Request) {
req.Header.Set("Authorization", "Bearer "+c.Identity.AccessToken)
req.Header.Set("x-user-access-token", c.Identity.AccessToken)
req.Header.Set("X-DWS-Client-Id", c.Identity.ClientID)
req.Header.Set("X-DWS-Source-Id", c.Identity.SourceID)
if c.Identity.CorpID != "" {
req.Header.Set("X-DWS-Corp-Id", c.Identity.CorpID)
}
req.Header.Set("Accept", "application/json")
}
type responseEnvelope struct {
Success *bool `json:"success"`
RequestID string `json:"request_id,omitempty"`
RequestID2 string `json:"requestId,omitempty"`
Result json.RawMessage `json:"result"`
Error *APIError `json:"error"`
ErrorCode string `json:"errorCode,omitempty"`
ErrorMsg string `json:"errorMsg,omitempty"`
}
func (e responseEnvelope) apiError() *APIError {
if e.Error != nil {
return e.Error
}
if e.ErrorCode != "" || e.ErrorMsg != "" {
return &APIError{Code: e.ErrorCode, Message: e.ErrorMsg}
}
return nil
}
func (e responseEnvelope) requestID() string {
return firstNonEmpty(e.RequestID, e.RequestID2)
}
func decodeResult(raw json.RawMessage, out any) error {
if len(raw) == 0 || string(raw) == "null" {
return nil
}
if sub, ok := out.(*Subscription); ok {
if decoded, ok, err := decodeSubscriptionResult(raw); ok || err != nil {
if err != nil {
return err
}
*sub = decoded
return nil
}
}
if err := json.Unmarshal(raw, out); err == nil {
return nil
}
return json.Unmarshal(raw, out)
}
func decodeSubscriptionResult(raw json.RawMessage) (Subscription, bool, error) {
var ids []string
if err := json.Unmarshal(raw, &ids); err == nil {
if len(ids) == 0 {
return Subscription{}, true, nil
}
return Subscription{SubscribeID: ids[0]}, true, nil
}
var item dwsSubscription
if err := json.Unmarshal(raw, &item); err == nil &&
(firstNonEmpty(item.SubID, item.SubscribeID) != "" ||
firstNonEmpty(item.EventKey, item.EventKeySnake) != "") {
return item.toSubscription(), true, nil
}
return Subscription{}, false, nil
}
func decodeAPIError(data []byte) *APIError {
var env responseEnvelope
if err := json.Unmarshal(data, &env); err == nil {
if env.Error != nil {
return env.Error
}
if env.ErrorCode != "" || env.ErrorMsg != "" {
return &APIError{Code: env.ErrorCode, Message: env.ErrorMsg}
}
}
var apiErr APIError
if err := json.Unmarshal(data, &apiErr); err == nil && (apiErr.Code != "" || apiErr.Message != "") {
return &apiErr
}
return nil
}
func responseRequestID(data []byte) string {
var env responseEnvelope
if err := json.Unmarshal(data, &env); err == nil {
return env.requestID()
}
return ""
}
func withRequestDetails(apiErr *APIError, method, path string, status int, requestID string) *APIError {
if apiErr == nil {
return nil
}
if apiErr.Details == nil {
apiErr.Details = make(map[string]any, 4)
}
apiErr.Details["method"] = method
apiErr.Details["path"] = path
apiErr.Details["http_status"] = status
if requestID != "" {
apiErr.Details["request_id"] = requestID
}
return apiErr
}
func logControlRequest(message, method, path string, q url.Values, status int, requestPayload, responsePayload, requestID string, apiErr *APIError) {
attrs := []any{
"method", method,
"path", path,
"http_status", status,
}
if query := redactedQueryString(q); query != "" {
attrs = append(attrs, "query", query)
}
if requestPayload != "" {
attrs = append(attrs, "request", requestPayload)
}
if responsePayload != "" {
attrs = append(attrs, "response", responsePayload)
}
if requestID != "" {
attrs = append(attrs, "request_id", requestID)
}
if apiErr != nil {
if apiErr.Code != "" {
attrs = append(attrs, "error_code", apiErr.Code)
}
if apiErr.Message != "" {
attrs = append(attrs, "error_msg", apiErr.Message)
}
}
slog.Debug(message, attrs...)
}
func sanitizeLogPayload(data []byte) string {
data = bytes.TrimSpace(data)
if len(data) == 0 {
return ""
}
var parsed any
if err := json.Unmarshal(data, &parsed); err == nil {
redacted := redactJSONValue(parsed)
if s, err := marshalLogJSON(redacted); err == nil {
return truncateLogPayload(s)
}
}
return truncateLogPayload(string(data))
}
func marshalLogJSON(v any) (string, error) {
var buf bytes.Buffer
enc := json.NewEncoder(&buf)
enc.SetEscapeHTML(false)
if err := enc.Encode(v); err != nil {
return "", err
}
return strings.TrimSpace(buf.String()), nil
}
func redactJSONValue(v any) any {
switch x := v.(type) {
case map[string]any:
out := make(map[string]any, len(x))
for k, value := range x {
if sensitiveLogKey(k) {
out[k] = "<redacted>"
continue
}
out[k] = redactJSONValue(value)
}
return out
case []any:
out := make([]any, len(x))
for i, value := range x {
out[i] = redactJSONValue(value)
}
return out
default:
return v
}
}
func sensitiveLogKey(key string) bool {
key = strings.ToLower(strings.TrimSpace(key))
return strings.Contains(key, "token") ||
strings.Contains(key, "secret") ||
strings.Contains(key, "ticket") ||
strings.Contains(key, "authorization")
}
func redactedQueryString(q url.Values) string {
if len(q) == 0 {
return ""
}
clone := make(url.Values, len(q))
for key, values := range q {
if sensitiveLogKey(key) {
clone[key] = []string{"<redacted>"}
continue
}
clone[key] = append([]string(nil), values...)
}
return clone.Encode()
}
func truncateLogPayload(s string) string {
if len(s) <= controlLogPayloadLimit {
return s
}
return s[:controlLogPayloadLimit] + "...<truncated>"
}
func isNotFound(err error) bool {
var apiErr *APIError
return errors.As(err, &apiErr) && (apiErr.Code == "PERSONAL_EVENT_NOT_FOUND" || apiErr.Code == "NOT_FOUND")
}
func dwsStatusString(raw json.RawMessage) string {
raw = bytes.TrimSpace(raw)
if len(raw) == 0 || string(raw) == "null" {
return ""
}
var s string
if err := json.Unmarshal(raw, &s); err == nil {
return s
}
var n int
if err := json.Unmarshal(raw, &n); err == nil {
switch n {
case 1:
return "active"
case 2:
return "paused"
case 3:
return "deleted"
default:
return fmt.Sprintf("%d", n)
}
}
return string(raw)
}
func firstNonEmpty(values ...string) string {
for _, v := range values {
if strings.TrimSpace(v) != "" {
return strings.TrimSpace(v)
}
}
return ""
}
+646
View File
@@ -0,0 +1,646 @@
// 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 personal
import (
"bytes"
"encoding/json"
"errors"
"fmt"
"log/slog"
"net/http"
"net/http/httptest"
"strconv"
"strings"
"testing"
)
func TestClientCreateSubscriptionDWSRequestAndArrayResponse(t *testing.T) {
var gotPath string
var gotReq dwsCreateSubscriptionRequest
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
t.Fatalf("method = %s", r.Method)
}
gotPath = r.URL.Path
if got := r.Header.Get("Authorization"); got != "Bearer token-1" {
t.Fatalf("Authorization = %q", got)
}
if got := r.Header.Get("x-user-access-token"); got != "token-1" {
t.Fatalf("x-user-access-token = %q", got)
}
if got := r.Header.Get("X-DWS-Client-Id"); got != "client-1" {
t.Fatalf("X-DWS-Client-Id = %q", got)
}
if got := r.Header.Get("X-DWS-Source-Id"); got != "open" {
t.Fatalf("X-DWS-Source-Id = %q", got)
}
if got := r.Header.Get("X-DWS-Corp-Id"); got != "corp-1" {
t.Fatalf("X-DWS-Corp-Id = %q", got)
}
if err := json.NewDecoder(r.Body).Decode(&gotReq); err != nil {
t.Fatalf("Decode body: %v", err)
}
_ = json.NewEncoder(w).Encode(map[string]any{
"success": true,
"result": []string{"sub-1"},
})
}))
defer srv.Close()
c := NewClient(srv.URL, Identity{
AccessToken: "token-1",
CorpID: "corp-1",
UserID: "user-1",
ClientID: "client-1",
SourceID: "open",
})
sub, err := c.CreateSubscription(t.Context(), CreateSubscriptionRequest{
EventKey: EventSingleChat,
RuleType: "singleChat",
Name: "test-o2o",
RuleParam: map[string]any{
"targetUid": "507971",
"targetUidType": "staffId",
},
Filter: map[string]any{"field": "payload.body.content", "op": "contains", "value": "P0"},
IdempotencyKey: "idem-1",
})
if err != nil {
t.Fatalf("CreateSubscription() error = %v", err)
}
if gotPath != "/subscription/user" {
t.Fatalf("path = %q, want /subscription/user", gotPath)
}
if gotReq.ClientID != "client-1" || gotReq.SourceID != "open" || gotReq.EventKey != EventSingleChat {
t.Fatalf("request identity/event = %#v", gotReq)
}
if gotReq.DeliveryPref != "realtime" {
t.Fatalf("deliveryPref = %q, want realtime", gotReq.DeliveryPref)
}
var filterRule map[string]any
if err := json.Unmarshal([]byte(gotReq.FilterRule), &filterRule); err != nil {
t.Fatalf("filterRule is not JSON: %q: %v", gotReq.FilterRule, err)
}
if filterRule["targetUid"] != "507971" || filterRule["targetUidType"] != "staffId" {
t.Fatalf("filterRule = %#v", filterRule)
}
if gotReq.Ext["ruleType"] != "singleChat" || gotReq.Ext["name"] != "test-o2o" || gotReq.Ext["idempotencyKey"] != "idem-1" {
t.Fatalf("ext = %#v", gotReq.Ext)
}
if sub.SubscribeID != "sub-1" {
t.Fatalf("subscribe_id = %q", sub.SubscribeID)
}
if sub.EventKey != EventSingleChat || sub.RuleType != "singleChat" || sub.Status != "active" || sub.SourceID != "open" {
t.Fatalf("subscription = %#v", sub)
}
}
func TestClientCreateSubscriptionObjectResponses(t *testing.T) {
cases := []map[string]any{
{"subId": "sub-camel", "eventKey": EventMention, "sourceId": "open", "status": 1},
{"subscribe_id": "sub-snake", "event_key": EventMention, "source_id": "open", "status": "active"},
}
for _, result := range cases {
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
_ = json.NewEncoder(w).Encode(map[string]any{
"success": true,
"result": result,
})
}))
c := NewClient(srv.URL, Identity{AccessToken: "token-1", ClientID: "client-1", SourceID: "open"})
sub, err := c.CreateSubscription(t.Context(), CreateSubscriptionRequest{
EventKey: EventMention,
RuleType: "at",
RuleParam: map[string]any{},
})
srv.Close()
if err != nil {
t.Fatalf("CreateSubscription() error = %v", err)
}
if sub.SubscribeID == "" || !strings.HasPrefix(sub.SubscribeID, "sub-") {
t.Fatalf("subscription = %#v", sub)
}
}
}
func TestClientDebugLogCreateSubscriptionRequestResponse(t *testing.T) {
logs := captureClientDebugLogs(t)
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
_ = json.NewEncoder(w).Encode(map[string]any{
"success": true,
"requestId": "req-ok",
"result": []string{"sub-1"},
})
}))
defer srv.Close()
c := NewClient(srv.URL, Identity{AccessToken: "secret-token", ClientID: "client-1", SourceID: "pre_open_source"})
if _, err := c.CreateSubscription(t.Context(), CreateSubscriptionRequest{
EventKey: EventSingleChat,
RuleType: "singleChat",
RuleParam: map[string]any{
"targetUid": "507971",
"targetUidType": "staffId",
},
IdempotencyKey: "idem-1",
}); err != nil {
t.Fatalf("CreateSubscription() error = %v", err)
}
out := logs.String()
for _, want := range []string{
"personal event control request",
"/subscription/user",
"client-1",
"pre_open_source",
EventSingleChat,
"filterRule",
"targetUid",
"507971",
"sub-1",
"req-ok",
} {
if !strings.Contains(out, want) {
t.Fatalf("debug log missing %q: %s", want, out)
}
}
if strings.Contains(out, "secret-token") {
t.Fatalf("debug log leaked access token: %s", out)
}
}
func TestClientBusinessErrorHTTP200(t *testing.T) {
logs := captureClientDebugLogs(t)
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
_ = json.NewEncoder(w).Encode(map[string]any{
"success": false,
"requestId": "req-1",
"errorCode": "INVALID_PARAM",
"errorMsg": "clientId is empty",
})
}))
defer srv.Close()
c := NewClient(srv.URL, Identity{AccessToken: "token-1", ClientID: "client-1", SourceID: "open"})
_, err := c.CreateSubscription(t.Context(), CreateSubscriptionRequest{
EventKey: EventMention,
RuleType: "at",
RuleParam: map[string]any{},
})
if err == nil || !strings.Contains(err.Error(), "INVALID_PARAM") || !strings.Contains(err.Error(), "clientId is empty") {
t.Fatalf("error = %v, want INVALID_PARAM business error", err)
}
var apiErr *APIError
if !errors.As(err, &apiErr) {
t.Fatalf("error type = %T, want *APIError", err)
}
if apiErr.Details["method"] != http.MethodPost || apiErr.Details["path"] != "/subscription/user" ||
apiErr.Details["http_status"] != http.StatusOK || apiErr.Details["request_id"] != "req-1" {
t.Fatalf("details = %#v", apiErr.Details)
}
out := logs.String()
for _, want := range []string{"/subscription/user", "INVALID_PARAM", "clientId is empty", "req-1", "request", "response"} {
if !strings.Contains(out, want) {
t.Fatalf("debug log missing %q: %s", want, out)
}
}
}
func TestClientOmitsCorpHeaderWhenUnknown(t *testing.T) {
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if got := r.Header.Get("X-DWS-Corp-Id"); got != "" {
t.Fatalf("X-DWS-Corp-Id = %q, want empty", got)
}
if got := r.Header.Get("Authorization"); got != "Bearer token-1" {
t.Fatalf("Authorization = %q", got)
}
_ = json.NewEncoder(w).Encode(map[string]any{
"success": true,
"result": []string{"sub-1"},
})
}))
defer srv.Close()
c := NewClient(srv.URL, Identity{AccessToken: "token-1", ClientID: "client-1", SourceID: "open"})
if _, err := c.CreateSubscription(t.Context(), CreateSubscriptionRequest{
EventKey: EventMention,
RuleType: "at",
RuleParam: map[string]any{},
}); err != nil {
t.Fatalf("CreateSubscription() error = %v", err)
}
}
func TestIdentityKeyUsesLocalSubjectFallback(t *testing.T) {
withCorpUser := Identity{CorpID: "corp-1", UserID: "user-1", ClientID: "client-1", SourceID: "open"}
if got := withCorpUser.Key(); got != "corp_user\x00corp-1\x00user-1\x00client-1\x00open" {
t.Fatalf("corp/user key = %q", got)
}
fallback := Identity{LocalSubject: "refresh:abc", ClientID: "client-1", SourceID: "open"}
if got := fallback.Key(); got != "local_subject\x00refresh:abc\x00client-1\x00open" {
t.Fatalf("fallback key = %q", got)
}
}
func TestClientDeleteSubscriptionTreatsNotFoundAsSuccess(t *testing.T) {
var gotBody map[string]string
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
t.Fatalf("method = %s", r.Method)
}
if r.URL.Path != "/subscription/cancel" {
t.Fatalf("path = %q", r.URL.Path)
}
if err := json.NewDecoder(r.Body).Decode(&gotBody); err != nil {
t.Fatalf("decode body: %v", err)
}
_ = json.NewEncoder(w).Encode(map[string]any{
"success": false,
"error": map[string]any{"code": "PERSONAL_EVENT_NOT_FOUND", "message": "not found"},
})
}))
defer srv.Close()
c := NewClient(srv.URL, Identity{AccessToken: "token", ClientID: "client", SourceID: "open"})
if err := c.DeleteSubscription(t.Context(), "sub-404"); err != nil {
t.Fatalf("DeleteSubscription() error = %v", err)
}
if gotBody["subId"] != "sub-404" {
t.Fatalf("cancel body = %#v", gotBody)
}
}
func TestClientDeleteSubscriptionBusinessError(t *testing.T) {
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
_ = json.NewEncoder(w).Encode(map[string]any{
"success": false,
"requestId": "req-cancel",
"errorCode": "INVALID_STATE",
"errorMsg": "subscription cannot be cancelled",
})
}))
defer srv.Close()
c := NewClient(srv.URL, Identity{AccessToken: "token-1", ClientID: "client-1", SourceID: "open"})
err := c.DeleteSubscription(t.Context(), "sub-1")
if err == nil || !strings.Contains(err.Error(), "INVALID_STATE") {
t.Fatalf("DeleteSubscription() error = %v, want INVALID_STATE", err)
}
var apiErr *APIError
if !errors.As(err, &apiErr) {
t.Fatalf("error type = %T, want *APIError", err)
}
if apiErr.Details["path"] != "/subscription/cancel" || apiErr.Details["request_id"] != "req-cancel" {
t.Fatalf("details = %#v", apiErr.Details)
}
}
func TestClientDebugLogListAndDeleteSubscription(t *testing.T) {
logs := captureClientDebugLogs(t)
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch r.URL.Path {
case "/event/sublist":
_ = json.NewEncoder(w).Encode(map[string]any{
"success": true,
"result": map[string]any{"items": []map[string]any{}},
})
case "/subscription/cancel":
_ = json.NewEncoder(w).Encode(map[string]any{"success": true})
default:
t.Fatalf("unexpected path %q", r.URL.Path)
}
}))
defer srv.Close()
c := NewClient(srv.URL, Identity{AccessToken: "token", ClientID: "client", SourceID: "open"})
if _, err := c.ListSubscriptions(t.Context(), ListOptions{Status: "active"}); err != nil {
t.Fatalf("ListSubscriptions() error = %v", err)
}
if err := c.DeleteSubscription(t.Context(), "sub-1"); err != nil {
t.Fatalf("DeleteSubscription() error = %v", err)
}
out := logs.String()
for _, want := range []string{"/event/sublist", "clientId=client", "sourceId=open", "/subscription/cancel", "sub-1"} {
if !strings.Contains(out, want) {
t.Fatalf("debug log missing %q: %s", want, out)
}
}
}
func TestClientDebugLogRedactsSensitivePayloadFields(t *testing.T) {
logs := captureClientDebugLogs(t)
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
_ = json.NewEncoder(w).Encode(map[string]any{
"success": true,
"requestId": "req-secret",
"result": map[string]any{
"access_token": "resp-access-token",
"client_secret": "resp-client-secret",
"ticket": "resp-ticket",
"Authorization": "Bearer resp-auth",
"safe": "ok",
},
})
}))
defer srv.Close()
c := NewClient(srv.URL, Identity{AccessToken: "header-token-secret", ClientID: "client", SourceID: "open"})
err := c.do(t.Context(), http.MethodPost, "/subscription/user", nil, map[string]any{
"access_token": "req-access-token",
"client_secret": "req-client-secret",
"ticket": "req-ticket",
"Authorization": "Bearer req-auth",
"safe": "ok",
}, nil)
if err != nil {
t.Fatalf("do() error = %v", err)
}
out := logs.String()
for _, leaked := range []string{
"header-token-secret",
"req-access-token",
"req-client-secret",
"req-ticket",
"Bearer req-auth",
"resp-access-token",
"resp-client-secret",
"resp-ticket",
"Bearer resp-auth",
} {
if strings.Contains(out, leaked) {
t.Fatalf("debug log leaked %q: %s", leaked, out)
}
}
for _, want := range []string{"<redacted>", "safe", "ok", "req-secret"} {
if !strings.Contains(out, want) {
t.Fatalf("debug log missing %q: %s", want, out)
}
}
}
func TestClientListSubscriptionsDWSSublist(t *testing.T) {
var gotQuery string
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path != "/event/sublist" {
t.Fatalf("path = %q", r.URL.Path)
}
gotQuery = r.URL.RawQuery
_ = json.NewEncoder(w).Encode(map[string]any{
"success": true,
"result": map[string]any{
"total": 2,
"pageNo": 1,
"pageSize": 20,
"items": []map[string]any{
{
"subId": "sub-1",
"eventKey": EventSingleChat,
"sourceId": "open",
"deliveryPref": "realtime",
"status": 1,
"gmtCreate": "2026-06-29T10:00:00Z",
},
{
"subId": "sub-2",
"eventKey": EventMention,
"sourceId": "open",
"status": 3,
},
},
},
})
}))
defer srv.Close()
c := NewClient(srv.URL, Identity{AccessToken: "token", ClientID: "client", SourceID: "open"})
subs, err := c.ListSubscriptions(t.Context(), ListOptions{Status: "active", EventKey: EventSingleChat})
if err != nil {
t.Fatalf("ListSubscriptions() error = %v", err)
}
if !strings.Contains(gotQuery, "clientId=client") || !strings.Contains(gotQuery, "sourceId=open") ||
!strings.Contains(gotQuery, "pageNo=1") || !strings.Contains(gotQuery, "pageSize=100") {
t.Fatalf("query = %q", gotQuery)
}
if len(subs) != 1 || subs[0].SubscribeID != "sub-1" || subs[0].Status != "active" || subs[0].CreatedAt == "" {
t.Fatalf("subs = %#v", subs)
}
}
func TestClientGetSubscriptionFiltersSublist(t *testing.T) {
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
_ = json.NewEncoder(w).Encode(map[string]any{
"success": true,
"result": map[string]any{
"items": []map[string]any{
{"subId": "sub-1", "eventKey": EventMention, "sourceId": "open", "status": 1},
{"subId": "sub-2", "eventKey": EventSingleChat, "sourceId": "open", "status": 1},
},
},
})
}))
defer srv.Close()
c := NewClient(srv.URL, Identity{AccessToken: "token", ClientID: "client", SourceID: "open"})
sub, err := c.GetSubscription(t.Context(), "sub-2")
if err != nil {
t.Fatalf("GetSubscription() error = %v", err)
}
if sub.SubscribeID != "sub-2" || sub.EventKey != EventSingleChat {
t.Fatalf("subscription = %#v", sub)
}
}
func TestClientListSubscriptionsPaginatesAllResults(t *testing.T) {
var pages []int
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
pageNo, err := strconv.Atoi(r.URL.Query().Get("pageNo"))
if err != nil {
t.Fatalf("pageNo = %q", r.URL.Query().Get("pageNo"))
}
pages = append(pages, pageNo)
start := (pageNo - 1) * subscriptionListPageSize
end := start + subscriptionListPageSize
if end > 205 {
end = 205
}
_ = json.NewEncoder(w).Encode(map[string]any{
"success": true,
"result": map[string]any{
"total": 205,
"items": dwsSubscriptionTestItems(start, end),
},
})
}))
defer srv.Close()
c := NewClient(srv.URL, Identity{AccessToken: "token", ClientID: "client", SourceID: "open"})
subs, err := c.ListSubscriptions(t.Context(), ListOptions{})
if err != nil {
t.Fatalf("ListSubscriptions() error = %v", err)
}
if len(subs) != 205 || subs[204].SubscribeID != "sub-204" {
t.Fatalf("subscriptions = %d, last = %#v", len(subs), subs[len(subs)-1])
}
if fmt.Sprint(pages) != "[1 2 3]" {
t.Fatalf("pages = %v, want [1 2 3]", pages)
}
}
func TestClientGetSubscriptionFindsLaterPage(t *testing.T) {
var calls int
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
calls++
pageNo, _ := strconv.Atoi(r.URL.Query().Get("pageNo"))
items := dwsSubscriptionTestItems(0, 100)
if pageNo == 2 {
items = dwsSubscriptionTestItems(100, 101)
}
_ = json.NewEncoder(w).Encode(map[string]any{
"success": true,
"result": map[string]any{"total": 101, "items": items},
})
}))
defer srv.Close()
c := NewClient(srv.URL, Identity{AccessToken: "token", ClientID: "client", SourceID: "open"})
sub, err := c.GetSubscription(t.Context(), "sub-100")
if err != nil {
t.Fatalf("GetSubscription() error = %v", err)
}
if sub.SubscribeID != "sub-100" || calls != 2 {
t.Fatalf("subscription = %#v, calls = %d", sub, calls)
}
}
func TestClientListSubscriptionsWithoutTotalStopsOnEmptyPage(t *testing.T) {
var calls int
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
calls++
items := dwsSubscriptionTestItems(0, 100)
if r.URL.Query().Get("pageNo") == "2" {
items = nil
}
_ = json.NewEncoder(w).Encode(map[string]any{
"success": true,
"result": map[string]any{"items": items},
})
}))
defer srv.Close()
c := NewClient(srv.URL, Identity{AccessToken: "token", ClientID: "client", SourceID: "open"})
subs, err := c.ListSubscriptions(t.Context(), ListOptions{})
if err != nil {
t.Fatal(err)
}
if len(subs) != 100 || calls != 2 {
t.Fatalf("subscriptions = %d, calls = %d", len(subs), calls)
}
}
func TestClientListSubscriptionsUsesServerPageSize(t *testing.T) {
var calls int
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
calls++
pageNo, _ := strconv.Atoi(r.URL.Query().Get("pageNo"))
start := (pageNo - 1) * 20
end := start + 20
if end > 45 {
end = 45
}
_ = json.NewEncoder(w).Encode(map[string]any{
"success": true,
"result": map[string]any{
"total": 45,
"pageSize": 20,
"items": dwsSubscriptionTestItems(start, end),
},
})
}))
defer srv.Close()
c := NewClient(srv.URL, Identity{AccessToken: "token", ClientID: "client", SourceID: "open"})
subs, err := c.ListSubscriptions(t.Context(), ListOptions{})
if err != nil {
t.Fatal(err)
}
if len(subs) != 45 || calls != 3 {
t.Fatalf("subscriptions = %d, calls = %d", len(subs), calls)
}
}
func TestClientListSubscriptionsDeduplicatesSubscribeID(t *testing.T) {
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
items := dwsSubscriptionTestItems(0, 100)
if r.URL.Query().Get("pageNo") == "2" {
items = append(dwsSubscriptionTestItems(99, 100), dwsSubscriptionTestItems(100, 101)...)
}
_ = json.NewEncoder(w).Encode(map[string]any{
"success": true,
"result": map[string]any{"total": 101, "items": items},
})
}))
defer srv.Close()
c := NewClient(srv.URL, Identity{AccessToken: "token", ClientID: "client", SourceID: "open"})
subs, err := c.ListSubscriptions(t.Context(), ListOptions{})
if err != nil {
t.Fatal(err)
}
if len(subs) != 101 || subs[100].SubscribeID != "sub-100" {
t.Fatalf("subscriptions = %#v", subs)
}
}
func TestClientListSubscriptionsRejectsRepeatedPage(t *testing.T) {
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
_ = json.NewEncoder(w).Encode(map[string]any{
"success": true,
"result": map[string]any{
"total": 200,
"items": dwsSubscriptionTestItems(0, 100),
},
})
}))
defer srv.Close()
c := NewClient(srv.URL, Identity{AccessToken: "token", ClientID: "client", SourceID: "open"})
_, err := c.ListSubscriptions(t.Context(), ListOptions{})
if err == nil || !strings.Contains(err.Error(), "pagination made no progress") {
t.Fatalf("ListSubscriptions() error = %v", err)
}
}
func dwsSubscriptionTestItems(start, end int) []map[string]any {
items := make([]map[string]any, 0, end-start)
for i := start; i < end; i++ {
items = append(items, map[string]any{
"subId": fmt.Sprintf("sub-%d", i),
"eventKey": EventSingleChat,
"sourceId": "open",
"status": 1,
})
}
return items
}
func captureClientDebugLogs(t *testing.T) *bytes.Buffer {
t.Helper()
var buf bytes.Buffer
previous := slog.Default()
slog.SetDefault(slog.New(slog.NewJSONHandler(&buf, &slog.HandlerOptions{Level: slog.LevelDebug})))
t.Cleanup(func() {
slog.SetDefault(previous)
})
return &buf
}
+401
View File
@@ -0,0 +1,401 @@
// 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 personal
import (
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"strings"
)
const (
EventMention = "user_im_message_receive_at"
EventSingleChat = "user_im_message_receive_o2o"
EventInChat = "user_im_message_receive_group"
EventFromUser = "user_im_message_receive_user"
)
const (
StatusEnabled = "enabled"
StatusPending = "pending"
)
type Definition struct {
EventKey string `json:"event_key"`
DisplayName string `json:"display_name"`
Description string `json:"description"`
Category string `json:"category"`
RuleType string `json:"rule_type"`
Status string `json:"status"`
RequiredParams []string `json:"required_params"`
Auth map[string]any `json:"auth,omitempty"`
Public bool `json:"-"`
}
type SchemaDocument struct {
EventKey string `json:"event_key"`
DisplayName string `json:"display_name"`
Description string `json:"description"`
Category string `json:"category"`
RuleType string `json:"rule_type"`
RequiredParams []string `json:"required_params"`
JQRootPath string `json:"jq_root_path"`
Schema map[string]any `json:"schema"`
}
type RuleOptions struct {
RuleType string
UserID string
GroupID string
}
type SchemaPendingError struct {
EventKey string
}
func (e *SchemaPendingError) Error() string {
return fmt.Sprintf("%s schema is pending; try user_im_message_receive_at or user_im_message_receive_o2o first", e.EventKey)
}
var definitions = []Definition{
{
EventKey: EventMention,
DisplayName: "@我的消息",
Description: "当前用户被 @ 的消息",
Category: "im",
RuleType: "at",
Status: StatusEnabled,
RequiredParams: nil,
Auth: map[string]any{"identity": "user"},
Public: true,
},
{
EventKey: EventSingleChat,
DisplayName: "指定单聊消息",
Description: "当前用户与指定用户的单聊消息",
Category: "im",
RuleType: "singleChat",
Status: StatusEnabled,
RequiredParams: []string{"user"},
Auth: map[string]any{"identity": "user"},
Public: true,
},
{
EventKey: EventInChat,
DisplayName: "指定群消息",
Description: "当前用户所在指定会话的消息",
Category: "im",
RuleType: "group",
Status: StatusEnabled,
RequiredParams: []string{"group"},
Auth: map[string]any{"identity": "user"},
Public: true,
},
{
EventKey: EventFromUser,
DisplayName: "指定发送人消息",
Description: "当前用户收到的特别关注用户的消息",
Category: "im",
RuleType: "sender",
Status: StatusEnabled,
RequiredParams: []string{"user"},
Auth: map[string]any{"identity": "user"},
Public: false,
},
}
func Definitions() []Definition {
out := append([]Definition(nil), definitions...)
return out
}
func Lookup(eventKey string) (Definition, bool) {
for _, def := range definitions {
if def.EventKey == eventKey {
return def, true
}
}
return Definition{}, false
}
func IsPublic(eventKey string) bool {
def, ok := Lookup(eventKey)
return ok && def.Public
}
func PublicAvailabilityError(eventKey string) error {
return fmt.Errorf("event %s is not publicly available yet", eventKey)
}
func Catalog(category string, enabledOnly, includePending bool) []Definition {
category = strings.TrimSpace(category)
var out []Definition
for _, def := range definitions {
if !def.Public {
continue
}
if category != "" && def.Category != category {
continue
}
if enabledOnly && def.Status != StatusEnabled {
continue
}
if !includePending && def.Status == StatusPending {
continue
}
out = append(out, def)
}
return out
}
func BuildSchemaDocument(def Definition) SchemaDocument {
requiredParams := make([]string, 0, len(def.RequiredParams))
requiredParams = append(requiredParams, def.RequiredParams...)
return SchemaDocument{
EventKey: def.EventKey,
DisplayName: def.DisplayName,
Description: def.Description,
Category: def.Category,
RuleType: def.RuleType,
RequiredParams: requiredParams,
JQRootPath: ".data | fromjson",
Schema: personalMessageSchema(def.EventKey),
}
}
func BuildRuleParam(eventKey string, opts RuleOptions) (ruleType string, ruleParam map[string]any, err error) {
def, ok := Lookup(eventKey)
if !ok {
return "", nil, fmt.Errorf("unknown personal event key %q", eventKey)
}
if opts.RuleType != "" && opts.RuleType != def.RuleType {
return "", nil, fmt.Errorf("--rule %q does not match %s rule %q", opts.RuleType, eventKey, def.RuleType)
}
if def.Status == StatusPending {
return "", nil, &SchemaPendingError{EventKey: eventKey}
}
userID := strings.TrimSpace(opts.UserID)
groupID := strings.TrimSpace(opts.GroupID)
switch def.RuleType {
case "at":
if userID != "" {
return "", nil, fmt.Errorf("--user is only supported for %s", EventSingleChat)
}
if groupID != "" {
return "", nil, fmt.Errorf("--group is only supported for %s", EventInChat)
}
return def.RuleType, map[string]any{}, nil
case "singleChat":
if groupID != "" {
return "", nil, fmt.Errorf("--group is only supported for %s", EventInChat)
}
if userID == "" {
return "", nil, fmt.Errorf("--user is required")
}
return def.RuleType, map[string]any{
"targetUid": userID,
"targetUidType": "staffId",
}, nil
case "sender":
if groupID != "" {
return "", nil, fmt.Errorf("--group is only supported for %s", EventInChat)
}
if userID == "" {
return "", nil, fmt.Errorf("--user is required")
}
return def.RuleType, map[string]any{
"targetUid": userID,
"targetUidType": "staffId",
}, nil
case "group":
if userID != "" {
return "", nil, fmt.Errorf("--user is only supported for %s", EventSingleChat)
}
if groupID == "" {
return "", nil, fmt.Errorf("--group is required")
}
return def.RuleType, map[string]any{
"openConversationId": groupID,
}, nil
default:
return "", nil, &SchemaPendingError{EventKey: eventKey}
}
}
func BuildFilter(filterJSON string, queryCSV string) (any, string, error) {
var parts []any
filterJSON = strings.TrimSpace(filterJSON)
if filterJSON != "" {
var v any
if err := json.Unmarshal([]byte(filterJSON), &v); err != nil {
return nil, "", fmt.Errorf("--filter-json must be valid JSON: %w", err)
}
v = normalizeFilterAliases(v)
parts = append(parts, v)
}
queries := splitCSV(queryCSV)
if len(queries) > 0 {
parts = append(parts, map[string]any{
"field": "payload.body.content",
"op": "contains_any",
"value": queries,
})
}
switch len(parts) {
case 0:
return nil, "", nil
case 1:
canon, err := CanonicalJSON(parts[0])
return parts[0], canon, err
default:
v := map[string]any{"and": parts}
canon, err := CanonicalJSON(v)
return v, canon, err
}
}
func IdempotencyKey(identity Identity, eventKey, ruleType string, ruleParam map[string]any, filterCanonical string) string {
ruleCanonical, _ := CanonicalJSON(ruleParam)
sum := sha256.Sum256([]byte(strings.Join([]string{
identity.Key(),
eventKey,
ruleType,
ruleCanonical,
filterCanonical,
}, "\x00")))
return "dws-cli-" + hex.EncodeToString(sum[:8])
}
func CanonicalJSON(v any) (string, error) {
if v == nil {
return "", nil
}
b, err := json.Marshal(v)
if err != nil {
return "", err
}
return string(b), nil
}
func splitCSV(raw string) []string {
var out []string
for _, part := range strings.Split(raw, ",") {
part = strings.TrimSpace(part)
if part != "" {
out = append(out, part)
}
}
return out
}
var filterFieldAliases = map[string]string{
"content": "payload.body.content",
"conversation_id": "payload.body.openConversationId",
"sender": "payload.body.sender",
"sender_open_dingtalk_id": "payload.body.senderOpenDingTalkId",
}
func normalizeFilterAliases(v any) any {
switch x := v.(type) {
case map[string]any:
out := make(map[string]any, len(x))
for k, value := range x {
if k == "field" {
if raw, ok := value.(string); ok {
if mapped, ok := filterFieldAliases[raw]; ok {
value = mapped
}
}
} else {
value = normalizeFilterAliases(value)
}
out[k] = value
}
return out
case []any:
out := make([]any, len(x))
for i, value := range x {
out[i] = normalizeFilterAliases(value)
}
return out
default:
return v
}
}
func personalMessageSchema(eventKey string) map[string]any {
return map[string]any{
"type": "object",
"properties": map[string]any{
"type": map[string]any{
"type": "string",
"description": "事件类型,固定为当前 event_key",
"enum": []string{eventKey},
},
"event_id": map[string]any{
"type": "string",
"description": "事件 ID,可用于去重",
},
"timestamp": map[string]any{
"type": "integer",
"description": "事件发生时间戳,对应 occurredAtMs",
"format": "timestamp_ms",
},
"subscribe_id": map[string]any{
"type": "string",
"description": "订阅 ID,对应 subId",
},
"message_id": map[string]any{
"type": "string",
"description": "开放消息 ID,对应 payload.body.openMessageId",
"format": "open_message_id",
},
"conversation_id": map[string]any{
"type": "string",
"description": "会话 ID,对应 payload.body.openConversationId",
"format": "open_conversation_id",
},
"sender": map[string]any{
"type": "string",
"description": "发送人展示名,对应 payload.body.sender",
},
"sender_open_dingtalk_id": map[string]any{
"type": "string",
"description": "发送人开放 ID,对应 payload.body.senderOpenDingTalkId",
"format": "open_dingtalk_id",
},
"content": map[string]any{
"type": "string",
"description": "消息正文,对应 payload.body.content",
},
"create_time": map[string]any{
"type": "string",
"description": "消息创建时间,对应 payload.body.createTime",
},
"event_time": map[string]any{
"type": "integer",
"description": "消息事件时间戳,对应 payload.event_time",
"format": "timestamp_ms",
},
},
}
}
func IsSchemaPending(err error) bool {
var pending *SchemaPendingError
return errors.As(err, &pending)
}
+271
View File
@@ -0,0 +1,271 @@
// 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 personal
import (
"encoding/json"
"reflect"
"strings"
"testing"
)
func TestCatalogEnabledEvents(t *testing.T) {
items := Catalog("", true, false)
keys := make([]string, 0, len(items))
for _, item := range items {
keys = append(keys, item.EventKey)
if item.Status != StatusEnabled {
t.Fatalf("%s status = %q, want enabled", item.EventKey, item.Status)
}
}
want := []string{
EventMention,
EventSingleChat,
EventInChat,
}
if !reflect.DeepEqual(keys, want) {
t.Fatalf("keys = %#v, want %#v", keys, want)
}
}
func TestEventFromUserRemainsInternalButNotPublic(t *testing.T) {
if _, ok := Lookup(EventFromUser); !ok {
t.Fatalf("Lookup(%q) failed, want internal definition retained", EventFromUser)
}
if IsPublic(EventFromUser) {
t.Fatalf("IsPublic(%q) = true, want hidden", EventFromUser)
}
}
func TestLegacyEventKeysAreUnknown(t *testing.T) {
legacyKeys := []string{
"im_message_receive_at",
"im_message_receive_o2o",
"im_message_receive_group",
"im_message_receive_user",
}
for _, key := range legacyKeys {
if _, ok := Lookup(key); ok {
t.Fatalf("Lookup(%q) succeeded, want unknown", key)
}
if _, _, err := BuildRuleParam(key, RuleOptions{}); err == nil || !strings.Contains(err.Error(), "unknown personal event key") {
t.Fatalf("BuildRuleParam(%q) error = %v, want unknown personal event key", key, err)
}
}
}
func TestDefinitionJSONHidesInternalSchemaIDs(t *testing.T) {
raw, err := json.Marshal(Definitions())
if err != nil {
t.Fatalf("Marshal() error = %v", err)
}
out := string(raw)
for _, leaked := range []string{"schema_ids", "im_msg_23", "im_msg_29"} {
if strings.Contains(out, leaked) {
t.Fatalf("definitions JSON leaked %q: %s", leaked, out)
}
}
}
func TestSchemaDocumentsUseSingleJSONSchema(t *testing.T) {
for _, eventKey := range []string{EventMention, EventSingleChat, EventInChat, EventFromUser} {
t.Run(eventKey, func(t *testing.T) {
def, ok := Lookup(eventKey)
if !ok {
t.Fatalf("Lookup(%q) failed", eventKey)
}
doc := BuildSchemaDocument(def)
raw, err := json.Marshal(doc)
if err != nil {
t.Fatalf("Marshal() error = %v", err)
}
out := string(raw)
for _, want := range []string{
"event_key",
"display_name",
"description",
"category",
"rule_type",
"required_params",
"jq_root_path",
"schema",
"event_id",
"timestamp",
"subscribe_id",
"content",
"sender",
"sender_open_dingtalk_id",
"conversation_id",
"message_id",
"create_time",
"event_time",
} {
if !strings.Contains(out, want) {
t.Fatalf("schema for %s missing %q: %s", eventKey, want, out)
}
}
for _, leaked := range []string{
"message.text",
"chat.openConversationId",
"sender.userId",
"sender.unionId",
"auth",
"resolved_output_schema",
"decoded_data_schema",
"filter_schema",
"payload_schema",
"output_schema",
"data_json_path",
"headers",
"audit",
"tenant",
"subject",
"traceId",
"msgIdMetaq",
"at_users",
"sender_user_id",
} {
if strings.Contains(out, leaked) {
t.Fatalf("schema for %s leaked %q: %s", eventKey, leaked, out)
}
}
if doc.JQRootPath != ".data | fromjson" {
t.Fatalf("jq_root_path = %q, want .data | fromjson", doc.JQRootPath)
}
if doc.RequiredParams == nil {
t.Fatalf("required_params = nil, want empty slice")
}
props, ok := doc.Schema["properties"].(map[string]any)
if !ok {
t.Fatalf("schema.properties = %#v, want object", doc.Schema["properties"])
}
if _, ok := props["content"].(map[string]any); !ok {
t.Fatalf("schema.properties.content = %#v, want object", props["content"])
}
})
}
}
func TestBuildRuleParamMention(t *testing.T) {
rule, param, err := BuildRuleParam(EventMention, RuleOptions{})
if err != nil {
t.Fatalf("BuildRuleParam() error = %v", err)
}
if rule != "at" {
t.Fatalf("rule = %q, want at", rule)
}
if len(param) != 0 {
t.Fatalf("param = %#v, want empty map", param)
}
}
func TestBuildRuleParamSingleChatRequiresPeer(t *testing.T) {
_, _, err := BuildRuleParam(EventSingleChat, RuleOptions{})
if err == nil || !strings.Contains(err.Error(), "--user is required") {
t.Fatalf("error = %v, want user requirement", err)
}
}
func TestBuildRuleParamSingleChatUserIDMapsToStaffID(t *testing.T) {
rule, param, err := BuildRuleParam(EventSingleChat, RuleOptions{UserID: "staff-1"})
if err != nil {
t.Fatalf("BuildRuleParam() error = %v", err)
}
if rule != "singleChat" {
t.Fatalf("rule = %q, want singleChat", rule)
}
if param["targetUidType"] != "staffId" || param["targetUid"] != "staff-1" {
t.Fatalf("param = %#v", param)
}
}
func TestBuildRuleParamSender(t *testing.T) {
_, _, err := BuildRuleParam(EventFromUser, RuleOptions{})
if err == nil || !strings.Contains(err.Error(), "--user is required") {
t.Fatalf("error = %v, want sender requirement", err)
}
rule, param, err := BuildRuleParam(EventFromUser, RuleOptions{UserID: "staff-1"})
if err != nil {
t.Fatalf("BuildRuleParam() error = %v", err)
}
if rule != "sender" {
t.Fatalf("rule = %q, want sender", rule)
}
if param["targetUidType"] != "staffId" || param["targetUid"] != "staff-1" {
t.Fatalf("param = %#v", param)
}
}
func TestBuildRuleParamGroup(t *testing.T) {
_, _, err := BuildRuleParam(EventInChat, RuleOptions{})
if err == nil || !strings.Contains(err.Error(), "--group is required") {
t.Fatalf("error = %v, want group requirement", err)
}
rule, param, err := BuildRuleParam(EventInChat, RuleOptions{GroupID: "cid-1"})
if err != nil {
t.Fatalf("BuildRuleParam() error = %v", err)
}
if rule != "group" {
t.Fatalf("rule = %q, want group", rule)
}
if param["openConversationId"] != "cid-1" {
t.Fatalf("param = %#v", param)
}
}
func TestBuildRuleParamRejectsWrongScopedFlags(t *testing.T) {
if _, _, err := BuildRuleParam(EventMention, RuleOptions{UserID: "staff-1"}); err == nil || !strings.Contains(err.Error(), "--user is only supported") {
t.Fatalf("mention with user error = %v, want unsupported user", err)
}
if _, _, err := BuildRuleParam(EventSingleChat, RuleOptions{UserID: "staff-1", GroupID: "cid-1"}); err == nil || !strings.Contains(err.Error(), "--group is only supported") {
t.Fatalf("singleChat with group error = %v, want unsupported group", err)
}
if _, _, err := BuildRuleParam(EventInChat, RuleOptions{UserID: "staff-1", GroupID: "cid-1"}); err == nil || !strings.Contains(err.Error(), "--user is only supported") {
t.Fatalf("group with user error = %v, want unsupported user", err)
}
}
func TestBuildFilterQueryAndJSON(t *testing.T) {
filter, canonical, err := BuildFilter(`{"field":"conversation_id","op":"eq","value":"cid1"}`, "P0, 故障")
if err != nil {
t.Fatalf("BuildFilter() error = %v", err)
}
m := filter.(map[string]any)
parts := m["and"].([]any)
if len(parts) != 2 {
t.Fatalf("parts len = %d, want 2", len(parts))
}
if !strings.Contains(canonical, "contains_any") {
t.Fatalf("canonical = %s, want contains_any", canonical)
}
if !strings.Contains(canonical, "payload.body.content") || strings.Contains(canonical, "message.text") {
t.Fatalf("canonical = %s, want query filter on payload.body.content only", canonical)
}
if !strings.Contains(canonical, "payload.body.openConversationId") || strings.Contains(canonical, "conversation_id") {
t.Fatalf("canonical = %s, want conversation_id alias mapped to payload.body.openConversationId", canonical)
}
}
func TestIdempotencyKeyUsesLocalIdentityKey(t *testing.T) {
left := Identity{LocalSubject: "refresh:left", ClientID: "client-1", SourceID: "open"}
right := Identity{LocalSubject: "refresh:right", ClientID: "client-1", SourceID: "open"}
ruleParam := map[string]any{"targetUid": "507971", "targetUidType": "staffId"}
leftKey := IdempotencyKey(left, EventSingleChat, "singleChat", ruleParam, "")
rightKey := IdempotencyKey(right, EventSingleChat, "singleChat", ruleParam, "")
if leftKey == rightKey {
t.Fatalf("idempotency key collapsed for different local subjects: %s", leftKey)
}
}
+172
View File
@@ -0,0 +1,172 @@
// 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 personal
import (
"encoding/json"
"errors"
"fmt"
"os"
"path/filepath"
"sort"
"time"
eventlock "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/lock"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
)
const (
StateFileName = "personal_subscriptions.json"
StateLockFileName = "personal_subscriptions.lock"
stateLockWaitTimeout = 5 * time.Second
stateLockRetryDelay = 25 * time.Millisecond
)
type RunState struct {
SubscribeID string `json:"subscribe_id"`
EventKey string `json:"event_key,omitempty"`
RuleType string `json:"rule_type,omitempty"`
ClientID string `json:"client_id,omitempty"`
SourceID string `json:"source_id,omitempty"`
IdentityHash string `json:"identity_hash,omitempty"`
CreatedAt time.Time `json:"created_at"`
}
func LoadRunStates(workDir string) ([]RunState, error) {
path := filepath.Join(workDir, StateFileName)
b, err := os.ReadFile(path)
if os.IsNotExist(err) {
return nil, nil
}
if err != nil {
return nil, err
}
var states []RunState
if err := json.Unmarshal(b, &states); err != nil {
return nil, err
}
sort.Slice(states, func(i, j int) bool {
return states[i].SubscribeID < states[j].SubscribeID
})
return states, nil
}
func UpsertRunState(workDir string, st RunState) error {
if st.SubscribeID == "" {
return nil
}
if st.CreatedAt.IsZero() {
st.CreatedAt = time.Now().UTC()
}
return withRunStateLock(workDir, stateLockWaitTimeout, func() error {
states, err := LoadRunStates(workDir)
if err != nil {
return err
}
replaced := false
for i := range states {
if states[i].SubscribeID == st.SubscribeID {
states[i] = st
replaced = true
break
}
}
if !replaced {
states = append(states, st)
}
return writeRunStates(workDir, states)
})
}
func RemoveRunStates(workDir string, subscribeIDs []string) error {
if len(subscribeIDs) == 0 {
return nil
}
remove := make(map[string]struct{}, len(subscribeIDs))
for _, id := range subscribeIDs {
if id != "" {
remove[id] = struct{}{}
}
}
return withRunStateLock(workDir, stateLockWaitTimeout, func() error {
states, err := LoadRunStates(workDir)
if err != nil {
return err
}
filtered := states[:0]
for _, st := range states {
if _, ok := remove[st.SubscribeID]; !ok {
filtered = append(filtered, st)
}
}
return writeRunStates(workDir, filtered)
})
}
func withRunStateLock(workDir string, wait time.Duration, fn func() error) error {
if err := os.MkdirAll(workDir, config.DirPerm); err != nil {
return err
}
lockPath := filepath.Join(workDir, StateLockFileName)
deadline := time.Now().Add(wait)
for {
held, err := eventlock.TryAcquire(lockPath)
if err == nil {
defer held.Close()
return fn()
}
if !errors.Is(err, eventlock.ErrBusy) {
return err
}
remaining := time.Until(deadline)
if remaining <= 0 {
return fmt.Errorf("personal event: timed out waiting for run-state lock after %s", wait)
}
delay := stateLockRetryDelay
if remaining < delay {
delay = remaining
}
time.Sleep(delay)
}
}
func writeRunStates(workDir string, states []RunState) error {
if err := os.MkdirAll(workDir, config.DirPerm); err != nil {
return err
}
sort.Slice(states, func(i, j int) bool {
return states[i].SubscribeID < states[j].SubscribeID
})
path := filepath.Join(workDir, StateFileName)
if len(states) == 0 {
if err := os.Remove(path); err != nil && !os.IsNotExist(err) {
return err
}
return nil
}
b, err := json.MarshalIndent(states, "", " ")
if err != nil {
return err
}
b = append(b, '\n')
tmp := path + ".tmp"
if err := os.WriteFile(tmp, b, config.FilePerm); err != nil {
return err
}
if err := os.Rename(tmp, path); err != nil {
_ = os.Remove(tmp)
return err
}
return nil
}
+141
View File
@@ -0,0 +1,141 @@
// 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 personal
import (
"encoding/json"
"fmt"
"os"
"path/filepath"
"strings"
"sync"
"testing"
"time"
eventlock "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/lock"
)
func TestRunStatesConcurrentUpsertPreservesEverySubscription(t *testing.T) {
workDir := t.TempDir()
const count = 64
var wg sync.WaitGroup
errs := make(chan error, count)
for i := 0; i < count; i++ {
wg.Add(1)
go func(i int) {
defer wg.Done()
errs <- UpsertRunState(workDir, RunState{SubscribeID: fmt.Sprintf("sub-%03d", i)})
}(i)
}
wg.Wait()
close(errs)
for err := range errs {
if err != nil {
t.Fatalf("UpsertRunState() error = %v", err)
}
}
states, err := LoadRunStates(workDir)
if err != nil {
t.Fatal(err)
}
if len(states) != count {
t.Fatalf("run states = %d, want %d", len(states), count)
}
for i, st := range states {
if want := fmt.Sprintf("sub-%03d", i); st.SubscribeID != want {
t.Fatalf("state %d = %q, want %q", i, st.SubscribeID, want)
}
}
}
func TestRunStatesConcurrentUpsertAndRemoveRemainConsistent(t *testing.T) {
workDir := t.TempDir()
for i := 0; i < 50; i++ {
if err := UpsertRunState(workDir, RunState{SubscribeID: fmt.Sprintf("base-%02d", i)}); err != nil {
t.Fatal(err)
}
}
var wg sync.WaitGroup
errs := make(chan error, 75)
for i := 0; i < 50; i++ {
wg.Add(1)
go func(i int) {
defer wg.Done()
errs <- UpsertRunState(workDir, RunState{SubscribeID: fmt.Sprintf("new-%02d", i)})
}(i)
}
for i := 0; i < 25; i++ {
wg.Add(1)
go func(i int) {
defer wg.Done()
errs <- RemoveRunStates(workDir, []string{fmt.Sprintf("base-%02d", i)})
}(i)
}
wg.Wait()
close(errs)
for err := range errs {
if err != nil {
t.Fatalf("run-state mutation error = %v", err)
}
}
raw, err := os.ReadFile(filepath.Join(workDir, StateFileName))
if err != nil {
t.Fatal(err)
}
var decoded []RunState
if err := json.Unmarshal(raw, &decoded); err != nil {
t.Fatalf("state file is invalid JSON: %v", err)
}
if len(decoded) != 75 {
t.Fatalf("run states = %d, want 75", len(decoded))
}
present := make(map[string]bool, len(decoded))
for _, st := range decoded {
present[st.SubscribeID] = true
}
for i := 25; i < 50; i++ {
if !present[fmt.Sprintf("base-%02d", i)] {
t.Fatalf("unremoved base subscription %d was lost", i)
}
}
for i := 0; i < 50; i++ {
if !present[fmt.Sprintf("new-%02d", i)] {
t.Fatalf("new subscription %d was lost", i)
}
}
}
func TestWithRunStateLockTimesOut(t *testing.T) {
workDir := t.TempDir()
held, err := eventlock.TryAcquire(filepath.Join(workDir, StateLockFileName))
if err != nil {
t.Fatal(err)
}
defer held.Close()
called := false
err = withRunStateLock(workDir, 20*time.Millisecond, func() error {
called = true
return nil
})
if err == nil || !strings.Contains(err.Error(), "timed out waiting for run-state lock") {
t.Fatalf("withRunStateLock() error = %v, want timeout", err)
}
if called {
t.Fatal("locked callback was called")
}
}
+23
View File
@@ -0,0 +1,23 @@
// 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 process provides cross-platform process-existence checks. Used by
// the bus daemon's stale-lock recovery path: bus.lock holds the daemon's PID;
// on startup a competing process reads that PID and calls Alive() to decide
// whether to abort (the holder is still running) or steal the lock (orphan).
//
// On Unix this is the standard signal-0 trick (syscall.Kill(pid, 0)) which
// errors with ESRCH if the process is gone. On Windows it uses OpenProcess
// with PROCESS_QUERY_LIMITED_INFORMATION and GetExitCodeProcess to detect
// STILL_ACTIVE; this avoids the Unix-only assumption in plan §1.
package process
@@ -0,0 +1,63 @@
// 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 process
import (
"os"
"os/exec"
"runtime"
"testing"
)
func TestAlive_SelfPID(t *testing.T) {
if !Alive(os.Getpid()) {
t.Fatal("Alive(self pid) should be true")
}
}
func TestAlive_NonPositivePID(t *testing.T) {
for _, pid := range []int{0, -1, -42} {
if Alive(pid) {
t.Fatalf("Alive(%d) should be false (non-positive)", pid)
}
}
}
func TestAlive_DeadChild(t *testing.T) {
// Spawn a short-lived child and wait for it to exit, then verify Alive
// returns false. We use sleep 0 which exits immediately on both Unix
// and Windows (Windows treats `sleep 0` as a built-in via cmd.exe but
// to avoid Windows shell quirks we fall back to `cmd /c exit 0`).
var cmd *exec.Cmd
if runtime.GOOS == "windows" {
cmd = exec.Command("cmd", "/c", "exit", "0")
} else {
cmd = exec.Command("sh", "-c", "exit 0")
}
if err := cmd.Start(); err != nil {
t.Fatalf("start child: %v", err)
}
pid := cmd.Process.Pid
if err := cmd.Wait(); err != nil {
t.Fatalf("wait child: %v", err)
}
// On Windows, even after Wait, the kernel may keep the handle table
// entry briefly; GetExitCodeProcess should still report non-STILL_ACTIVE
// so Alive returns false. On Unix, the parent has reaped via Wait so
// signal 0 returns ESRCH.
if Alive(pid) {
t.Fatalf("Alive(%d) should be false after child exit", pid)
}
}
@@ -0,0 +1,52 @@
// 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.
//go:build !windows
package process
import (
"errors"
"syscall"
)
// Alive reports whether a process with the given PID is currently running.
// Returns false for non-positive PIDs without performing any syscall.
//
// On Unix this uses signal 0 (syscall.Kill(pid, 0)), which performs the
// normal permission/existence check without delivering a signal:
// - nil error → process exists and we have permission
// - syscall.EPERM → process exists but owned by another user
// (still counts as alive for our purpose)
// - syscall.ESRCH → no such process
//
// Any other error is treated as "alive" defensively, so we never steal a
// lock based on an ambiguous syscall failure.
func Alive(pid int) bool {
if pid <= 0 {
return false
}
err := syscall.Kill(pid, 0)
if err == nil {
return true
}
if errors.Is(err, syscall.EPERM) {
// Process exists, we just can't signal it.
return true
}
if errors.Is(err, syscall.ESRCH) {
return false
}
// Unknown error: be conservative, treat as alive.
return true
}
@@ -0,0 +1,60 @@
// 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.
//go:build windows
package process
import (
"golang.org/x/sys/windows"
)
// stillActive is the WinAPI sentinel returned by GetExitCodeProcess for
// running processes (STILL_ACTIVE = 259).
const stillActive uint32 = 259
// Alive reports whether a process with the given PID is currently running.
// Returns false for non-positive PIDs without performing any syscall.
//
// On Windows this opens the process with PROCESS_QUERY_LIMITED_INFORMATION
// (least privilege required) and calls GetExitCodeProcess. A return code of
// STILL_ACTIVE means the process is running; any other value means it has
// exited. If OpenProcess fails with ERROR_INVALID_PARAMETER the PID is
// unknown to the system (process never existed or has been recycled).
//
// Any other unexpected error is treated as "alive" defensively, so we never
// steal a lock based on an ambiguous syscall failure.
func Alive(pid int) bool {
if pid <= 0 {
return false
}
h, err := windows.OpenProcess(windows.PROCESS_QUERY_LIMITED_INFORMATION, false, uint32(pid))
if err != nil {
if err == windows.ERROR_INVALID_PARAMETER {
return false
}
// Access denied means the process exists but we can't query it.
if err == windows.ERROR_ACCESS_DENIED {
return true
}
// Other errors: be conservative, treat as alive.
return true
}
defer windows.CloseHandle(h)
var code uint32
if err := windows.GetExitCodeProcess(h, &code); err != nil {
return true
}
return code == stillActive
}
+54
View File
@@ -0,0 +1,54 @@
// 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 registry
// CatchAllEventTypes returns the default event-type wildcards passed in
// the Hello frame when `dws event consume` is invoked without
// --event-types.
//
// v1 returns nil = "bus-wide catch-all" — the bus already subscribes
// upstream to everything the open-platform UI ticks. The Hub's matcher
// treats nil/empty as match-everything, so the consumer gets every event
// the open platform pushes without needing an explicit list.
//
// Why we don't ship a curated default (yet): the exact event_type
// strings DingTalk emits over Stream are not yet confirmed by a P0
// run against a real app (escape-hatch row #3). Shipping a curated
// list that differs by one character from what DingTalk actually sends
// would silently filter out events users expect to see. The bus-wide
// catch-all behaviour is conservative — too inclusive rather than too
// exclusive — and avoids that failure mode.
//
// Known/expected event type strings (from open-platform docs; pending
// P0 SDK confirmation — DO NOT enable until verified):
//
// im.message.receive_v1 receive any IM message
// im.message.read_v1 message read receipt
// im.message.reaction.created_v1 reaction added
// im.message.reaction.deleted_v1 reaction removed
// im.chat.member.bot.added_v1 bot added to a chat
// im.chat.member.bot.deleted_v1 bot removed from a chat
// im.chat.member.user.added_v1 user added to a chat
// im.chat.disbanded_v1 chat disbanded
// contact.user.created_v3 / updated / deleted
// contact.department.created_v3 / updated / deleted
// cal.event.created_v1 / updated / deleted
// approval.instance.status_changed
// approval.task.created
// attendance.check_v1
//
// Once verified, switch CatchAllEventTypes to return the slice above (and
// document the change so users can override with --event-types '*' to
// regain literal "everything the open platform sends").
func CatchAllEventTypes() []string { return nil }
@@ -0,0 +1,243 @@
// 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 registry
import (
"encoding/json"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/transport"
)
// CompactView is the JSON-encoded shape `dws event consume --compact`
// writes per event. Implementations populate the common header fields
// then layer event-type-specific semantic fields on top via "extra".
type CompactView struct {
// Type echoes EventType for routing/filtering at the agent layer.
Type string `json:"type"`
// EventID is the SDK-assigned identifier (DataFrameHeader "eventId").
EventID string `json:"event_id,omitempty"`
// Timestamp is the event-born-time milliseconds value, surfaced under a
// shorter name agents typically expect.
Timestamp int64 `json:"timestamp,omitempty"`
// CorpID propagates the SDK's eventCorpId when present.
CorpID string `json:"corp_id,omitempty"`
// AppID propagates eventUnifiedAppId when present.
AppID string `json:"app_id,omitempty"`
// Extra holds the event-type-specific flattened fields (e.g.
// chat_id/sender_id for IM messages). Marshalled inline.
Extra map[string]any `json:"-"`
}
// MarshalJSON serialises CompactView with the Extra fields lifted to the
// top level, dropping the wrapping "extra" key. This produces the flat
// map output that agents prefer.
func (v CompactView) MarshalJSON() ([]byte, error) {
out := make(map[string]any, 6+len(v.Extra))
out["type"] = v.Type
if v.EventID != "" {
out["event_id"] = v.EventID
}
if v.Timestamp != 0 {
out["timestamp"] = v.Timestamp
}
if v.CorpID != "" {
out["corp_id"] = v.CorpID
}
if v.AppID != "" {
out["app_id"] = v.AppID
}
for k, val := range v.Extra {
// Header fields win over Extra on collision — we never let a
// processor accidentally overwrite the documented top-level shape.
if _, taken := out[k]; taken {
continue
}
out[k] = val
}
return json.Marshal(out)
}
// Processor transforms one Event frame into a CompactView. Implementations
// MUST NOT modify the input. Returning a zero CompactView is allowed and
// falls back to the generic processor at the caller's discretion.
type Processor func(ev transport.Event) CompactView
// processors holds the registered specialised processors. Lookup is by
// exact EventType match; wildcard support is intentionally absent because
// compact rendering is identity-specific (per event_type schema).
//
// To register more event types, add entries here. Each processor receives
// the full transport.Event and returns a CompactView with the per-type
// semantic fields lifted into Extra. See GenericProcessor for the
// fall-through behaviour applied to unregistered types.
var processors = map[string]Processor{
// IM
"im.message.receive_v1": compactIMMessage,
"im.message.read_v1": compactIMMessageRead,
// Approval
"approval.instance.status_changed": compactApprovalInstance,
"approval.task.created": compactApprovalTask,
// Contact
"contact.user.created_v3": compactContactUser,
"contact.user.updated_v3": compactContactUser,
"contact.user.deleted_v3": compactContactUser,
// Calendar
"cal.event.created_v1": compactCalendarEvent,
"cal.event.updated_v1": compactCalendarEvent,
"cal.event.deleted_v1": compactCalendarEvent,
// Attendance
"attendance.check_v1": compactAttendanceCheck,
}
// LookupProcessor returns the registered processor for eventType, or
// GenericProcessor when no specialised one exists.
func LookupProcessor(eventType string) Processor {
if p, ok := processors[eventType]; ok {
return p
}
return GenericProcessor
}
// GenericProcessor is the fallback used when no specialised processor is
// registered for the event type. It surfaces the 5 SDK header fields and
// tries to parse the JSON payload — on parse failure it embeds the raw
// payload string under "data" so nothing is lost.
func GenericProcessor(ev transport.Event) CompactView {
v := CompactView{
Type: ev.EventType,
EventID: ev.EventID,
Timestamp: ev.EventBornTime,
CorpID: ev.EventCorpID,
AppID: ev.EventUnifiedAppID,
Extra: map[string]any{},
}
var parsed map[string]any
if err := json.Unmarshal([]byte(ev.Data), &parsed); err == nil {
for k, val := range parsed {
v.Extra[k] = val
}
} else if ev.Data != "" {
v.Extra["data"] = ev.Data
}
return v
}
// liftFromNested copies named keys from a nested map at top level on the
// CompactView. No-op when nested is missing. Used by the per-type
// processors to lift "sender.sender_id.open_id" style nested fields.
func liftFromNested(v *CompactView, nestedKey string, keys ...string) {
nested, ok := v.Extra[nestedKey].(map[string]any)
if !ok {
return
}
for _, k := range keys {
if val, ok := nested[k]; ok {
if _, taken := v.Extra[k]; !taken {
v.Extra[k] = val
}
}
}
}
// compactIMMessage is the specialised processor for IM message receipts
// (event_type: im.message.receive_v1).
//
// The SDK payload for IM message events is a nested JSON document. The
// canonical agent fields we expose are: message_id, chat_id, chat_type,
// message_type, content (already a string for text messages), sender_id.
// Any field DingTalk adds later flows through via the generic merge —
// we never drop unknown fields, we just promote the common ones.
func compactIMMessage(ev transport.Event) CompactView {
v := GenericProcessor(ev)
liftFromNested(&v, "message", "message_id", "chat_id", "chat_type", "message_type", "content")
// Sender open_id may be doubly nested: sender.sender_id.open_id.
if sender, ok := v.Extra["sender"].(map[string]any); ok {
if sid, ok := sender["sender_id"].(map[string]any); ok {
if open, ok := sid["open_id"]; ok {
if _, taken := v.Extra["sender_id"]; !taken {
v.Extra["sender_id"] = open
}
}
}
}
return v
}
// compactIMMessageRead lifts message read receipts (open_id of the reader
// + the message_ids they acknowledged).
func compactIMMessageRead(ev transport.Event) CompactView {
v := GenericProcessor(ev)
liftFromNested(&v, "reader", "open_id", "read_time")
liftFromNested(&v, "context", "open_message_ids")
return v
}
// compactApprovalInstance promotes the approval-instance-level identity
// fields (instance_id / status / business_id / process_code) and the
// requestor name.
func compactApprovalInstance(ev transport.Event) CompactView {
v := GenericProcessor(ev)
liftFromNested(&v, "instance", "instance_id", "process_code", "business_id", "title")
liftFromNested(&v, "status_change", "from_status", "to_status", "operate_time")
if op, ok := v.Extra["operator"].(map[string]any); ok {
if id, ok := op["userid"]; ok {
v.Extra["operator_userid"] = id
}
}
return v
}
// compactApprovalTask lifts task assignment info (task_id, assignee user_id).
func compactApprovalTask(ev transport.Event) CompactView {
v := GenericProcessor(ev)
liftFromNested(&v, "task", "task_id", "instance_id", "status", "action_type")
if assignee, ok := v.Extra["assignee"].(map[string]any); ok {
if id, ok := assignee["userid"]; ok {
v.Extra["assignee_userid"] = id
}
}
return v
}
// compactContactUser lifts the user identity for user lifecycle events.
func compactContactUser(ev transport.Event) CompactView {
v := GenericProcessor(ev)
liftFromNested(&v, "user", "open_id", "userid", "name", "active")
if v.Extra["open_id"] == nil {
// older payload may put open_id at top level under "open_ids" array
if arr, ok := v.Extra["open_ids"].([]any); ok && len(arr) > 0 {
v.Extra["open_id"] = arr[0]
}
}
return v
}
// compactCalendarEvent lifts calendar event identity + organiser.
func compactCalendarEvent(ev transport.Event) CompactView {
v := GenericProcessor(ev)
liftFromNested(&v, "event", "event_id", "calendar_id", "title", "start_time", "end_time", "location")
if org, ok := v.Extra["organizer"].(map[string]any); ok {
if id, ok := org["userid"]; ok {
v.Extra["organizer_userid"] = id
}
}
return v
}
// compactAttendanceCheck lifts the punch info (userid, check_time, location).
func compactAttendanceCheck(ev transport.Event) CompactView {
v := GenericProcessor(ev)
liftFromNested(&v, "punch", "userid", "check_time", "check_type", "location_method")
return v
}
@@ -0,0 +1,238 @@
// 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 registry
import (
"encoding/json"
"testing"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/transport"
)
func mustJSON(t *testing.T, v any) map[string]any {
t.Helper()
b, err := json.Marshal(v)
if err != nil {
t.Fatalf("marshal: %v", err)
}
var out map[string]any
if err := json.Unmarshal(b, &out); err != nil {
t.Fatalf("unmarshal: %v", err)
}
return out
}
func TestGenericProcessor_HeaderFieldsPromoted(t *testing.T) {
ev := transport.Event{
EventID: "ev_1",
EventBornTime: 1700000000123,
EventCorpID: "corp_x",
EventType: "approval.task",
EventUnifiedAppID: "app_y",
Data: `{"task_id":"t1","status":"approved"}`,
}
cv := GenericProcessor(ev)
out := mustJSON(t, cv)
wants := map[string]any{
"type": "approval.task",
"event_id": "ev_1",
"corp_id": "corp_x",
"app_id": "app_y",
"task_id": "t1",
"status": "approved",
"timestamp": float64(1700000000123),
}
for k, want := range wants {
if got := out[k]; got != want {
t.Errorf("compact[%q] = %v, want %v", k, got, want)
}
}
}
func TestGenericProcessor_PayloadParseFailEmbedsRaw(t *testing.T) {
ev := transport.Event{
EventType: "x",
Data: "not json",
}
out := mustJSON(t, GenericProcessor(ev))
if out["data"] != "not json" {
t.Fatalf("non-JSON Data should land under 'data', got: %v", out["data"])
}
}
func TestGenericProcessor_EmptyDataNoExtraData(t *testing.T) {
ev := transport.Event{EventType: "x", Data: ""}
out := mustJSON(t, GenericProcessor(ev))
if _, ok := out["data"]; ok {
t.Errorf("empty Data should not produce a 'data' key; got: %v", out)
}
}
func TestCompactView_HeaderFieldsWinOverExtra(t *testing.T) {
// A malicious or buggy processor that tries to overwrite "type" must
// be ignored — header fields are the documented contract.
cv := CompactView{
Type: "the_real_type",
Extra: map[string]any{"type": "fake"},
}
out := mustJSON(t, cv)
if out["type"] != "the_real_type" {
t.Fatalf("Extra 'type' must not override header; got %v", out["type"])
}
}
func TestLookupProcessor_FallsBackToGeneric(t *testing.T) {
p := LookupProcessor("never.registered.event")
if p == nil {
t.Fatal("LookupProcessor should never return nil")
}
cv := p(transport.Event{EventType: "never.registered.event", Data: "{}"})
if cv.Type != "never.registered.event" {
t.Fatalf("fallback processor type = %q", cv.Type)
}
}
func TestCompactIMMessage_LiftsNestedFields(t *testing.T) {
// Simulate the nested SDK shape: {"message": {...}, "sender": {...}}.
ev := transport.Event{
EventType: "im.message.receive_v1",
EventID: "ev_abc",
Data: `{
"message": {
"message_id": "om_x",
"chat_id": "oc_y",
"chat_type": "p2p",
"message_type": "text",
"content": "{\"text\":\"hello\"}"
},
"sender": {
"sender_id": {"open_id": "ou_z"},
"sender_type": "user"
}
}`,
}
out := mustJSON(t, compactIMMessage(ev))
if out["message_id"] != "om_x" {
t.Errorf("message_id = %v", out["message_id"])
}
if out["chat_id"] != "oc_y" {
t.Errorf("chat_id = %v", out["chat_id"])
}
if out["chat_type"] != "p2p" {
t.Errorf("chat_type = %v", out["chat_type"])
}
if out["sender_id"] != "ou_z" {
t.Errorf("sender_id = %v", out["sender_id"])
}
// Original nested fields should still be present too (generic merge).
if _, ok := out["message"]; !ok {
t.Error("nested 'message' should also be retained for downstream consumers")
}
}
func TestCompactIMMessage_FlatPayloadAlsoWorks(t *testing.T) {
// If the SDK ever flattens these to the top level, we should still
// surface them — generic merge already does this.
ev := transport.Event{
EventType: "im.message.receive_v1",
Data: `{"message_id":"om_x","chat_id":"oc_y","content":"hi"}`,
}
out := mustJSON(t, compactIMMessage(ev))
if out["message_id"] != "om_x" || out["chat_id"] != "oc_y" || out["content"] != "hi" {
t.Fatalf("flat fields lost: %+v", out)
}
}
func TestCatchAllEventTypes_EmptyForV1(t *testing.T) {
if got := CatchAllEventTypes(); got != nil {
t.Fatalf("CatchAllEventTypes v1 should be nil, got %v", got)
}
}
func TestSpecialisedProcessorsRegistered(t *testing.T) {
// Sentinel: every specialised event type should resolve to a non-
// GenericProcessor function. If a P7 entry was accidentally dropped
// from the registry map this catches it before runtime users do.
wantRegistered := []string{
"im.message.receive_v1",
"im.message.read_v1",
"approval.instance.status_changed",
"approval.task.created",
"contact.user.created_v3",
"contact.user.updated_v3",
"contact.user.deleted_v3",
"cal.event.created_v1",
"cal.event.updated_v1",
"cal.event.deleted_v1",
"attendance.check_v1",
}
for _, et := range wantRegistered {
if _, ok := processors[et]; !ok {
t.Errorf("event type %q should have a registered processor", et)
}
}
}
func TestCompactApprovalInstance(t *testing.T) {
ev := transport.Event{
EventType: "approval.instance.status_changed",
Data: `{
"instance": {
"instance_id": "inst_123",
"process_code": "proc_abc",
"title": "请假申请"
},
"status_change": {
"from_status": "PENDING",
"to_status": "APPROVED"
},
"operator": {"userid": "user_xyz"}
}`,
}
out := mustJSON(t, compactApprovalInstance(ev))
for k, want := range map[string]any{
"instance_id": "inst_123",
"process_code": "proc_abc",
"title": "请假申请",
"from_status": "PENDING",
"to_status": "APPROVED",
"operator_userid": "user_xyz",
} {
if out[k] != want {
t.Errorf("approval[%q] = %v, want %v", k, out[k], want)
}
}
}
func TestCompactContactUser(t *testing.T) {
ev := transport.Event{
EventType: "contact.user.created_v3",
Data: `{"user":{"open_id":"ou_x","userid":"u_x","name":"Alice","active":true}}`,
}
out := mustJSON(t, compactContactUser(ev))
if out["open_id"] != "ou_x" || out["userid"] != "u_x" || out["name"] != "Alice" {
t.Errorf("contact user fields missing: %+v", out)
}
}
func TestCompactCalendarEvent(t *testing.T) {
ev := transport.Event{
EventType: "cal.event.created_v1",
Data: `{"event":{"event_id":"e1","title":"Standup","start_time":"2026-01-01T09:00:00Z"},"organizer":{"userid":"u_x"}}`,
}
out := mustJSON(t, compactCalendarEvent(ev))
if out["event_id"] != "e1" || out["title"] != "Standup" || out["organizer_userid"] != "u_x" {
t.Errorf("calendar event fields missing: %+v", out)
}
}
+28
View File
@@ -0,0 +1,28 @@
// 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 registry holds two cross-cutting reference data sets used by
// daemon, consume, and the cobra command layer:
//
// - CatchAll: the default event_type list passed to `consume`'s Hello
// when --event-types is omitted. v1 (P4) ships an empty list so the
// consumer falls back to bus-wide catch-all; P7 fills in the curated
// DingTalk default set after P0 SDK verification confirms the exact
// event_type string values.
//
// - CompactProcessors: per-event_type formatters that flatten the SDK
// RawEvent into an agent-friendly map[string]any. A generic processor
// handles unknown types by surfacing the top-level header fields plus
// a parsed payload; specialised processors (im.message.* etc.) extract
// semantically meaningful fields.
package registry
+178
View File
@@ -0,0 +1,178 @@
// 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 source wraps the open-dingtalk Stream SDK and exposes a small
// blocking Start interface plus a connection state machine. The SDK is the
// only place in the bus that talks to the cloud; the rest of the bus stays
// vendor-agnostic and would be drop-in replaceable for a different Stream
// provider.
package source
import (
"context"
"errors"
"time"
dwsevent "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event"
"github.com/open-dingtalk/dingtalk-stream-sdk-go/client"
"github.com/open-dingtalk/dingtalk-stream-sdk-go/event"
"github.com/open-dingtalk/dingtalk-stream-sdk-go/payload"
)
// Config carries the credentials and behavioural knobs needed to construct a
// DingtalkSource. ClientID/ClientSecret are required for app-credential SDK
// mode. Portal ticket normal mode uses portal-side managed credentials, so it
// does not require local app credentials.
type Config struct {
ClientID string
ClientSecret string
// PortalTicket switches the source to the portal-managed user Stream
// ticket flow. When set, Start fetches endpoint+ticket over HTTP and
// dials the returned WebSocket directly instead of asking the SDK to
// open an app-credential connection.
PortalTicket *PortalTicketConfig
// Now is injected for tests. Defaults to time.Now when nil.
Now func() time.Time
}
// SourceOptions is the functional-option type reserved for future overlay
// extensions (e.g. inject trace IDs into emit, swap the underlying client
// for a fake). v1 takes no options; the type exists so callers can write
// `New(cfg)` today and `New(cfg, WithFoo())` tomorrow without an API break.
// See plan §7 Edition 扩展点.
type SourceOption func(*sourceConfig)
type sourceConfig struct {
// reserved for v2 hooks (pre-emit interceptor, etc.)
}
// DingtalkSource is the cloud-side adapter. It owns one StreamClient and one
// state Machine; lifecycle is bounded by the context passed to Start.
type DingtalkSource struct {
cfg Config
machine *Machine
cli *client.StreamClient
}
// New constructs a DingtalkSource. Returns an error if required Config
// fields are missing — keep the boundary tight so misconfiguration fails
// loudly rather than at first-event time.
func New(cfg Config, _ ...SourceOption) (*DingtalkSource, error) {
if cfg.PortalTicket != nil {
if err := cfg.PortalTicket.Valid(); err != nil {
return nil, err
}
if cfg.PortalTicket.normalizedMode() == PortalTicketModeCustom && cfg.ClientID == "" {
return nil, errors.New("source: ClientID is required")
}
} else {
if cfg.ClientID == "" {
return nil, errors.New("source: ClientID is required")
}
if cfg.ClientSecret == "" {
return nil, errors.New("source: ClientSecret is required")
}
}
if cfg.Now == nil {
cfg.Now = time.Now
}
m := NewMachine()
m.now = cfg.Now
return &DingtalkSource{cfg: cfg, machine: m}, nil
}
// State returns the current Snapshot of the connection state machine.
func (s *DingtalkSource) State() Snapshot { return s.machine.Snapshot() }
// Start opens the Stream WebSocket and blocks until ctx is cancelled or a
// fatal SDK error occurs. Events are delivered to emit synchronously from
// the SDK callback goroutine — emit MUST return immediately (the SDK's
// processLoop is single-goroutine; see plan invariant #1 and P0 §1 row 2).
//
// Blocking semantics: the underlying StreamClient.Start returns as soon as
// dial succeeds (it spawns its own processLoop goroutine), so we wait on
// ctx.Done() before returning. cli.Close() is always called on exit; ctx.Err
// is returned for context cancellation, otherwise the underlying error.
func (s *DingtalkSource) Start(ctx context.Context, emit dwsevent.EmitFn) error {
if emit == nil {
return errors.New("source: emit is required")
}
if s.cli != nil {
return errors.New("source: Start called twice")
}
if s.cfg.PortalTicket != nil {
return s.startPortalTicket(ctx, emit)
}
options := []client.ClientOption{
client.WithAppCredential(client.NewAppCredentialConfig(s.cfg.ClientID, s.cfg.ClientSecret)),
}
s.cli = client.NewStreamClient(options...)
s.cli.RegisterAllEventRouter(s.makeHandler(emit))
s.machine.OnConnecting()
if err := s.cli.Start(ctx); err != nil {
s.machine.OnStopped()
return err
}
s.machine.OnConnected()
<-ctx.Done()
s.cli.Close()
s.machine.OnStopped()
return ctx.Err()
}
// makeHandler returns the IFrameHandler closure the SDK invokes for every
// inbound event. The closure:
// 1. parses the EventHeader from the raw DataFrame,
// 2. builds a RawEvent with all 5 SDK header fields + payload + receive
// time + full header map (passthrough),
// 3. calls emit (non-blocking by contract),
// 4. updates the state machine,
// 5. returns SUCCESS (LATER is reserved for future explicit retry policy;
// v1 always Success and relies on dedup, see P0 §1 row 6).
func (s *DingtalkSource) makeHandler(emit dwsevent.EmitFn) func(context.Context, *payload.DataFrame) (*payload.DataFrameResponse, error) {
return func(_ context.Context, df *payload.DataFrame) (*payload.DataFrameResponse, error) {
hdr := event.NewEventHeaderFromDataFrame(df)
raw := &dwsevent.RawEvent{
EventID: hdr.EventId,
EventBornTime: hdr.EventBornTime,
EventCorpID: hdr.EventCorpId,
EventType: hdr.EventType,
EventUnifiedAppID: hdr.EventUnifiedAppId,
Data: df.Data,
Headers: copyHeaders(df.Headers),
ReceivedAt: s.cfg.Now().UTC(),
}
emit(raw)
s.machine.OnEvent()
resp := payload.NewSuccessDataFrameResponse()
if err := resp.SetJson(event.NewEventProcessResultSuccess()); err != nil {
return nil, err
}
return resp, nil
}
}
func copyHeaders(h payload.DataFrameHeader) map[string]string {
if len(h) == 0 {
return nil
}
out := make(map[string]string, len(h))
for k, v := range h {
out[k] = v
}
return out
}
+235
View File
@@ -0,0 +1,235 @@
// 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 source
import (
"context"
"encoding/json"
"net/http"
"net/http/httptest"
"testing"
"time"
dwsevent "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event"
"github.com/open-dingtalk/dingtalk-stream-sdk-go/event"
"github.com/open-dingtalk/dingtalk-stream-sdk-go/payload"
)
func TestNew_RequiresClientID(t *testing.T) {
if _, err := New(Config{ClientSecret: "secret"}); err == nil {
t.Fatal("expected error when ClientID empty")
}
}
func TestNew_RequiresClientSecret(t *testing.T) {
if _, err := New(Config{ClientID: "id"}); err == nil {
t.Fatal("expected error when ClientSecret empty")
}
}
func TestNew_DefaultsNow(t *testing.T) {
s, err := New(Config{ClientID: "id", ClientSecret: "secret"})
if err != nil {
t.Fatalf("New: %v", err)
}
if s.cfg.Now == nil {
t.Fatal("Now should default to time.Now")
}
}
func TestNew_AcceptsPortalTicketNormalWithoutClientSecret(t *testing.T) {
if _, err := New(Config{
PortalTicket: &PortalTicketConfig{
TicketURL: "https://example.com/stream/connections/ticket",
AccessToken: "token",
SourceID: "pre_open_source",
Mode: "normal",
},
}); err != nil {
t.Fatalf("New: %v", err)
}
}
func TestNew_RejectsPortalTicketCustomWithoutSecret(t *testing.T) {
if _, err := New(Config{
ClientID: "id",
PortalTicket: &PortalTicketConfig{
TicketURL: "https://example.com/stream/connections/ticket",
AccessToken: "token",
SourceID: "pre_open_source",
Mode: "custom",
ClientID: "custom_client",
},
}); err == nil {
t.Fatal("expected error when custom portal ticket secret is empty")
}
}
func TestRequestPortalTicketCustomBody(t *testing.T) {
var got struct {
SourceID string `json:"sourceId"`
ChannelType string `json:"channelType"`
Mode string `json:"mode"`
ClientID string `json:"clientId"`
ClientSecret string `json:"clientSecret"`
}
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if token := r.Header.Get("x-user-access-token"); token != "token-123" {
t.Fatalf("x-user-access-token = %q", token)
}
if err := json.NewDecoder(r.Body).Decode(&got); err != nil {
t.Fatalf("decode request: %v", err)
}
_ = json.NewEncoder(w).Encode(map[string]any{
"success": true,
"result": map[string]string{
"endpoint": "wss://pre-wss-open-connection.dingtalk.com/connect",
"ticket": "ticket-123",
},
})
}))
defer srv.Close()
ticket, err := requestPortalTicket(context.Background(), &PortalTicketConfig{
TicketURL: srv.URL,
AccessToken: "token-123",
SourceID: "pre_open_source",
Mode: "custom",
ClientID: "custom_client",
ClientSecret: "custom_secret",
})
if err != nil {
t.Fatalf("requestPortalTicket: %v", err)
}
if ticket.Endpoint == "" || ticket.Ticket == "" {
t.Fatalf("ticket = %#v", ticket)
}
if got.SourceID != "pre_open_source" || got.ChannelType != "pre_open_source" {
t.Fatalf("source fields = %#v", got)
}
if got.Mode != "custom" || got.ClientID != "custom_client" || got.ClientSecret != "custom_secret" {
t.Fatalf("custom fields = %#v", got)
}
}
func TestState_InitiallyDisconnected(t *testing.T) {
s, _ := New(Config{ClientID: "id", ClientSecret: "secret"})
if got := s.State().State; got != StateDisconnected {
t.Fatalf("initial state = %s, want %s", got, StateDisconnected)
}
}
func TestMakeHandler_TranslatesDataFrameToRawEvent(t *testing.T) {
fixed := time.Date(2026, 5, 28, 12, 0, 0, 0, time.UTC)
s, err := New(Config{
ClientID: "id",
ClientSecret: "secret",
Now: func() time.Time { return fixed },
})
if err != nil {
t.Fatalf("New: %v", err)
}
var got *dwsevent.RawEvent
handler := s.makeHandler(func(e *dwsevent.RawEvent) { got = e })
df := &payload.DataFrame{
SpecVersion: "1.0",
Type: "EVENT",
Time: 1700000000,
Headers: payload.DataFrameHeader{
"eventId": "ev_abc",
"eventBornTime": "1700000000123",
"eventCorpId": "corp_x",
"eventType": "im.message.receive_v1",
"eventUnifiedAppId": "app_y",
"extra": "passthrough",
},
Data: `{"chat":"hello"}`,
}
resp, err := handler(context.Background(), df)
if err != nil {
t.Fatalf("handler: %v", err)
}
if resp == nil {
t.Fatal("handler returned nil resp")
}
if got == nil {
t.Fatal("emit was not called")
}
if got.EventID != "ev_abc" {
t.Errorf("EventID = %q, want %q", got.EventID, "ev_abc")
}
if got.EventType != "im.message.receive_v1" {
t.Errorf("EventType = %q", got.EventType)
}
if got.EventCorpID != "corp_x" {
t.Errorf("EventCorpID = %q", got.EventCorpID)
}
if got.EventUnifiedAppID != "app_y" {
t.Errorf("EventUnifiedAppID = %q", got.EventUnifiedAppID)
}
if got.EventBornTime != 1700000000123 {
t.Errorf("EventBornTime = %d", got.EventBornTime)
}
if got.Data != `{"chat":"hello"}` {
t.Errorf("Data = %q", got.Data)
}
if !got.ReceivedAt.Equal(fixed) {
t.Errorf("ReceivedAt = %v, want %v (injected Now)", got.ReceivedAt, fixed)
}
if got.Headers["extra"] != "passthrough" {
t.Error("passthrough header lost")
}
// Also verify the connection state machine ticked: handler should have
// called OnEvent, but since we never called OnConnected the visible
// state remains disconnected. We only verify the lastEventAt timestamp.
snap := s.State()
if snap.LastEventAt.IsZero() {
t.Error("OnEvent should have updated lastEventAt")
}
}
func TestStart_RejectsNilEmit(t *testing.T) {
s, _ := New(Config{ClientID: "id", ClientSecret: "secret"})
err := s.Start(context.Background(), nil)
if err == nil || err.Error() == "" {
t.Fatalf("Start with nil emit should error, got %v", err)
}
}
func TestMakeHandler_NilHeaders(t *testing.T) {
s, _ := New(Config{ClientID: "id", ClientSecret: "secret"})
handler := s.makeHandler(func(*dwsevent.RawEvent) {})
df := &payload.DataFrame{
// no Headers at all
Data: "{}",
}
resp, err := handler(context.Background(), df)
if err != nil {
t.Fatalf("handler: %v", err)
}
if resp == nil {
t.Fatal("nil resp")
}
}
// Compile-time guard: the EventProcessResult we return must be the SDK's
// EventProcessResultSuccess type. If the SDK ever changes the shape, this
// will fail to compile.
var _ = event.NewEventProcessResultSuccess
+627
View File
@@ -0,0 +1,627 @@
// 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 source
import (
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"io"
"log/slog"
"net/http"
"net/url"
"strings"
"sync/atomic"
"time"
dwsevent "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
"github.com/gorilla/websocket"
"github.com/open-dingtalk/dingtalk-stream-sdk-go/event"
"github.com/open-dingtalk/dingtalk-stream-sdk-go/payload"
)
const (
personalRawDebugPayloadLimit = 8192
personalReconnectMinBackoff = time.Second
personalReconnectMaxBackoff = 30 * time.Second
)
type PersonalConfig struct {
AccessToken string
ClientID string
ClientSecret string
SourceID string
TicketURL string
TicketMode string
HTTPClient *http.Client
WebSocketDialer *websocket.Dialer
Now func() time.Time
ReconnectMin time.Duration
ReconnectMax time.Duration
}
type PersonalSource struct {
cfg PersonalConfig
machine *Machine
started atomic.Bool
}
type retryablePersonalError struct {
err error
}
func (e *retryablePersonalError) Error() string { return e.err.Error() }
func (e *retryablePersonalError) Unwrap() error { return e.err }
type ticketResponse struct {
Endpoint string `json:"endpoint"`
Ticket string `json:"ticket"`
}
func NewPersonal(cfg PersonalConfig) (*PersonalSource, error) {
if strings.TrimSpace(cfg.AccessToken) == "" {
return nil, errors.New("personal source: AccessToken is required")
}
if strings.TrimSpace(cfg.ClientID) == "" {
return nil, errors.New("personal source: ClientID is required")
}
if strings.TrimSpace(cfg.SourceID) == "" {
return nil, errors.New("personal source: SourceID is required")
}
if strings.TrimSpace(cfg.TicketURL) == "" {
return nil, errors.New("personal source: TicketURL is required")
}
if cfg.TicketMode == "" {
cfg.TicketMode = "normal"
}
if cfg.TicketMode == "custom" && strings.TrimSpace(cfg.ClientSecret) == "" {
return nil, errors.New("personal source: ClientSecret is required when stream ticket mode is custom")
}
if cfg.HTTPClient == nil {
cfg.HTTPClient = &http.Client{Timeout: 30 * time.Second}
}
if cfg.WebSocketDialer == nil {
cfg.WebSocketDialer = websocket.DefaultDialer
}
if cfg.Now == nil {
cfg.Now = time.Now
}
if cfg.ReconnectMin <= 0 {
cfg.ReconnectMin = personalReconnectMinBackoff
}
if cfg.ReconnectMax <= 0 {
cfg.ReconnectMax = personalReconnectMaxBackoff
}
if cfg.ReconnectMax < cfg.ReconnectMin {
cfg.ReconnectMax = cfg.ReconnectMin
}
m := NewMachine()
m.now = cfg.Now
return &PersonalSource{cfg: cfg, machine: m}, nil
}
func (s *PersonalSource) State() Snapshot { return s.machine.Snapshot() }
func (s *PersonalSource) Start(ctx context.Context, emit dwsevent.EmitFn) error {
if emit == nil {
return errors.New("personal source: emit is required")
}
if !s.started.CompareAndSwap(false, true) {
return errors.New("personal source: Start called twice")
}
s.machine.OnConnecting()
defer s.machine.OnStopped()
backoff := s.cfg.ReconnectMin
for {
acked, err := s.runAttempt(ctx, emit)
if err == nil {
return nil
}
if ctx.Err() != nil {
return ctx.Err()
}
if !isRetryablePersonalError(err) {
return err
}
if acked {
backoff = s.cfg.ReconnectMin
}
s.machine.OnReconnect()
slog.Warn("personal source reconnecting",
"error", personalRetryLogError(err),
"retry_in", backoff,
"reconnect_count", s.machine.Snapshot().ReconnectCount,
)
if err := waitPersonalReconnect(ctx, backoff); err != nil {
return err
}
backoff = nextPersonalBackoff(backoff, s.cfg.ReconnectMax)
}
}
func (s *PersonalSource) runAttempt(ctx context.Context, emit dwsevent.EmitFn) (bool, error) {
ticket, err := s.fetchTicket(ctx)
if err != nil {
return false, err
}
wsURL, err := endpointWithTicket(ticket.Endpoint, ticket.Ticket)
if err != nil {
return false, err
}
conn, _, err := s.cfg.WebSocketDialer.DialContext(ctx, wsURL, nil)
if err != nil {
return false, retryPersonal(fmt.Errorf("personal source: dial websocket: %w", err))
}
attemptCtx, cancel := context.WithCancel(ctx)
defer func() {
cancel()
_ = conn.Close()
}()
closePersonalWebSocketOnContext(attemptCtx, conn)
s.machine.OnConnected()
acked := false
for {
messageType, data, err := conn.ReadMessage()
if err != nil {
if ctx.Err() != nil {
return acked, ctx.Err()
}
return acked, retryPersonal(fmt.Errorf("personal source: read websocket: %w", err))
}
if messageType != websocket.TextMessage {
continue
}
if err := s.handleFrame(conn, data, emit); err != nil {
return acked, err
}
acked = true
}
}
func (s *PersonalSource) fetchTicket(ctx context.Context) (*ticketResponse, error) {
body := map[string]any{
"sourceId": s.cfg.SourceID,
"mode": s.cfg.TicketMode,
}
if s.cfg.TicketMode == "custom" {
body["clientId"] = s.cfg.ClientID
body["clientSecret"] = s.cfg.ClientSecret
}
b, err := json.Marshal(body)
if err != nil {
return nil, err
}
req, err := http.NewRequestWithContext(ctx, http.MethodPost, s.cfg.TicketURL, bytes.NewReader(b))
if err != nil {
return nil, fmt.Errorf("personal source: create ticket request: %w", err)
}
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Accept", "application/json")
req.Header.Set("x-user-access-token", s.cfg.AccessToken)
req.Header.Set("Authorization", "Bearer "+s.cfg.AccessToken)
req.Header.Set("X-DWS-Client-Id", s.cfg.ClientID)
req.Header.Set("X-DWS-Source-Id", s.cfg.SourceID)
resp, err := s.cfg.HTTPClient.Do(req)
if err != nil {
return nil, retryPersonal(fmt.Errorf("personal source: fetch ticket: %w", err))
}
defer resp.Body.Close()
data, err := io.ReadAll(io.LimitReader(resp.Body, config.MaxResponseBodySize))
if err != nil {
return nil, retryPersonal(fmt.Errorf("personal source: read ticket response: %w", err))
}
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
err := fmt.Errorf("personal source: ticket HTTP %d", resp.StatusCode)
if retryableTicketStatus(resp.StatusCode) {
return nil, retryPersonal(err)
}
return nil, err
}
ticket, err := decodeTicket(data)
if err != nil {
return nil, err
}
if ticket.Endpoint == "" || ticket.Ticket == "" {
return nil, errors.New("personal source: ticket response missing endpoint or ticket")
}
return ticket, nil
}
func (s *PersonalSource) handleFrame(conn *websocket.Conn, data []byte, emit dwsevent.EmitFn) error {
df, err := payload.DecodeDataFrame(data)
if err != nil {
return fmt.Errorf("personal source: decode dataframe: %w", err)
}
raw := s.rawEventFromDataFrame(df)
logPersonalDataFrame(raw, df.Data)
emit(raw)
resp := payload.NewSuccessDataFrameResponse()
resp.SetHeader(payload.DataFrameHeaderKMessageId, df.GetMessageId())
resp.SetHeader(payload.DataFrameHeaderKContentType, payload.DataFrameContentTypeKJson)
if err := resp.SetJson(event.NewEventProcessResultSuccess()); err != nil {
return err
}
if err := conn.WriteJSON(resp); err != nil {
return retryPersonal(fmt.Errorf("personal source: write ack: %w", err))
}
s.machine.OnEvent()
return nil
}
func (s *PersonalSource) rawEventFromDataFrame(df *payload.DataFrame) *dwsevent.RawEvent {
hdr := event.NewEventHeaderFromDataFrame(df)
dataFields := parsePersonalDataFields(df.Data)
now := time.Now
if s != nil && s.cfg.Now != nil {
now = s.cfg.Now
}
sourceID := ""
if s != nil {
sourceID = s.cfg.SourceID
}
raw := &dwsevent.RawEvent{
EventID: firstNonEmpty(hdr.EventId, df.GetMessageId(), headerAny(df.Headers, "MESSAGE_ID", "messageId", "event_id", "eventId"), dataFields.string("eventId", "event_id")),
EventBornTime: firstInt64(hdr.EventBornTime, df.GetTimestamp(), df.Time),
EventCorpID: hdr.EventCorpId,
EventType: firstNonEmpty(hdr.EventType, eventTypeFromHeaders(df.Headers), dataFields.string("eventKey", "event_key"), dataFields.nestedString([]string{"source", "tag"})),
EventUnifiedAppID: hdr.EventUnifiedAppId,
EventScope: firstNonEmpty(headerAny(df.Headers, "event_scope", "eventScope"), "personal"),
SubscribeID: firstNonEmpty(headerAny(df.Headers, "SUB_ID", "subscribe_id", "subscribeId", "sub_id", "subId"), dataFields.string("subId", "sub_id", "subscribeId", "subscribe_id")),
SourceID: firstNonEmpty(headerAny(df.Headers, "SOURCE_ID", "source_id", "sourceId"), sourceID),
RuleType: firstNonEmpty(headerAny(df.Headers, "RULE_TYPE", "rule_type", "ruleType"), dataFields.string("ruleType", "rule_type"), dataFields.nestedString([]string{"ext", "ruleType"}, []string{"ext", "rule_type"})),
Data: df.Data,
Headers: copyHeaders(df.Headers),
ReceivedAt: now().UTC(),
}
return raw
}
func logPersonalDataFrame(raw *dwsevent.RawEvent, data string) {
if raw == nil {
return
}
slog.Debug("personal source received dataframe",
"event_type", raw.EventType,
"event_key", raw.EventType,
"event_id", raw.EventID,
"subscribe_id", raw.SubscribeID,
"source_id", raw.SourceID,
"rule_type", raw.RuleType,
"headers", redactPersonalRawStringMap(raw.Headers),
"data", sanitizePersonalRawPayload([]byte(data)),
)
}
func decodeTicket(data []byte) (*ticketResponse, error) {
var env struct {
Success bool `json:"success"`
Result json.RawMessage `json:"result"`
Data json.RawMessage `json:"data"`
Error any `json:"error"`
}
if err := json.Unmarshal(data, &env); err == nil && (env.Result != nil || env.Data != nil || env.Error != nil) {
raw := env.Result
if len(raw) == 0 || string(raw) == "null" {
raw = env.Data
}
var tr ticketResponse
if err := json.Unmarshal(raw, &tr); err != nil {
return nil, fmt.Errorf("personal source: parse ticket result: %w", err)
}
return &tr, nil
}
var tr ticketResponse
if err := json.Unmarshal(data, &tr); err != nil {
return nil, fmt.Errorf("personal source: parse ticket response: %w", err)
}
return &tr, nil
}
func endpointWithTicket(endpoint, ticket string) (string, error) {
u, err := url.Parse(strings.TrimSpace(endpoint))
if err != nil {
return "", fmt.Errorf("personal source: parse endpoint: %w", err)
}
if (u.Scheme != "ws" && u.Scheme != "wss") || u.Host == "" {
return "", errors.New("personal source: ticket response contains invalid websocket endpoint")
}
q := u.Query()
if q.Get("ticket") == "" {
q.Set("ticket", ticket)
u.RawQuery = q.Encode()
}
return u.String(), nil
}
func retryPersonal(err error) error {
if err == nil {
return nil
}
return &retryablePersonalError{err: err}
}
func isRetryablePersonalError(err error) bool {
var retryable *retryablePersonalError
return errors.As(err, &retryable)
}
func personalRetryLogError(err error) string {
message := err.Error()
switch {
case strings.Contains(message, "ticket HTTP"):
return message
case strings.Contains(message, "fetch ticket"):
return "personal source: fetch ticket: network error"
case strings.Contains(message, "read ticket response"):
return "personal source: read ticket response: network error"
case strings.Contains(message, "dial websocket"):
return "personal source: dial websocket: network error"
case strings.Contains(message, "read websocket"):
return "personal source: read websocket: connection closed"
case strings.Contains(message, "write ack"):
return "personal source: write ack: connection error"
default:
return "personal source: retryable stream error"
}
}
func retryableTicketStatus(status int) bool {
return status == http.StatusRequestTimeout ||
status == http.StatusTooManyRequests ||
status >= http.StatusInternalServerError
}
func waitPersonalReconnect(ctx context.Context, delay time.Duration) error {
timer := time.NewTimer(delay)
defer timer.Stop()
select {
case <-ctx.Done():
return ctx.Err()
case <-timer.C:
return nil
}
}
func nextPersonalBackoff(current, maximum time.Duration) time.Duration {
if current >= maximum || current > maximum/2 {
return maximum
}
return current * 2
}
func closePersonalWebSocketOnContext(ctx context.Context, conn *websocket.Conn) {
go func() {
<-ctx.Done()
_ = conn.Close()
}()
}
func headerAny(h payload.DataFrameHeader, names ...string) string {
for _, name := range names {
if v := strings.TrimSpace(headerValue(h, name)); v != "" {
return v
}
}
return ""
}
func headerValue(h payload.DataFrameHeader, name string) string {
if len(h) == 0 {
return ""
}
if v := h.Get(name); strings.TrimSpace(v) != "" {
return v
}
for k, v := range h {
if strings.EqualFold(k, name) {
return v
}
}
return ""
}
func eventTypeFromHeaders(h payload.DataFrameHeader) string {
if v := headerAny(h, "EVENT_TYPE", "event_type", "eventType", "EVENT_KEY", "event_key", "eventKey"); v != "" {
return v
}
if v := headerAny(h, "TOPIC", "topic"); v != "" && v != "*" {
return v
}
return ""
}
type personalDataFields map[string]any
func parsePersonalDataFields(raw string) personalDataFields {
raw = strings.TrimSpace(raw)
if raw == "" {
return nil
}
var value any
if err := json.Unmarshal([]byte(raw), &value); err != nil {
return nil
}
if s, ok := value.(string); ok {
var nested any
if err := json.Unmarshal([]byte(strings.TrimSpace(s)), &nested); err == nil {
value = nested
}
}
m, ok := value.(map[string]any)
if !ok {
return nil
}
return personalDataFields(m)
}
func (f personalDataFields) string(names ...string) string {
if len(f) == 0 {
return ""
}
for _, name := range names {
if v, ok := f[name].(string); ok && strings.TrimSpace(v) != "" {
return strings.TrimSpace(v)
}
for k, raw := range f {
if strings.EqualFold(k, name) {
if v, ok := raw.(string); ok && strings.TrimSpace(v) != "" {
return strings.TrimSpace(v)
}
}
}
}
return ""
}
func (f personalDataFields) nestedString(paths ...[]string) string {
if len(f) == 0 {
return ""
}
for _, path := range paths {
if v := nestedStringCI(map[string]any(f), path...); v != "" {
return v
}
}
return ""
}
func nestedStringCI(m map[string]any, path ...string) string {
var cur any = m
for _, segment := range path {
obj, ok := cur.(map[string]any)
if !ok {
return ""
}
next, ok := obj[segment]
if !ok {
for k, v := range obj {
if strings.EqualFold(k, segment) {
next = v
ok = true
break
}
}
}
if !ok {
return ""
}
cur = next
}
if v, ok := cur.(string); ok {
return strings.TrimSpace(v)
}
return ""
}
func firstNonEmpty(values ...string) string {
for _, v := range values {
if strings.TrimSpace(v) != "" {
return strings.TrimSpace(v)
}
}
return ""
}
func firstInt64(values ...int64) int64 {
for _, v := range values {
if v != 0 {
return v
}
}
return 0
}
func redactPersonalRawStringMap(in map[string]string) map[string]string {
if len(in) == 0 {
return nil
}
out := make(map[string]string, len(in))
for k, v := range in {
if personalRawSensitiveKey(k) {
out[k] = "<redacted>"
continue
}
out[k] = v
}
return out
}
func sanitizePersonalRawPayload(data []byte) string {
data = bytes.TrimSpace(data)
if len(data) == 0 {
return ""
}
var parsed any
if err := json.Unmarshal(data, &parsed); err == nil {
redacted := redactPersonalRawJSONValue(parsed)
if s, err := marshalPersonalRawJSON(redacted); err == nil {
return truncatePersonalRawLog(s)
}
}
return truncatePersonalRawLog(string(data))
}
func redactPersonalRawJSONValue(v any) any {
switch x := v.(type) {
case map[string]any:
out := make(map[string]any, len(x))
for k, value := range x {
if personalRawSensitiveKey(k) {
out[k] = "<redacted>"
continue
}
out[k] = redactPersonalRawJSONValue(value)
}
return out
case []any:
out := make([]any, len(x))
for i, value := range x {
out[i] = redactPersonalRawJSONValue(value)
}
return out
default:
return v
}
}
func marshalPersonalRawJSON(v any) (string, error) {
var buf bytes.Buffer
enc := json.NewEncoder(&buf)
enc.SetEscapeHTML(false)
if err := enc.Encode(v); err != nil {
return "", err
}
return strings.TrimSpace(buf.String()), nil
}
func personalRawSensitiveKey(key string) bool {
key = strings.ToLower(strings.TrimSpace(key))
return strings.Contains(key, "token") ||
strings.Contains(key, "secret") ||
strings.Contains(key, "ticket") ||
strings.Contains(key, "authorization")
}
func truncatePersonalRawLog(s string) string {
if len(s) <= personalRawDebugPayloadLimit {
return s
}
return s[:personalRawDebugPayloadLimit] + "...<truncated>"
}
+649
View File
@@ -0,0 +1,649 @@
// 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 source
import (
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"log/slog"
"net/http"
"net/http/httptest"
"strconv"
"strings"
"sync"
"sync/atomic"
"testing"
"time"
dwsevent "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/bus"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/transport"
"github.com/gorilla/websocket"
streamevent "github.com/open-dingtalk/dingtalk-stream-sdk-go/event"
"github.com/open-dingtalk/dingtalk-stream-sdk-go/payload"
)
func TestPersonalSourceFetchTicketConnectsAndACKs(t *testing.T) {
logs := capturePersonalSourceDebugLogs(t)
var wsEndpoint string
ackCh := make(chan payload.DataFrameResponse, 1)
upgrader := websocket.Upgrader{}
mux := http.NewServeMux()
mux.HandleFunc("/ticket", func(w http.ResponseWriter, r *http.Request) {
if got := r.Header.Get("x-user-access-token"); got != "token-1" {
t.Fatalf("x-user-access-token = %q", got)
}
var req map[string]any
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
t.Fatalf("decode ticket request: %v", err)
}
if req["sourceId"] != "open" || req["mode"] != "normal" {
t.Fatalf("ticket request = %#v", req)
}
_ = json.NewEncoder(w).Encode(map[string]any{
"success": true,
"result": map[string]any{
"endpoint": wsEndpoint,
"ticket": "ticket-1",
},
})
})
mux.HandleFunc("/ws", func(w http.ResponseWriter, r *http.Request) {
if got := r.URL.Query().Get("ticket"); got != "ticket-1" {
t.Fatalf("ticket query = %q", got)
}
conn, err := upgrader.Upgrade(w, r, nil)
if err != nil {
t.Fatalf("upgrade: %v", err)
}
defer conn.Close()
df := payload.DataFrame{
Type: "event",
Headers: payload.DataFrameHeader{
payload.DataFrameHeaderKMessageId: "msg-1",
streamevent.DataFrameHeaderKEventId: "evt-1",
streamevent.DataFrameHeaderKEventBornTime: "1234",
streamevent.DataFrameHeaderKEventCorpId: "corp-1",
streamevent.DataFrameHeaderKEventType: "user_im_message_receive_at",
"subscribeId": "sub-1",
"ruleType": "at",
"sourceId": "open",
"accessToken": "header-secret-token",
},
Data: `{"message":{"text":"hi"},"access_token":"data-secret-token","client_secret":"data-secret","ticket":"data-ticket","Authorization":"Bearer data-auth"}`,
}
if err := conn.WriteJSON(df); err != nil {
t.Fatalf("write dataframe: %v", err)
}
var ack payload.DataFrameResponse
if err := conn.ReadJSON(&ack); err != nil {
t.Fatalf("read ack: %v", err)
}
ackCh <- ack
})
srv := httptest.NewServer(mux)
defer srv.Close()
wsEndpoint = "ws" + strings.TrimPrefix(srv.URL, "http") + "/ws"
src, err := NewPersonal(PersonalConfig{
AccessToken: "token-1",
ClientID: "client-1",
SourceID: "open",
TicketURL: srv.URL + "/ticket",
TicketMode: "normal",
HTTPClient: srv.Client(),
Now: func() time.Time { return time.Unix(10, 0) },
})
if err != nil {
t.Fatalf("NewPersonal() error = %v", err)
}
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
evCh := make(chan *dwsevent.RawEvent, 1)
done := make(chan error, 1)
go func() {
done <- src.Start(ctx, func(ev *dwsevent.RawEvent) { evCh <- ev })
}()
select {
case ev := <-evCh:
if ev.SubscribeID != "sub-1" || ev.RuleType != "at" || ev.SourceID != "open" {
t.Fatalf("event personal fields = %#v", ev)
}
if ev.EventType != "user_im_message_receive_at" || ev.EventID != "evt-1" {
t.Fatalf("event = %#v", ev)
}
out := logs.String()
for _, want := range []string{"personal source received dataframe", "user_im_message_receive_at", "evt-1", "sub-1", "sourceId", "message", "<redacted>"} {
if !strings.Contains(out, want) {
t.Fatalf("debug log missing %q: %s", want, out)
}
}
for _, leaked := range []string{"header-secret-token", "data-secret-token", "data-secret", "data-ticket", "Bearer data-auth"} {
if strings.Contains(out, leaked) {
t.Fatalf("debug log leaked %q: %s", leaked, out)
}
}
case <-time.After(2 * time.Second):
t.Fatal("timed out waiting for event")
}
select {
case ack := <-ackCh:
if ack.Code != payload.DataFrameResponseStatusCodeKOK {
t.Fatalf("ack code = %d", ack.Code)
}
if ack.GetHeader(payload.DataFrameHeaderKMessageId) != "msg-1" {
t.Fatalf("ack message id = %q", ack.GetHeader(payload.DataFrameHeaderKMessageId))
}
case <-time.After(2 * time.Second):
t.Fatal("timed out waiting for ack")
}
cancel()
select {
case <-done:
case <-time.After(2 * time.Second):
t.Fatal("source did not stop after cancel")
}
}
func TestPersonalSourceReconnectsWithFreshTicket(t *testing.T) {
var wsEndpoint string
var ticketCalls atomic.Int32
ackCh := make(chan int, 2)
holdSecond := make(chan struct{})
var releaseSecondOnce sync.Once
releaseSecond := func() { releaseSecondOnce.Do(func() { close(holdSecond) }) }
upgrader := websocket.Upgrader{}
mux := http.NewServeMux()
mux.HandleFunc("/ticket", func(w http.ResponseWriter, _ *http.Request) {
attempt := int(ticketCalls.Add(1))
_ = json.NewEncoder(w).Encode(map[string]any{
"success": true,
"result": map[string]any{
"endpoint": wsEndpoint,
"ticket": fmt.Sprintf("ticket-%d", attempt),
},
})
})
mux.HandleFunc("/ws", func(w http.ResponseWriter, r *http.Request) {
attempt, err := strconv.Atoi(strings.TrimPrefix(r.URL.Query().Get("ticket"), "ticket-"))
if err != nil {
http.Error(w, "bad ticket", http.StatusBadRequest)
return
}
conn, err := upgrader.Upgrade(w, r, nil)
if err != nil {
return
}
defer conn.Close()
if err := conn.WriteJSON(personalTestDataFrame(attempt)); err != nil {
return
}
var ack payload.DataFrameResponse
if err := conn.ReadJSON(&ack); err != nil {
return
}
ackCh <- attempt
if attempt == 2 {
<-holdSecond
}
})
srv := httptest.NewServer(mux)
defer srv.Close()
defer releaseSecond()
wsEndpoint = "ws" + strings.TrimPrefix(srv.URL, "http") + "/ws"
src, err := NewPersonal(PersonalConfig{
AccessToken: "token",
ClientID: "client",
SourceID: "open",
TicketURL: srv.URL + "/ticket",
HTTPClient: srv.Client(),
ReconnectMin: 5 * time.Millisecond,
ReconnectMax: 10 * time.Millisecond,
})
if err != nil {
t.Fatal(err)
}
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
events := make(chan *dwsevent.RawEvent, 2)
done := make(chan error, 1)
go func() { done <- src.Start(ctx, func(ev *dwsevent.RawEvent) { events <- ev }) }()
for i := 1; i <= 2; i++ {
select {
case ev := <-events:
if ev.EventID != fmt.Sprintf("evt-%d", i) {
t.Fatalf("event %d ID = %q", i, ev.EventID)
}
case <-time.After(2 * time.Second):
t.Fatalf("timed out waiting for event %d", i)
}
select {
case attempt := <-ackCh:
if attempt != i {
t.Fatalf("ack attempt = %d, want %d", attempt, i)
}
case <-time.After(2 * time.Second):
t.Fatalf("timed out waiting for ACK %d", i)
}
}
if got := ticketCalls.Load(); got != 2 {
t.Fatalf("ticket calls = %d, want 2", got)
}
if got := src.State().ReconnectCount; got != 1 {
t.Fatalf("reconnect count = %d, want 1", got)
}
cancel()
releaseSecond()
select {
case err := <-done:
if !errors.Is(err, context.Canceled) {
t.Fatalf("Start() error = %v, want context canceled", err)
}
case <-time.After(2 * time.Second):
t.Fatal("source did not stop after cancel")
}
}
func TestPersonalSourceRetriesTicket500And429(t *testing.T) {
var wsEndpoint string
var ticketCalls atomic.Int32
acked := make(chan struct{}, 1)
hold := make(chan struct{})
var releaseOnce sync.Once
release := func() { releaseOnce.Do(func() { close(hold) }) }
upgrader := websocket.Upgrader{}
mux := http.NewServeMux()
mux.HandleFunc("/ticket", func(w http.ResponseWriter, _ *http.Request) {
switch ticketCalls.Add(1) {
case 1:
w.WriteHeader(http.StatusInternalServerError)
return
case 2:
w.WriteHeader(http.StatusTooManyRequests)
return
default:
_ = json.NewEncoder(w).Encode(map[string]any{
"success": true,
"result": map[string]any{"endpoint": wsEndpoint, "ticket": "ticket-ok"},
})
}
})
mux.HandleFunc("/ws", func(w http.ResponseWriter, r *http.Request) {
conn, err := upgrader.Upgrade(w, r, nil)
if err != nil {
return
}
defer conn.Close()
if err := conn.WriteJSON(personalTestDataFrame(1)); err != nil {
return
}
var ack payload.DataFrameResponse
if err := conn.ReadJSON(&ack); err != nil {
return
}
acked <- struct{}{}
<-hold
})
srv := httptest.NewServer(mux)
defer srv.Close()
defer release()
wsEndpoint = "ws" + strings.TrimPrefix(srv.URL, "http") + "/ws"
src, err := NewPersonal(PersonalConfig{
AccessToken: "token",
ClientID: "client",
SourceID: "open",
TicketURL: srv.URL + "/ticket",
HTTPClient: srv.Client(),
ReconnectMin: time.Millisecond,
ReconnectMax: 2 * time.Millisecond,
})
if err != nil {
t.Fatal(err)
}
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
done := make(chan error, 1)
go func() { done <- src.Start(ctx, func(*dwsevent.RawEvent) {}) }()
select {
case <-acked:
case <-time.After(2 * time.Second):
t.Fatal("timed out waiting for recovered connection")
}
if got := ticketCalls.Load(); got != 3 {
t.Fatalf("ticket calls = %d, want 3", got)
}
if got := src.State().ReconnectCount; got != 2 {
t.Fatalf("reconnect count = %d, want 2", got)
}
cancel()
release()
select {
case <-done:
case <-time.After(2 * time.Second):
t.Fatal("source did not stop after cancel")
}
}
func TestPersonalSourceTicketFatalResponsesDoNotRetry(t *testing.T) {
tests := []struct {
name string
status int
body string
}{
{name: "unauthorized", status: http.StatusUnauthorized},
{name: "forbidden", status: http.StatusForbidden},
{name: "malformed success", status: http.StatusOK, body: `{"success":true,"result":{"endpoint":"","ticket":""}}`},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
var calls atomic.Int32
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
calls.Add(1)
w.WriteHeader(tt.status)
_, _ = w.Write([]byte(tt.body))
}))
defer srv.Close()
src, err := NewPersonal(PersonalConfig{
AccessToken: "token",
ClientID: "client",
SourceID: "open",
TicketURL: srv.URL,
HTTPClient: srv.Client(),
ReconnectMin: time.Millisecond,
})
if err != nil {
t.Fatal(err)
}
if err := src.Start(t.Context(), func(*dwsevent.RawEvent) {}); err == nil {
t.Fatal("Start() error = nil, want fatal ticket error")
}
if got := calls.Load(); got != 1 {
t.Fatalf("ticket calls = %d, want 1", got)
}
})
}
}
func TestPersonalSourceRetriesTicketNetworkError(t *testing.T) {
var calls atomic.Int32
secondCall := make(chan struct{}, 1)
httpClient := &http.Client{Transport: roundTripFunc(func(*http.Request) (*http.Response, error) {
if calls.Add(1) == 2 {
secondCall <- struct{}{}
}
return nil, errors.New("network unavailable")
})}
src, err := NewPersonal(PersonalConfig{
AccessToken: "token",
ClientID: "client",
SourceID: "open",
TicketURL: "https://ticket.invalid",
HTTPClient: httpClient,
ReconnectMin: time.Millisecond,
ReconnectMax: time.Millisecond,
})
if err != nil {
t.Fatal(err)
}
ctx, cancel := context.WithCancel(context.Background())
done := make(chan error, 1)
go func() { done <- src.Start(ctx, func(*dwsevent.RawEvent) {}) }()
select {
case <-secondCall:
case <-time.After(time.Second):
t.Fatal("ticket network error was not retried")
}
cancel()
select {
case err := <-done:
if !errors.Is(err, context.Canceled) {
t.Fatalf("Start() error = %v, want context canceled", err)
}
case <-time.After(time.Second):
t.Fatal("source did not stop after cancel")
}
}
func TestPersonalSourceCancelDuringReconnectBackoff(t *testing.T) {
requestDone := make(chan struct{}, 1)
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
requestDone <- struct{}{}
w.WriteHeader(http.StatusInternalServerError)
}))
defer srv.Close()
src, err := NewPersonal(PersonalConfig{
AccessToken: "token",
ClientID: "client",
SourceID: "open",
TicketURL: srv.URL,
HTTPClient: srv.Client(),
ReconnectMin: 5 * time.Second,
ReconnectMax: 5 * time.Second,
})
if err != nil {
t.Fatal(err)
}
ctx, cancel := context.WithCancel(context.Background())
done := make(chan error, 1)
go func() { done <- src.Start(ctx, func(*dwsevent.RawEvent) {}) }()
select {
case <-requestDone:
case <-time.After(time.Second):
t.Fatal("ticket request did not arrive")
}
deadline := time.Now().Add(time.Second)
for src.State().ReconnectCount == 0 && time.Now().Before(deadline) {
time.Sleep(time.Millisecond)
}
if src.State().ReconnectCount == 0 {
t.Fatal("source did not enter reconnect backoff")
}
cancel()
select {
case err := <-done:
if !errors.Is(err, context.Canceled) {
t.Fatalf("Start() error = %v, want context canceled", err)
}
case <-time.After(500 * time.Millisecond):
t.Fatal("source did not exit promptly during reconnect backoff")
}
}
func TestNextPersonalBackoffCapsAtMaximum(t *testing.T) {
if got := nextPersonalBackoff(time.Second, 30*time.Second); got != 2*time.Second {
t.Fatalf("next backoff = %s, want 2s", got)
}
if got := nextPersonalBackoff(20*time.Second, 30*time.Second); got != 30*time.Second {
t.Fatalf("capped backoff = %s, want 30s", got)
}
}
func TestPersonalRetryLogErrorDoesNotExposeWebSocketURL(t *testing.T) {
err := retryPersonal(errors.New("personal source: dial websocket: wss://example.test/ws?ticket=secret-ticket"))
got := personalRetryLogError(err)
if strings.Contains(got, "secret-ticket") || strings.Contains(got, "example.test") {
t.Fatalf("retry log exposed websocket details: %q", got)
}
}
func TestPersonalSourceParsesUppercaseHeaders(t *testing.T) {
src := personalSourceForRawEventTests()
data := `{"payload":true}`
raw := src.rawEventFromDataFrame(&payload.DataFrame{
Time: 12345,
Headers: payload.DataFrameHeader{
"EVENT_TYPE": "user_im_message_receive_o2o",
"SUB_ID": "sub-1",
"SOURCE_ID": "pre_open_source",
"MESSAGE_ID": "evt-1",
},
Data: data,
})
if raw.EventType != "user_im_message_receive_o2o" {
t.Fatalf("EventType = %q", raw.EventType)
}
if raw.SubscribeID != "sub-1" {
t.Fatalf("SubscribeID = %q", raw.SubscribeID)
}
if raw.SourceID != "pre_open_source" {
t.Fatalf("SourceID = %q", raw.SourceID)
}
if raw.EventID != "evt-1" {
t.Fatalf("EventID = %q", raw.EventID)
}
if raw.EventScope != "personal" {
t.Fatalf("EventScope = %q", raw.EventScope)
}
if raw.Data != data {
t.Fatalf("Data changed: %q", raw.Data)
}
}
func TestPersonalSourceParsesTopicFallbackAndIgnoresWildcard(t *testing.T) {
src := personalSourceForRawEventTests()
raw := src.rawEventFromDataFrame(&payload.DataFrame{
Headers: payload.DataFrameHeader{
"TOPIC": "user_im_message_receive_o2o",
"topic": "*",
},
})
if raw.EventType != "user_im_message_receive_o2o" {
t.Fatalf("EventType from TOPIC = %q", raw.EventType)
}
raw = src.rawEventFromDataFrame(&payload.DataFrame{
Headers: payload.DataFrameHeader{"topic": "*"},
})
if raw.EventType != "" {
t.Fatalf("EventType from wildcard topic = %q, want empty", raw.EventType)
}
}
func TestPersonalSourceParsesDataFallbackWithoutChangingData(t *testing.T) {
src := personalSourceForRawEventTests()
payloadJSON := `{"eventKey":"user_im_message_receive_o2o","subId":"sub-data","eventId":"evt-data","ext":{"ruleType":"singleChat"}}`
encodedPayload, err := json.Marshal(payloadJSON)
if err != nil {
t.Fatal(err)
}
data := string(encodedPayload)
raw := src.rawEventFromDataFrame(&payload.DataFrame{Data: data})
if raw.EventType != "user_im_message_receive_o2o" {
t.Fatalf("EventType = %q", raw.EventType)
}
if raw.SubscribeID != "sub-data" {
t.Fatalf("SubscribeID = %q", raw.SubscribeID)
}
if raw.EventID != "evt-data" {
t.Fatalf("EventID = %q", raw.EventID)
}
if raw.RuleType != "singleChat" {
t.Fatalf("RuleType = %q", raw.RuleType)
}
if raw.Data != data {
t.Fatalf("Data changed: %q", raw.Data)
}
}
func TestPersonalSourceParsesSourceTagDataFallback(t *testing.T) {
src := personalSourceForRawEventTests()
data := `{"source":{"tag":"user_im_message_receive_group"},"subId":"sub-group"}`
raw := src.rawEventFromDataFrame(&payload.DataFrame{Data: data})
if raw.EventType != "user_im_message_receive_group" {
t.Fatalf("EventType = %q", raw.EventType)
}
if raw.SubscribeID != "sub-group" {
t.Fatalf("SubscribeID = %q", raw.SubscribeID)
}
}
func TestPersonalSourceParsedHeadersPassNormalBusFilter(t *testing.T) {
src := personalSourceForRawEventTests()
raw := src.rawEventFromDataFrame(&payload.DataFrame{
Headers: payload.DataFrameHeader{
"EVENT_TYPE": "user_im_message_receive_o2o",
"SUB_ID": "sub-1",
},
})
h := bus.NewHub(10)
c, err := h.Register(transport.Hello{
EventTypes: []string{"user_im_message_receive_o2o"},
SubscribeID: "sub-1",
})
if err != nil {
t.Fatal(err)
}
h.Deliver(raw)
select {
case frame := <-c.SendCh:
ev, ok := frame.(transport.Event)
if !ok {
t.Fatalf("frame = %T, want transport.Event", frame)
}
if ev.EventType != "user_im_message_receive_o2o" || ev.SubscribeID != "sub-1" {
t.Fatalf("event = %#v", ev)
}
case <-time.After(time.Second):
t.Fatal("timed out waiting for filtered event")
}
}
func personalSourceForRawEventTests() *PersonalSource {
return &PersonalSource{cfg: PersonalConfig{
SourceID: "fallback_source",
Now: func() time.Time { return time.Unix(20, 0) },
}}
}
func personalTestDataFrame(attempt int) payload.DataFrame {
return payload.DataFrame{
Type: "event",
Headers: payload.DataFrameHeader{
payload.DataFrameHeaderKMessageId: fmt.Sprintf("msg-%d", attempt),
streamevent.DataFrameHeaderKEventId: fmt.Sprintf("evt-%d", attempt),
streamevent.DataFrameHeaderKEventType: "user_im_message_receive_o2o",
"SUB_ID": fmt.Sprintf("sub-%d", attempt),
},
Data: fmt.Sprintf(`{"eventKey":"user_im_message_receive_o2o","eventId":"evt-%d"}`, attempt),
}
}
type roundTripFunc func(*http.Request) (*http.Response, error)
func (f roundTripFunc) RoundTrip(req *http.Request) (*http.Response, error) {
return f(req)
}
func capturePersonalSourceDebugLogs(t *testing.T) *bytes.Buffer {
t.Helper()
var buf bytes.Buffer
previous := slog.Default()
slog.SetDefault(slog.New(slog.NewJSONHandler(&buf, &slog.HandlerOptions{Level: slog.LevelDebug})))
t.Cleanup(func() {
slog.SetDefault(previous)
})
return &buf
}
+275
View File
@@ -0,0 +1,275 @@
// 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 source
import (
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"io"
"net/http"
"net/url"
"strings"
"time"
dwsevent "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event"
"github.com/gorilla/websocket"
"github.com/open-dingtalk/dingtalk-stream-sdk-go/payload"
)
const (
PortalTicketModeNormal = "normal"
PortalTicketModeCustom = "custom"
)
// PortalTicketConfig describes the portal-managed user Stream ticket flow.
// normal mode uses portal-side managed credentials; custom mode asks portal to
// open the user connection with the caller-provided clientId/clientSecret.
type PortalTicketConfig struct {
TicketURL string
AccessToken string
SourceID string
Mode string
ClientID string
ClientSecret string
UserAgent string
HTTPClient *http.Client
}
func (c *PortalTicketConfig) Valid() error {
if c == nil {
return errors.New("source: PortalTicketConfig is nil")
}
if strings.TrimSpace(c.TicketURL) == "" {
return errors.New("source: portal ticket URL is required")
}
if strings.TrimSpace(c.AccessToken) == "" {
return errors.New("source: portal access token is required")
}
if strings.TrimSpace(c.SourceID) == "" {
return errors.New("source: portal sourceId is required")
}
mode := normalizePortalTicketMode(c.Mode)
if mode == "" {
return fmt.Errorf("source: unsupported portal ticket mode %q", c.Mode)
}
if mode == PortalTicketModeCustom &&
(strings.TrimSpace(c.ClientID) == "" || strings.TrimSpace(c.ClientSecret) == "") {
return errors.New("source: custom portal ticket mode requires clientId/clientSecret")
}
return nil
}
func (c *PortalTicketConfig) normalizedMode() string {
return normalizePortalTicketMode(c.Mode)
}
func normalizePortalTicketMode(mode string) string {
mode = strings.ToLower(strings.TrimSpace(mode))
if mode == "" {
return PortalTicketModeNormal
}
switch mode {
case PortalTicketModeNormal, PortalTicketModeCustom:
return mode
default:
return ""
}
}
func (s *DingtalkSource) startPortalTicket(ctx context.Context, emit dwsevent.EmitFn) error {
s.machine.OnConnecting()
ticket, err := requestPortalTicket(ctx, s.cfg.PortalTicket)
if err != nil {
s.machine.OnStopped()
return err
}
wsURL, err := websocketURL(ticket)
if err != nil {
s.machine.OnStopped()
return err
}
userAgent := strings.TrimSpace(s.cfg.PortalTicket.UserAgent)
if userAgent == "" {
userAgent = "dws-event-consume"
}
conn, resp, err := (&websocket.Dialer{HandshakeTimeout: 20 * time.Second}).DialContext(ctx, wsURL, http.Header{
"User-Agent": []string{userAgent},
})
if err != nil {
s.machine.OnStopped()
if resp != nil {
defer resp.Body.Close()
raw, _ := io.ReadAll(io.LimitReader(resp.Body, 1<<20))
return fmt.Errorf("source: portal stream connect HTTP %d: %s: %w",
resp.StatusCode, truncatePortalTicketLog(string(raw), 300), err)
}
return fmt.Errorf("source: portal stream connect: %w", err)
}
defer conn.Close()
s.machine.OnConnected()
closeOnContext(ctx, conn)
handler := s.makeHandler(emit)
for {
messageType, message, err := conn.ReadMessage()
if err != nil {
s.machine.OnStopped()
if isContextDone(ctx) {
return ctx.Err()
}
return fmt.Errorf("source: portal stream read: %w", err)
}
if messageType != websocket.TextMessage && messageType != websocket.BinaryMessage {
continue
}
df, err := payload.DecodeDataFrame(message)
if err != nil {
continue
}
resp, err := handler(ctx, df)
if err != nil {
s.machine.OnStopped()
return err
}
if resp == nil {
continue
}
ensurePortalAckHeaders(resp, df)
if err := conn.WriteMessage(websocket.TextMessage, resp.Encode()); err != nil {
s.machine.OnStopped()
if isContextDone(ctx) {
return ctx.Err()
}
return fmt.Errorf("source: portal stream ack: %w", err)
}
}
}
type portalStreamTicket struct {
Endpoint string `json:"endpoint"`
Ticket string `json:"ticket"`
}
func requestPortalTicket(ctx context.Context, cfg *PortalTicketConfig) (portalStreamTicket, error) {
httpClient := cfg.HTTPClient
if httpClient == nil {
httpClient = &http.Client{Timeout: 20 * time.Second}
}
body := map[string]string{
"sourceId": strings.TrimSpace(cfg.SourceID),
"channelType": strings.TrimSpace(cfg.SourceID),
"mode": cfg.normalizedMode(),
}
if cfg.normalizedMode() == PortalTicketModeCustom {
body["clientId"] = strings.TrimSpace(cfg.ClientID)
body["clientSecret"] = strings.TrimSpace(cfg.ClientSecret)
}
rawBody, err := json.Marshal(body)
if err != nil {
return portalStreamTicket{}, err
}
req, err := http.NewRequestWithContext(ctx, http.MethodPost, strings.TrimSpace(cfg.TicketURL), bytes.NewReader(rawBody))
if err != nil {
return portalStreamTicket{}, err
}
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Accept", "application/json")
if ua := strings.TrimSpace(cfg.UserAgent); ua != "" {
req.Header.Set("User-Agent", ua)
}
req.Header.Set("x-user-access-token", cfg.AccessToken)
resp, err := httpClient.Do(req)
if err != nil {
return portalStreamTicket{}, fmt.Errorf("source: portal ticket request: %w", err)
}
defer resp.Body.Close()
raw, _ := io.ReadAll(io.LimitReader(resp.Body, 1<<20))
if resp.StatusCode >= 400 {
return portalStreamTicket{}, fmt.Errorf("source: portal ticket HTTP %d: %s",
resp.StatusCode, truncatePortalTicketLog(string(raw), 300))
}
var direct portalStreamTicket
if err := json.Unmarshal(raw, &direct); err == nil && direct.Endpoint != "" && direct.Ticket != "" {
return direct, nil
}
var envelope struct {
Success bool `json:"success"`
Result portalStreamTicket `json:"result"`
ErrorCode string `json:"errorCode"`
ErrorMsg string `json:"errorMsg"`
}
if err := json.Unmarshal(raw, &envelope); err != nil {
return portalStreamTicket{}, fmt.Errorf("source: portal ticket parse: %w", err)
}
if !envelope.Success {
return portalStreamTicket{}, fmt.Errorf("source: portal ticket failed: %s %s",
envelope.ErrorCode, envelope.ErrorMsg)
}
if envelope.Result.Endpoint == "" || envelope.Result.Ticket == "" {
return portalStreamTicket{}, errors.New("source: portal ticket result missing endpoint/ticket")
}
return envelope.Result, nil
}
func websocketURL(ticket portalStreamTicket) (string, error) {
u, err := url.Parse(ticket.Endpoint)
if err != nil {
return "", err
}
q := u.Query()
q.Set("ticket", ticket.Ticket)
u.RawQuery = q.Encode()
return u.String(), nil
}
func ensurePortalAckHeaders(resp *payload.DataFrameResponse, df *payload.DataFrame) {
if resp.GetHeader(payload.DataFrameHeaderKMessageId) == "" {
resp.SetHeader(payload.DataFrameHeaderKMessageId, df.GetMessageId())
}
if resp.GetHeader(payload.DataFrameHeaderKContentType) == "" {
resp.SetHeader(payload.DataFrameHeaderKContentType, payload.DataFrameContentTypeKJson)
}
}
func closeOnContext(ctx context.Context, conn *websocket.Conn) {
go func() {
<-ctx.Done()
_ = conn.Close()
}()
}
func isContextDone(ctx context.Context) bool {
select {
case <-ctx.Done():
return true
default:
return false
}
}
func truncatePortalTicketLog(s string, max int) string {
s = strings.TrimSpace(s)
if len(s) <= max {
return s
}
return s[:max] + "..."
}
+181
View File
@@ -0,0 +1,181 @@
// 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 source
import (
"sync"
"time"
)
// State enumerates the connection states surfaced by `dws event status`.
//
// Stream SDK v0.9.1 does not expose a user-registerable ConnectionState
// hook (verified P0 §1 escape-hatch table row #4), so the state machine
// runs in "inferred" mode: it derives state from observable signals —
// Start() return value, last-event-received timestamp, last-error timestamp,
// SDK reconnect log lines parsed by the source layer. The Source field
// `state_source` (env-style annotation) is emitted alongside the state so
// downstream tools can distinguish authoritative from inferred values when
// SDK hook support arrives in a future release.
type State string
const (
// StateDisconnected is the zero state before Start has been called.
StateDisconnected State = "disconnected"
// StateConnecting is set on entry to Start, before SDK dial completes.
StateConnecting State = "connecting"
// StateConnected is set when Start returns nil (dial succeeded).
StateConnected State = "connected"
// StateIdle is set when the connection is still up but no event has been
// observed for a quiet window (default 60s). Distinguishes "everything
// fine, just no traffic" from "looks connected but broken".
StateIdle State = "idle"
// StateReconnecting is set when the SDK has signalled a reconnect (via
// log line parse or future hook). Cleared on next event received.
StateReconnecting State = "reconnecting"
// StateDegraded is set when reconnects keep happening but no events flow.
// Useful to surface "looks alive, actually broken" cases via status.
StateDegraded State = "degraded"
// StateStopped is set after Close() / context cancellation.
StateStopped State = "stopped"
)
// SourceKind identifies whether the state was reported by an authoritative
// SDK hook or inferred from observable side-channels.
type SourceKind string
const (
// SourceHook is reserved for future SDK versions that expose
// ConnectionState / keepalive callbacks. v1 does not emit this.
SourceHook SourceKind = "hook"
// SourceInferred is the only value emitted by v1.
SourceInferred SourceKind = "inferred"
)
// IdleAfter is the duration of no events after which StateConnected
// transitions to StateIdle. Exposed for tests to bypass real-time waits.
var IdleAfter = 60 * time.Second
// Machine tracks the connection state. Methods are safe for concurrent use.
type Machine struct {
mu sync.RWMutex
state State
source SourceKind
lastEventAt time.Time
lastReconnectAt time.Time
reconnectCount int
now func() time.Time
}
// NewMachine returns a Machine in the StateDisconnected initial state.
func NewMachine() *Machine {
return &Machine{
state: StateDisconnected,
source: SourceInferred,
now: time.Now,
}
}
// State returns the current effective state, applying time-based idle
// detection on top of the stored state. This is the public read accessor
// used by `event status` formatting.
func (m *Machine) State() State {
m.mu.RLock()
defer m.mu.RUnlock()
return m.effectiveStateLocked()
}
// Snapshot returns the full state view for status rendering: state +
// source + last event/reconnect timestamps + reconnect count.
type Snapshot struct {
State State `json:"state"`
StateSource SourceKind `json:"state_source"`
LastEventAt time.Time `json:"last_event_at,omitempty"`
LastReconnectAt time.Time `json:"last_reconnect_at,omitempty"`
ReconnectCount int `json:"reconnect_count"`
}
// Snapshot returns the current state view atomically.
func (m *Machine) Snapshot() Snapshot {
m.mu.RLock()
defer m.mu.RUnlock()
return Snapshot{
State: m.effectiveStateLocked(),
StateSource: m.source,
LastEventAt: m.lastEventAt,
LastReconnectAt: m.lastReconnectAt,
ReconnectCount: m.reconnectCount,
}
}
// OnConnecting moves to StateConnecting. Called by Source.Start before the
// SDK dial completes.
func (m *Machine) OnConnecting() {
m.mu.Lock()
defer m.mu.Unlock()
m.state = StateConnecting
}
// OnConnected moves to StateConnected. Called when SDK Start returns nil.
func (m *Machine) OnConnected() {
m.mu.Lock()
defer m.mu.Unlock()
m.state = StateConnected
}
// OnEvent records the timestamp of the most recently received event. Resets
// any Reconnecting or Idle state back to Connected (the connection is
// demonstrably alive).
func (m *Machine) OnEvent() {
m.mu.Lock()
defer m.mu.Unlock()
m.lastEventAt = m.now()
if m.state == StateReconnecting || m.state == StateIdle || m.state == StateDegraded {
m.state = StateConnected
}
}
// OnReconnect marks a reconnect signal observed (from SDK log or future
// hook). Increments the counter and stamps the time.
func (m *Machine) OnReconnect() {
m.mu.Lock()
defer m.mu.Unlock()
m.lastReconnectAt = m.now()
m.reconnectCount++
m.state = StateReconnecting
}
// OnStopped marks the SDK as closed (graceful exit or Close()).
func (m *Machine) OnStopped() {
m.mu.Lock()
defer m.mu.Unlock()
m.state = StateStopped
}
// effectiveStateLocked applies idle-detection on top of the stored state.
// Caller must hold m.mu (read or write).
func (m *Machine) effectiveStateLocked() State {
switch m.state {
case StateConnected:
// Connected → Idle after IdleAfter with no events. If we've also
// seen recent reconnects with no recovered events, surface Degraded.
if !m.lastEventAt.IsZero() && m.now().Sub(m.lastEventAt) >= IdleAfter {
if m.reconnectCount > 0 && m.lastReconnectAt.After(m.lastEventAt) {
return StateDegraded
}
return StateIdle
}
}
return m.state
}
+131
View File
@@ -0,0 +1,131 @@
// 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 source
import (
"testing"
"time"
)
func newMachineAt(now time.Time) *Machine {
m := NewMachine()
m.now = func() time.Time { return now }
return m
}
func TestMachine_InitialState(t *testing.T) {
m := NewMachine()
if m.State() != StateDisconnected {
t.Fatalf("initial state = %s, want %s", m.State(), StateDisconnected)
}
snap := m.Snapshot()
if snap.StateSource != SourceInferred {
t.Fatalf("v1 must emit StateSource=inferred, got %s", snap.StateSource)
}
}
func TestMachine_TransitionsHappyPath(t *testing.T) {
now := time.Date(2026, 5, 28, 12, 0, 0, 0, time.UTC)
m := newMachineAt(now)
m.OnConnecting()
if m.State() != StateConnecting {
t.Fatalf("after OnConnecting: %s", m.State())
}
m.OnConnected()
if m.State() != StateConnected {
t.Fatalf("after OnConnected: %s", m.State())
}
m.OnEvent()
if m.State() != StateConnected {
t.Fatalf("after OnEvent: %s (should remain connected)", m.State())
}
m.OnStopped()
if m.State() != StateStopped {
t.Fatalf("after OnStopped: %s", m.State())
}
}
func TestMachine_IdleAfterQuietWindow(t *testing.T) {
now := time.Date(2026, 5, 28, 12, 0, 0, 0, time.UTC)
m := newMachineAt(now)
m.OnConnected()
m.OnEvent() // lastEventAt = now
// Advance time past IdleAfter
m.now = func() time.Time { return now.Add(IdleAfter + 5*time.Second) }
if got := m.State(); got != StateIdle {
t.Fatalf("state after quiet window: %s, want %s", got, StateIdle)
}
// New event resets to connected
m.OnEvent()
if got := m.State(); got != StateConnected {
t.Fatalf("state after new event: %s, want %s", got, StateConnected)
}
}
func TestMachine_DegradedWhenReconnectsButNoEvents(t *testing.T) {
now := time.Date(2026, 5, 28, 12, 0, 0, 0, time.UTC)
m := newMachineAt(now)
m.OnConnected()
m.OnEvent() // baseline
// Advance + reconnect signal (no event recovery)
later := now.Add(IdleAfter + 30*time.Second)
m.now = func() time.Time { return later }
m.OnReconnect() // sets state=reconnecting, lastReconnectAt=later
// State while reconnecting
if got := m.State(); got != StateReconnecting {
t.Fatalf("after OnReconnect: %s, want %s", got, StateReconnecting)
}
// Push state back to connected (as if reconnect succeeded but no traffic)
m.OnConnected()
// Now idle window passes with no events but a recent reconnect → degraded
muchLater := later.Add(IdleAfter + 5*time.Second)
m.now = func() time.Time { return muchLater }
if got := m.State(); got != StateDegraded {
t.Fatalf("state after reconnect-but-no-events: %s, want %s", got, StateDegraded)
}
}
func TestMachine_ReconnectCount(t *testing.T) {
m := NewMachine()
m.OnConnected()
m.OnReconnect()
m.OnReconnect()
m.OnReconnect()
if snap := m.Snapshot(); snap.ReconnectCount != 3 {
t.Fatalf("ReconnectCount = %d, want 3", snap.ReconnectCount)
}
}
func TestMachine_OnEventClearsReconnecting(t *testing.T) {
m := NewMachine()
m.OnConnected()
m.OnReconnect()
if m.State() != StateReconnecting {
t.Fatal("expected StateReconnecting after OnReconnect")
}
m.OnEvent()
if m.State() != StateConnected {
t.Fatalf("OnEvent should clear Reconnecting, got %s", m.State())
}
}
+29
View File
@@ -0,0 +1,29 @@
// 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 implements the cross-platform local IPC between the bus
// daemon and consume clients. The wire is \n-delimited JSON frames over a
// Unix domain socket (darwin/linux) or Windows Named Pipe; the abstraction
// keeps both endpoints (bus side / consume side) and both platforms behind
// one Listener/Dialer interface.
//
// Frame protocol (see plan §9):
//
// consume → bus hello, bye(reason=client_done), heartbeat, status_req
// bus → consume hello_ack, event, source_state, bye(reason=shutdown),
// heartbeat, status_resp
//
// Frames are independent JSON objects; no length prefix, no checksum.
// Maximum single-frame size is enforced by the Reader to bound memory and
// prevent a buggy/malicious peer from causing OOM (MaxFrameBytes).
package transport
+117
View File
@@ -0,0 +1,117 @@
// 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 (
"bufio"
"encoding/json"
"errors"
"fmt"
"io"
)
// MaxFrameBytes caps a single frame's wire size. 1 MiB accommodates large
// event payloads (kanban cards, doc snapshots) with headroom while still
// preventing a buggy/malicious peer from causing OOM via an unbounded line.
// Callers writing frames larger than this should split their work, not bump
// the cap — the SDK enforces its own cloud-side cap that is well below this.
const MaxFrameBytes = 1 << 20 // 1 MiB
// ErrFrameTooLarge is returned by Reader.Read when an incoming frame exceeds
// MaxFrameBytes. The underlying connection is left in an indeterminate
// state (the rest of the oversized frame is NOT drained) — callers should
// close the connection on receipt of this error.
var ErrFrameTooLarge = errors.New("transport: frame exceeds MaxFrameBytes")
// Reader reads \n-delimited JSON frames from an underlying byte stream.
// Wrap each connection (one per direction) in a Reader; Reader is not
// safe for concurrent use.
type Reader struct {
br *bufio.Reader
}
// NewReader returns a Reader with a buffer sized to MaxFrameBytes so a
// single frame can fit in the buffer without growing.
func NewReader(r io.Reader) *Reader {
return &Reader{br: bufio.NewReaderSize(r, MaxFrameBytes)}
}
// Read returns the next frame's raw bytes (the terminating \n is stripped).
// Returns io.EOF cleanly when the peer closes the connection between
// frames. Returns ErrFrameTooLarge when a single frame exceeds the cap.
func (r *Reader) Read() ([]byte, error) {
line, err := r.br.ReadSlice('\n')
if err == nil {
// Strip trailing \n. Safe because ReadSlice always includes the
// delimiter when err == nil.
out := make([]byte, len(line)-1)
copy(out, line[:len(line)-1])
return out, nil
}
if errors.Is(err, bufio.ErrBufferFull) {
// Read past current contents (without retaining) to skip the
// remainder of the oversize frame would be the friendly path,
// but plan says callers should close the connection — keep it
// simple and surface ErrFrameTooLarge immediately.
return nil, ErrFrameTooLarge
}
// EOF on a partial line: report it as EOF (clean close between frames
// returns a fresh EOF with an empty line, already handled above).
if errors.Is(err, io.EOF) && len(line) == 0 {
return nil, io.EOF
}
if err != nil {
return nil, err
}
return nil, nil // unreachable
}
// ReadJSON decodes the next frame into dst. dst must be a non-nil pointer.
func (r *Reader) ReadJSON(dst any) error {
raw, err := r.Read()
if err != nil {
return err
}
if err := json.Unmarshal(raw, dst); err != nil {
return fmt.Errorf("transport: decode frame: %w", err)
}
return nil
}
// Writer marshals values to JSON and writes them as \n-terminated frames.
// Writer is not safe for concurrent use; callers wanting a fan-in writer
// must serialise externally.
type Writer struct {
w io.Writer
}
// NewWriter returns a Writer wrapping w. The Writer does NOT buffer —
// every WriteJSON call performs a single Write to the underlying stream
// so backpressure is immediate.
func NewWriter(w io.Writer) *Writer { return &Writer{w: w} }
// WriteJSON serialises v and writes it as one frame. Refuses to write
// frames larger than MaxFrameBytes with ErrFrameTooLarge.
func (w *Writer) WriteJSON(v any) error {
buf, err := json.Marshal(v)
if err != nil {
return fmt.Errorf("transport: encode frame: %w", err)
}
if len(buf)+1 > MaxFrameBytes {
return ErrFrameTooLarge
}
buf = append(buf, '\n')
_, err = w.w.Write(buf)
return err
}
+114
View File
@@ -0,0 +1,114 @@
// 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 (
"bytes"
"errors"
"io"
"strings"
"testing"
)
func TestWriterReader_Roundtrip(t *testing.T) {
var buf bytes.Buffer
w := NewWriter(&buf)
if err := w.WriteJSON(Hello{Type: FrameTypeHello, ConsumerPID: 42, EventTypes: []string{"im.*"}}); err != nil {
t.Fatalf("WriteJSON: %v", err)
}
if err := w.WriteJSON(Heartbeat{Type: FrameTypeHeartbeat}); err != nil {
t.Fatalf("WriteJSON 2: %v", err)
}
r := NewReader(&buf)
var h Hello
if err := r.ReadJSON(&h); err != nil {
t.Fatalf("ReadJSON 1: %v", err)
}
if h.ConsumerPID != 42 || h.Type != FrameTypeHello || len(h.EventTypes) != 1 || h.EventTypes[0] != "im.*" {
t.Fatalf("decoded Hello = %+v", h)
}
var hb Heartbeat
if err := r.ReadJSON(&hb); err != nil {
t.Fatalf("ReadJSON 2: %v", err)
}
if hb.Type != FrameTypeHeartbeat {
t.Fatalf("hb type = %s", hb.Type)
}
}
func TestReader_CleanEOFBetweenFrames(t *testing.T) {
var buf bytes.Buffer
NewWriter(&buf).WriteJSON(Heartbeat{Type: FrameTypeHeartbeat})
r := NewReader(&buf)
if _, err := r.Read(); err != nil {
t.Fatalf("first read: %v", err)
}
_, err := r.Read()
if !errors.Is(err, io.EOF) {
t.Fatalf("second read after clean close = %v, want EOF", err)
}
}
func TestReader_FrameTooLarge(t *testing.T) {
// Synthesise a > MaxFrameBytes line with no \n
big := strings.Repeat("A", MaxFrameBytes+10)
r := NewReader(strings.NewReader(big + "\n"))
_, err := r.Read()
if !errors.Is(err, ErrFrameTooLarge) {
t.Fatalf("Read oversized = %v, want ErrFrameTooLarge", err)
}
}
func TestWriter_FrameTooLarge(t *testing.T) {
// Build a Hello with a filter string that would marshal larger than the cap.
huge := strings.Repeat("x", MaxFrameBytes)
w := NewWriter(io.Discard)
err := w.WriteJSON(Hello{Type: FrameTypeHello, Filter: huge})
if !errors.Is(err, ErrFrameTooLarge) {
t.Fatalf("WriteJSON oversized = %v, want ErrFrameTooLarge", err)
}
}
func TestReader_MalformedJSONErrors(t *testing.T) {
r := NewReader(strings.NewReader("{not json}\n"))
var x Hello
if err := r.ReadJSON(&x); err == nil {
t.Fatal("expected decode error for malformed JSON")
}
}
func TestPeekType(t *testing.T) {
cases := []struct {
raw string
want FrameType
err bool
}{
{`{"type":"hello"}`, FrameTypeHello, false},
{`{"type":"event","seq":1}`, FrameTypeEvent, false},
{`{"foo":"bar"}`, "", true},
{`not json`, "", true},
}
for _, tc := range cases {
got, err := PeekType([]byte(tc.raw))
if (err != nil) != tc.err {
t.Errorf("PeekType(%q) err = %v, wantErr=%v", tc.raw, err, tc.err)
continue
}
if got != tc.want {
t.Errorf("PeekType(%q) = %q, want %q", tc.raw, got, tc.want)
}
}
}
+199
View File
@@ -0,0 +1,199 @@
// 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 (
"encoding/json"
"fmt"
dwsevent "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event"
)
// FrameType discriminates the JSON frames flowing over the IPC channel.
// Wire values are stable strings — bumping any of these is a protocol
// breaking change requiring a coordinated bus + consume rollout.
type FrameType string
const (
FrameTypeHello FrameType = "hello" // consume → bus
FrameTypeHelloAck FrameType = "hello_ack" // bus → consume
FrameTypeEvent FrameType = "event" // bus → consume
FrameTypeHeartbeat FrameType = "heartbeat" // bidirectional
FrameTypeSourceState FrameType = "source_state" // bus → consume
FrameTypeBye FrameType = "bye" // bidirectional
FrameTypeStatusReq FrameType = "status_req" // consume/ad-hoc → bus
FrameTypeStatusResp FrameType = "status_resp" // bus → consume/ad-hoc
)
// Hello is the first frame a consumer sends after dialing the bus. The bus
// uses event_types + filter to do server-side pushdown (plan §1 decision:
// only events matching this filter are written to this consumer's sendCh).
type Hello struct {
Type FrameType `json:"type"`
ConsumerPID int `json:"consumer_pid"`
EventTypes []string `json:"event_types,omitempty"` // wildcard list ("im.*", "approval.*"); empty = catch-all
Filter string `json:"filter,omitempty"` // optional regex over event_type
SubscribeID string `json:"subscribe_id,omitempty"` // optional personal subscription isolation key; empty = no subscribe_id filter
Compact bool `json:"compact,omitempty"` // hint to status output; bus does not transform payloads
// Role distinguishes a real consumer (registered for events) from an
// ad-hoc tooling connection (status/list/stop). Ad-hoc connections do
// NOT register with the Hub.
Role HelloRole `json:"role,omitempty"`
}
// HelloRole tags the purpose of a Hello connection.
type HelloRole string
const (
HelloRoleConsumer HelloRole = "" // default
HelloRoleStatus HelloRole = "status" // event list / event status
HelloRoleStop HelloRole = "stop" // event stop (graceful trigger)
)
// HelloAck is the bus's reply on accepted Hello. SourceState/StateSource
// mirror the source.Machine snapshot at connect time; CredentialsSource
// fields tell users which channel actually supplied the credentials in use
// (env / app_config / keychain / plain_config), see plan "凭证来源拆字段".
type HelloAck struct {
Type FrameType `json:"type"`
BusPID int `json:"bus_pid"`
SourceState string `json:"source_state"` // mirrors source.State string value
StateSource string `json:"state_source"` // "hook" | "inferred"
ClientIDSource string `json:"client_id_source"` // auth.CredentialSource string
ClientSecretSource string `json:"client_secret_source"` // auth.CredentialSource string
IdleTimeoutSecs int `json:"idle_timeout_secs,omitempty"` // bus's IdleTimeout for diagnostics
}
// Event wraps one delivered RawEvent for the wire. We keep the payload as
// a string (not nested JSON) so the bus does not need to parse / re-encode.
type Event struct {
Type FrameType `json:"type"`
Seq uint64 `json:"seq"` // per-consumer monotonic, restarts at 1 on reconnect
EventID string `json:"event_id"`
EventBornTime int64 `json:"event_born_time"`
EventCorpID string `json:"event_corp_id,omitempty"`
EventType string `json:"event_type"`
EventUnifiedAppID string `json:"event_unified_app_id,omitempty"`
EventScope string `json:"event_scope,omitempty"`
SubscribeID string `json:"subscribe_id,omitempty"`
SourceID string `json:"source_id,omitempty"`
RuleType string `json:"rule_type,omitempty"`
Data string `json:"data"`
Headers map[string]string `json:"headers,omitempty"`
ReceivedAtUnixMS int64 `json:"received_at_unix_ms"`
}
// Heartbeat is bidirectional and stateless. It exists only to give both
// endpoints a chance to notice a dead peer via Read failure.
type Heartbeat struct {
Type FrameType `json:"type"`
}
// SourceState is pushed bus → consume whenever the connection state machine
// transitions to / from connected. Consumers may render it; v1 they just
// forward to stderr when not --quiet.
type SourceState struct {
Type FrameType `json:"type"`
State string `json:"state"` // source.State string
StateSource string `json:"state_source"` // hook | inferred
Attempt int `json:"attempt,omitempty"`
}
// Bye is sent by either side at graceful shutdown. Reason is free-form
// for logs; structured shutdown causes are encoded in the value:
//
// "client_done" — consume reached --max-events/--duration or SIGINT
// "shutdown" — bus SIGTERM/SIGINT
// "idle_timeout" — bus IdleTimeout fired with no consumers
// "stop_request" — bus received explicit Stop RPC
type Bye struct {
Type FrameType `json:"type"`
Reason string `json:"reason"`
}
// StatusReq is an empty JSON frame ad-hoc tooling sends after Hello to
// request a full StatusResp. Bus replies with one StatusResp then closes
// the connection.
type StatusReq struct {
Type FrameType `json:"type"`
}
// StatusResp is the bus's snapshot view for `dws event status` rendering.
// Counts are accumulated since bus start; per-consumer entries are sorted
// by PID for deterministic output.
type StatusResp struct {
Type FrameType `json:"type"`
Bus StatusBus `json:"bus"`
SourceState StatusSource `json:"source_state"`
Consumers []StatusConsumer `json:"consumers"`
PerEventTypeCounters map[string]Counters `json:"per_event_type"`
}
// StatusBus is the bus daemon's identity / lifecycle view.
type StatusBus struct {
PID int `json:"pid"`
UptimeSecs int64 `json:"uptime_secs"`
IdleTimeoutSec int `json:"idle_timeout_secs"`
ClientID string `json:"client_id"`
Edition string `json:"edition"`
SourceKind dwsevent.SourceKind `json:"source_kind,omitempty"`
IdentityHash string `json:"identity_hash,omitempty"`
SourceID string `json:"source_id,omitempty"`
}
// StatusSource is the source.Machine snapshot at status RPC time.
type StatusSource struct {
State string `json:"state"`
Source string `json:"source"`
LastEventAtMS int64 `json:"last_event_at_ms,omitempty"`
LastReconnectMS int64 `json:"last_reconnect_at_ms,omitempty"`
ReconnectCount int `json:"reconnect_count"`
}
// StatusConsumer is one consumer's per-IPC-connection view.
type StatusConsumer struct {
PID int `json:"pid"`
EventTypes []string `json:"event_types,omitempty"`
Filter string `json:"filter,omitempty"`
SubscribeID string `json:"subscribe_id,omitempty"`
SubscribedAtMS int64 `json:"subscribed_at_ms"`
Received uint64 `json:"received"`
Dropped uint64 `json:"dropped"`
}
// Counters is the shared "received / dropped" pair used by per-consumer
// and per-event-type rollups.
type Counters struct {
Received uint64 `json:"received"`
Dropped uint64 `json:"dropped"`
}
// typeOnly is the minimal envelope used by DecodeFrame to peek at the
// "type" field before deciding which concrete struct to unmarshal into.
type typeOnly struct {
Type FrameType `json:"type"`
}
// PeekType returns the FrameType of a raw frame without fully decoding it.
// Returns an error when the JSON is malformed or the type field is missing.
func PeekType(raw []byte) (FrameType, error) {
var t typeOnly
if err := json.Unmarshal(raw, &t); err != nil {
return "", fmt.Errorf("transport: peek frame type: %w", err)
}
if t.Type == "" {
return "", fmt.Errorf("transport: frame missing 'type' field")
}
return t.Type, nil
}
+237
View File
@@ -0,0 +1,237 @@
// 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 (
"encoding/json"
"strings"
"testing"
)
// Each frame type is round-tripped through encoding/json then PeekType
// to guard against accidental wire-format changes (the JSON field tags
// double as part of the cross-version protocol contract).
func TestFrameType_StableWireValues(t *testing.T) {
// If any of these strings change we've made a protocol-breaking
// change. The test value list is duplicated here on purpose so a
// reviewer renaming a constant is forced to also update the test.
wants := map[FrameType]string{
FrameTypeHello: "hello",
FrameTypeHelloAck: "hello_ack",
FrameTypeEvent: "event",
FrameTypeHeartbeat: "heartbeat",
FrameTypeSourceState: "source_state",
FrameTypeBye: "bye",
FrameTypeStatusReq: "status_req",
FrameTypeStatusResp: "status_resp",
}
for ft, want := range wants {
if string(ft) != want {
t.Errorf("FrameType wire value drift: %v != %q", ft, want)
}
}
}
func TestHelloRole_StableWireValues(t *testing.T) {
wants := map[HelloRole]string{
HelloRoleConsumer: "",
HelloRoleStatus: "status",
HelloRoleStop: "stop",
}
for r, want := range wants {
if string(r) != want {
t.Errorf("HelloRole wire value drift: %v != %q", r, want)
}
}
}
func roundTrip(t *testing.T, in any, dst any) {
t.Helper()
b, err := json.Marshal(in)
if err != nil {
t.Fatalf("marshal %T: %v", in, err)
}
if err := json.Unmarshal(b, dst); err != nil {
t.Fatalf("unmarshal %T: %v", in, err)
}
}
func TestHello_Roundtrip(t *testing.T) {
in := Hello{
Type: FrameTypeHello,
ConsumerPID: 42,
EventTypes: []string{"im.*", "approval.task"},
Filter: `^im\.`,
Compact: true,
Role: HelloRoleStatus,
}
var out Hello
roundTrip(t, in, &out)
if out.Type != in.Type || out.ConsumerPID != in.ConsumerPID || out.Filter != in.Filter ||
out.Compact != in.Compact || out.Role != in.Role || len(out.EventTypes) != len(in.EventTypes) {
t.Fatalf("roundtrip mismatch: %+v != %+v", out, in)
}
}
func TestHello_OmitemptyForDefaults(t *testing.T) {
// Default values should NOT appear in the wire form so old/new readers
// stay tolerant of each other (each new field comes in with its
// zero value by default).
in := Hello{Type: FrameTypeHello, ConsumerPID: 1}
b, _ := json.Marshal(in)
s := string(b)
for _, k := range []string{`"event_types"`, `"filter"`, `"compact"`, `"role"`} {
if strings.Contains(s, k) {
t.Errorf("zero-value field %s leaked into wire form: %s", k, s)
}
}
}
func TestHelloAck_Roundtrip(t *testing.T) {
in := HelloAck{
Type: FrameTypeHelloAck,
BusPID: 12345,
SourceState: "connected",
StateSource: "inferred",
ClientIDSource: "env",
ClientSecretSource: "env",
IdleTimeoutSecs: 300,
}
var out HelloAck
roundTrip(t, in, &out)
if out != in {
t.Fatalf("HelloAck roundtrip: %+v != %+v", out, in)
}
}
func TestEvent_Roundtrip(t *testing.T) {
in := Event{
Type: FrameTypeEvent,
Seq: 7,
EventID: "ev_x",
EventBornTime: 1700000000123,
EventCorpID: "corp_x",
EventType: "im.message.receive_v1",
EventUnifiedAppID: "app_y",
Data: `{"text":"hi"}`,
Headers: map[string]string{"extra": "v"},
ReceivedAtUnixMS: 1700000000999,
}
var out Event
roundTrip(t, in, &out)
if out.EventID != in.EventID || out.Seq != in.Seq || out.Data != in.Data ||
out.Headers["extra"] != "v" {
t.Fatalf("Event roundtrip: %+v != %+v", out, in)
}
}
func TestEvent_OmitsZeroOptionalHeader(t *testing.T) {
in := Event{Type: FrameTypeEvent, EventID: "ev_1", EventType: "x", Data: "{}"}
b, _ := json.Marshal(in)
if strings.Contains(string(b), `"headers"`) {
t.Errorf("nil Headers should not appear: %s", b)
}
}
func TestHeartbeat_Roundtrip(t *testing.T) {
in := Heartbeat{Type: FrameTypeHeartbeat}
var out Heartbeat
roundTrip(t, in, &out)
if out.Type != FrameTypeHeartbeat {
t.Fatalf("Heartbeat type lost: %v", out)
}
}
func TestSourceState_Roundtrip(t *testing.T) {
in := SourceState{
Type: FrameTypeSourceState,
State: "reconnecting",
StateSource: "inferred",
Attempt: 3,
}
var out SourceState
roundTrip(t, in, &out)
if out != in {
t.Fatalf("SourceState mismatch: %+v != %+v", out, in)
}
}
func TestBye_Roundtrip(t *testing.T) {
for _, reason := range []string{"client_done", "shutdown", "idle_timeout", "stop_request"} {
in := Bye{Type: FrameTypeBye, Reason: reason}
var out Bye
roundTrip(t, in, &out)
if out.Reason != reason {
t.Errorf("Bye reason %q lost: %+v", reason, out)
}
}
}
func TestStatusResp_Roundtrip(t *testing.T) {
in := StatusResp{
Type: FrameTypeStatusResp,
Bus: StatusBus{
PID: 99, UptimeSecs: 600, IdleTimeoutSec: 300,
ClientID: "ding_x", Edition: "open",
},
SourceState: StatusSource{
State: "connected", Source: "inferred", ReconnectCount: 1,
LastEventAtMS: 1700000000000, LastReconnectMS: 1700000001000,
},
Consumers: []StatusConsumer{
{PID: 12350, EventTypes: []string{"im.*"}, Filter: `^im\.`,
SubscribedAtMS: 1700000000000, Received: 100, Dropped: 2},
},
PerEventTypeCounters: map[string]Counters{
"im.message.receive_v1": {Received: 100, Dropped: 2},
},
}
var out StatusResp
roundTrip(t, in, &out)
if out.Bus.ClientID != in.Bus.ClientID || len(out.Consumers) != 1 ||
out.PerEventTypeCounters["im.message.receive_v1"].Received != 100 {
t.Fatalf("StatusResp roundtrip mismatch: %+v", out)
}
}
// PeekType is used by the daemon to dispatch incoming frames before
// fully decoding into the typed struct. Behaviour at boundaries matters.
func TestPeekType_EachFrameVariant(t *testing.T) {
cases := []struct {
v any
typ FrameType
}{
{Hello{Type: FrameTypeHello}, FrameTypeHello},
{HelloAck{Type: FrameTypeHelloAck}, FrameTypeHelloAck},
{Event{Type: FrameTypeEvent, EventID: "x"}, FrameTypeEvent},
{Heartbeat{Type: FrameTypeHeartbeat}, FrameTypeHeartbeat},
{SourceState{Type: FrameTypeSourceState, State: "connected"}, FrameTypeSourceState},
{Bye{Type: FrameTypeBye, Reason: "x"}, FrameTypeBye},
{StatusReq{Type: FrameTypeStatusReq}, FrameTypeStatusReq},
{StatusResp{Type: FrameTypeStatusResp}, FrameTypeStatusResp},
}
for _, c := range cases {
b, _ := json.Marshal(c.v)
got, err := PeekType(b)
if err != nil {
t.Errorf("PeekType(%T): %v", c.v, err)
continue
}
if got != c.typ {
t.Errorf("PeekType(%T) = %q, want %q", c.v, got, c.typ)
}
}
}
+47
View File
@@ -0,0 +1,47 @@
// 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"
// Listener is the bus-side accept loop. Implementations bind a Unix socket
// (darwin/linux) or Windows Named Pipe at the given Endpoint; calling Close
// removes the underlying socket file on Unix.
type Listener interface {
// Accept blocks until a peer connects or the listener is closed.
Accept() (net.Conn, error)
// Close unbinds the endpoint and unblocks pending Accept calls.
Close() error
// Endpoint returns the human-readable address (Unix path or Windows pipe
// name) — only used in log messages and status output.
Endpoint() string
}
// Listen binds an IPC endpoint at path. On Unix path is a filesystem socket
// path (caller must ensure the parent directory exists with mode 0700).
// On Windows path is a Named Pipe name like `\\.\pipe\dws-event-<edition>-<hash>`.
//
// Stale Unix sockets (left behind by a crashed bus) are unlinked
// automatically before bind. Caller MUST hold the bus.lock before calling
// Listen so this unlink is race-safe against a still-running bus.
func Listen(path string) (Listener, error) {
return listen(path)
}
// Dial connects to an IPC endpoint at path. Returns a standard net.Conn
// the caller can wrap with NewReader/NewWriter. Closes are independent —
// closing the conn does not close the listener.
func Dial(path string) (net.Conn, error) {
return dial(path)
}
@@ -0,0 +1,79 @@
// 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.
//go:build !windows
package transport
import (
"fmt"
"net"
"os"
dwsevent "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
)
type unixListener struct {
l net.Listener
path string
}
func (u *unixListener) Accept() (net.Conn, error) { return u.l.Accept() }
func (u *unixListener) Endpoint() string { return u.path }
func (u *unixListener) Close() error {
err := u.l.Close()
// Best-effort unlink. Ignored errors here because the bus may have
// already been replaced by a competing bus that unlinked first.
_ = os.Remove(u.path)
return err
}
// checkSocketPath rejects paths over the sun_path budget up front so the
// caller sees the actual problem instead of the syscall's bare EINVAL
// ("invalid argument").
func checkSocketPath(path string) error {
if max := dwsevent.MaxUnixSocketPath(); len(path) > max {
return fmt.Errorf("transport: unix socket path too long (%d > %d bytes): %s", len(path), max, path)
}
return nil
}
func listen(path string) (Listener, error) {
if err := checkSocketPath(path); err != nil {
return nil, err
}
// Stale socket cleanup. Caller holds bus.lock so this is race-safe.
if _, err := os.Stat(path); err == nil {
if err := os.Remove(path); err != nil {
return nil, fmt.Errorf("transport: remove stale socket %s: %w", path, err)
}
}
l, err := net.Listen("unix", path)
if err != nil {
return nil, fmt.Errorf("transport: listen %s: %w", path, err)
}
if err := os.Chmod(path, config.FilePerm); err != nil {
_ = l.Close()
_ = os.Remove(path)
return nil, fmt.Errorf("transport: chmod %s: %w", path, err)
}
return &unixListener{l: l, path: path}, nil
}
func dial(path string) (net.Conn, error) {
if err := checkSocketPath(path); err != nil {
return nil, err
}
return net.Dial("unix", path)
}

Some files were not shown because too many files have changed in this diff Show More