Compare commits
42
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
0137a1950b | ||
|
|
ed0ec3103a | ||
|
|
8f0bb91ddb | ||
|
|
b5be003330 | ||
|
|
dd29805063 | ||
|
|
497b627080 | ||
|
|
4233105963 | ||
|
|
2541081f97 | ||
|
|
3a6e5fb4b3 | ||
|
|
83783972a3 | ||
|
|
919b9d698f | ||
|
|
ea77270030 | ||
|
|
1a62ef3616 | ||
|
|
7d37c15a46 | ||
|
|
9fa763e213 | ||
|
|
ee0e629c72 | ||
|
|
2c8f970703 | ||
|
|
20d0eaa489 | ||
|
|
67c88828f0 | ||
|
|
b015c27034 | ||
|
|
604ec5f50a | ||
|
|
27296ec426 | ||
|
|
6a38a168dd | ||
|
|
9771053d81 | ||
|
|
10c0c5083e | ||
|
|
a0187b5297 | ||
|
|
81991f1c07 | ||
|
|
836670ef50 | ||
|
|
53ce0a8303 | ||
|
|
78867f3601 | ||
|
|
37438659e6 | ||
|
|
3c12c835a3 | ||
|
|
389f83241f | ||
|
|
3714adc2db | ||
|
|
f37d0569a1 | ||
|
|
798b58bf3c | ||
|
|
d926bed3cc | ||
|
|
5ac180d3dd | ||
|
|
35e60407d3 | ||
|
|
ea46132cf6 | ||
|
|
478dc155e8 | ||
|
|
08ecb38a42 |
@@ -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>
|
||||
|
||||
|
||||
@@ -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>
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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))
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -356,6 +356,7 @@ func NewRootCommandWithEngine(rootCtx context.Context, engine *pipeline.Engine)
|
||||
newCatalogCommand(loader),
|
||||
newConfigCommand(),
|
||||
newDoctorCommand(),
|
||||
newEventCommand(),
|
||||
newCompletionCommand(root),
|
||||
newRecoveryCommand(rootCtx, loader, flags),
|
||||
newUpgradeCommand(),
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
@@ -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")
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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,
|
||||
}
|
||||
}
|
||||
@@ -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,
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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 }
|
||||
@@ -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 }
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
@@ -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())
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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())
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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())
|
||||
}
|
||||
}
|
||||
@@ -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 }
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
@@ -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")
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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 ""
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
@@ -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>"
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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] + "..."
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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())
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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
Reference in New Issue
Block a user