fix(commands): avoid runtime-missing regression in session handlers

This commit is contained in:
mingmxren 2026-03-01 15:25:03 +08:00
parent be3010efc1
commit bb4a0dc766
3 changed files with 233 additions and 60 deletions

View file

@ -9,25 +9,30 @@ import (
"github.com/sipeed/picoclaw/pkg/config" "github.com/sipeed/picoclaw/pkg/config"
) )
type runtimeContextKey struct{}
// WithRuntime attaches command runtime capabilities to ctx for command handlers.
func WithRuntime(ctx context.Context, runtime Runtime) context.Context {
if ctx == nil {
ctx = context.Background()
}
return context.WithValue(ctx, runtimeContextKey{}, runtime)
}
func runtimeFromContext(ctx context.Context) Runtime {
if ctx == nil {
return nil
}
runtime, _ := ctx.Value(runtimeContextKey{}).(Runtime)
return runtime
}
func BuiltinDefinitions(cfg *config.Config) []Definition { func BuiltinDefinitions(cfg *config.Config) []Definition {
return builtinDefinitions(cfg, nil)
}
// BuiltinDefinitionsWithRuntime returns builtin command definitions with runtime-backed
// session command handlers enabled only when runtime is usable.
func BuiltinDefinitionsWithRuntime(cfg *config.Config, runtime Runtime) []Definition {
return builtinDefinitions(cfg, runtime)
}
func builtinDefinitions(cfg *config.Config, runtime Runtime) []Definition {
sessionRuntime := runtimeIfUsable(runtime)
var newHandler Handler
var sessionHandler Handler
if sessionRuntime != nil {
newHandler = func(_ context.Context, req Request) error {
return handleNewCommand(req, sessionRuntime, cfg)
}
sessionHandler = func(_ context.Context, req Request) error {
return handleSessionCommand(req, sessionRuntime)
}
}
return []Definition{ return []Definition{
{ {
Name: "start", Name: "start",
@ -55,18 +60,14 @@ func BuiltinDefinitions(cfg *config.Config) []Definition {
Description: "Start a new chat session", Description: "Start a new chat session",
Usage: "/new", Usage: "/new",
Channels: []string{"telegram", "whatsapp", "whatsapp_native"}, Channels: []string{"telegram", "whatsapp", "whatsapp_native"},
Handler: func(ctx context.Context, req Request) error { Handler: newHandler,
return handleNewCommand(ctx, req, cfg)
},
}, },
{ {
Name: "session", Name: "session",
Description: "Manage chat sessions", Description: "Manage chat sessions",
Usage: "/session [list|resume <index>]", Usage: "/session [list|resume <index>]",
Channels: []string{"telegram", "whatsapp", "whatsapp_native"}, Channels: []string{"telegram", "whatsapp", "whatsapp_native"},
Handler: func(ctx context.Context, req Request) error { Handler: sessionHandler,
return handleSessionCommand(ctx, req)
},
}, },
{ {
Name: "show", Name: "show",
@ -172,12 +173,7 @@ func replyText(text string) Handler {
} }
} }
func handleNewCommand(ctx context.Context, req Request, fallbackCfg *config.Config) error { func handleNewCommand(req Request, runtime Runtime, fallbackCfg *config.Config) error {
runtime := runtimeFromContext(ctx)
if runtime == nil || runtime.SessionOps() == nil || strings.TrimSpace(runtime.ScopeKey()) == "" {
return reply(req, "Command unavailable in current context.")
}
scopeKey := runtime.ScopeKey() scopeKey := runtime.ScopeKey()
newSessionKey, err := runtime.SessionOps().StartNew(scopeKey) newSessionKey, err := runtime.SessionOps().StartNew(scopeKey)
if err != nil { if err != nil {
@ -208,12 +204,7 @@ func handleNewCommand(ctx context.Context, req Request, fallbackCfg *config.Conf
return reply(req, fmt.Sprintf("Started new session: %s (pruned %d old session(s))", newSessionKey, len(pruned))) return reply(req, fmt.Sprintf("Started new session: %s (pruned %d old session(s))", newSessionKey, len(pruned)))
} }
func handleSessionCommand(ctx context.Context, req Request) error { func handleSessionCommand(req Request, runtime Runtime) error {
runtime := runtimeFromContext(ctx)
if runtime == nil || runtime.SessionOps() == nil || strings.TrimSpace(runtime.ScopeKey()) == "" {
return reply(req, "Command unavailable in current context.")
}
args := strings.Fields(commandArgs(req.Text)) args := strings.Fields(commandArgs(req.Text))
if len(args) < 1 { if len(args) < 1 {
return reply(req, "Usage: /session [list|resume <index>]") return reply(req, "Usage: /session [list|resume <index>]")
@ -278,6 +269,19 @@ func reply(req Request, text string) error {
return req.Reply(text) return req.Reply(text)
} }
func runtimeIfUsable(runtime Runtime) Runtime {
if runtime == nil {
return nil
}
if runtime.SessionOps() == nil {
return nil
}
if strings.TrimSpace(runtime.ScopeKey()) == "" {
return nil
}
return runtime
}
func enabledChannels(cfg *config.Config) []string { func enabledChannels(cfg *config.Config) []string {
enabled := make([]string, 0, 8) enabled := make([]string, 0, 8)
if cfg.Channels.Telegram.Enabled { if cfg.Channels.Telegram.Enabled {

View file

@ -1,6 +1,11 @@
package commands package commands
import "testing" import (
"testing"
"github.com/sipeed/picoclaw/pkg/config"
"github.com/sipeed/picoclaw/pkg/session"
)
func TestBuiltinDefinitions_ContainsTelegramDefaults(t *testing.T) { func TestBuiltinDefinitions_ContainsTelegramDefaults(t *testing.T) {
defs := BuiltinDefinitions(nil) defs := BuiltinDefinitions(nil)
@ -32,7 +37,7 @@ func TestBuiltinDefinitions_WhatsAppOnlyHasBasicCommands(t *testing.T) {
} }
} }
func TestBuiltinDefinitions_SessionCommandsHaveHandlers(t *testing.T) { func TestBuiltinDefinitions_DefaultSessionCommandsArePassthrough(t *testing.T) {
defs := BuiltinDefinitions(nil) defs := BuiltinDefinitions(nil)
defByName := map[string]Definition{} defByName := map[string]Definition{}
@ -44,8 +49,8 @@ func TestBuiltinDefinitions_SessionCommandsHaveHandlers(t *testing.T) {
if !ok { if !ok {
t.Fatalf("missing /new definition") t.Fatalf("missing /new definition")
} }
if newDef.Handler == nil { if newDef.Handler != nil {
t.Fatalf("/new should provide a runtime-backed handler") t.Fatalf("/new should be passthrough without runtime wiring")
} }
if !contains(newDef.Aliases, "reset") { if !contains(newDef.Aliases, "reset") {
t.Fatalf("/new aliases=%v, want alias \"reset\"", newDef.Aliases) t.Fatalf("/new aliases=%v, want alias \"reset\"", newDef.Aliases)
@ -55,7 +60,44 @@ func TestBuiltinDefinitions_SessionCommandsHaveHandlers(t *testing.T) {
if !ok { if !ok {
t.Fatalf("missing /session definition") t.Fatalf("missing /session definition")
} }
if sessionDef.Handler == nil { if sessionDef.Handler != nil {
t.Fatalf("/session should provide a runtime-backed handler") t.Fatalf("/session should be passthrough without runtime wiring")
}
}
type builtinTestSessionOps struct{}
func (f *builtinTestSessionOps) ResolveActive(scopeKey string) (string, error) { return "", nil }
func (f *builtinTestSessionOps) StartNew(scopeKey string) (string, error) { return "", nil }
func (f *builtinTestSessionOps) List(scopeKey string) ([]session.SessionMeta, error) {
return nil, nil
}
func (f *builtinTestSessionOps) Resume(scopeKey string, index int) (string, error) { return "", nil }
func (f *builtinTestSessionOps) Prune(scopeKey string, limit int) ([]string, error) { return nil, nil }
type builtinTestRuntime struct {
scope string
ops SessionOps
}
func (f *builtinTestRuntime) Channel() string { return "whatsapp" }
func (f *builtinTestRuntime) ScopeKey() string { return f.scope }
func (f *builtinTestRuntime) SessionOps() SessionOps { return f.ops }
func (f *builtinTestRuntime) Config() *config.Config { return nil }
func TestBuiltinDefinitionsWithRuntime_EnablesSessionHandlers(t *testing.T) {
runtime := &builtinTestRuntime{scope: "scope", ops: &builtinTestSessionOps{}}
defs := BuiltinDefinitionsWithRuntime(nil, runtime)
defByName := map[string]Definition{}
for _, def := range defs {
defByName[def.Name] = def
}
if defByName["new"].Handler == nil {
t.Fatalf("/new should provide runtime-backed handler when runtime is available")
}
if defByName["session"].Handler == nil {
t.Fatalf("/session should provide runtime-backed handler when runtime is available")
} }
} }

View file

@ -2,6 +2,7 @@ package commands
import ( import (
"context" "context"
"errors"
"testing" "testing"
"time" "time"
@ -79,8 +80,6 @@ func (f *sessionHandlerFakeRuntime) Config() *config.Config {
} }
func TestSessionHandlers_New_UsesRuntimeSessionOps(t *testing.T) { func TestSessionHandlers_New_UsesRuntimeSessionOps(t *testing.T) {
t.Helper()
ops := &sessionHandlerFakeSessionOps{ ops := &sessionHandlerFakeSessionOps{
startNewValue: "scope#2", startNewValue: "scope#2",
pruneValue: []string{"scope#1"}, pruneValue: []string{"scope#1"},
@ -94,11 +93,9 @@ func TestSessionHandlers_New_UsesRuntimeSessionOps(t *testing.T) {
}, },
} }
ctx := WithRuntime(context.Background(), runtime)
var reply string var reply string
ex := NewExecutor(NewRegistry(BuiltinDefinitions(nil))) ex := NewExecutor(NewRegistry(BuiltinDefinitionsWithRuntime(nil, runtime)))
res := ex.Execute(ctx, Request{ res := ex.Execute(context.Background(), Request{
Channel: "whatsapp", Channel: "whatsapp",
Text: "/new", Text: "/new",
Reply: func(text string) error { Reply: func(text string) error {
@ -125,8 +122,6 @@ func TestSessionHandlers_New_UsesRuntimeSessionOps(t *testing.T) {
} }
func TestSessionHandlers_SessionResume_UsesRuntimeSessionOps(t *testing.T) { func TestSessionHandlers_SessionResume_UsesRuntimeSessionOps(t *testing.T) {
t.Helper()
ops := &sessionHandlerFakeSessionOps{resumeValue: "scope#3"} ops := &sessionHandlerFakeSessionOps{resumeValue: "scope#3"}
runtime := &sessionHandlerFakeRuntime{ runtime := &sessionHandlerFakeRuntime{
channel: "whatsapp", channel: "whatsapp",
@ -135,11 +130,9 @@ func TestSessionHandlers_SessionResume_UsesRuntimeSessionOps(t *testing.T) {
cfg: &config.Config{}, cfg: &config.Config{},
} }
ctx := WithRuntime(context.Background(), runtime)
var reply string var reply string
ex := NewExecutor(NewRegistry(BuiltinDefinitions(nil))) ex := NewExecutor(NewRegistry(BuiltinDefinitionsWithRuntime(nil, runtime)))
res := ex.Execute(ctx, Request{ res := ex.Execute(context.Background(), Request{
Channel: "whatsapp", Channel: "whatsapp",
Text: "/session resume 3", Text: "/session resume 3",
Reply: func(text string) error { Reply: func(text string) error {
@ -163,8 +156,6 @@ func TestSessionHandlers_SessionResume_UsesRuntimeSessionOps(t *testing.T) {
} }
func TestSessionHandlers_SessionList_UsesRuntimeSessionOps(t *testing.T) { func TestSessionHandlers_SessionList_UsesRuntimeSessionOps(t *testing.T) {
t.Helper()
ops := &sessionHandlerFakeSessionOps{ ops := &sessionHandlerFakeSessionOps{
listValue: []session.SessionMeta{ listValue: []session.SessionMeta{
{ {
@ -183,11 +174,9 @@ func TestSessionHandlers_SessionList_UsesRuntimeSessionOps(t *testing.T) {
cfg: &config.Config{}, cfg: &config.Config{},
} }
ctx := WithRuntime(context.Background(), runtime)
var reply string var reply string
ex := NewExecutor(NewRegistry(BuiltinDefinitions(nil))) ex := NewExecutor(NewRegistry(BuiltinDefinitionsWithRuntime(nil, runtime)))
res := ex.Execute(ctx, Request{ res := ex.Execute(context.Background(), Request{
Channel: "whatsapp", Channel: "whatsapp",
Text: "/session list", Text: "/session list",
Reply: func(text string) error { Reply: func(text string) error {
@ -206,3 +195,141 @@ func TestSessionHandlers_SessionList_UsesRuntimeSessionOps(t *testing.T) {
t.Fatalf("reply=%q", reply) t.Fatalf("reply=%q", reply)
} }
} }
func TestSessionHandlers_MissingRuntime_Passthrough(t *testing.T) {
ex := NewExecutor(NewRegistry(BuiltinDefinitionsWithRuntime(nil, nil)))
for _, input := range []string{"/new", "/session list"} {
res := ex.Execute(context.Background(), Request{
Channel: "whatsapp",
Text: input,
})
if res.Outcome != OutcomePassthrough {
t.Fatalf("text=%q outcome=%v, want=%v", input, res.Outcome, OutcomePassthrough)
}
}
}
func TestSessionHandlers_NilSessionOps_Passthrough(t *testing.T) {
runtime := &sessionHandlerFakeRuntime{
channel: "whatsapp",
scope: "scope",
ops: nil,
cfg: &config.Config{},
}
ex := NewExecutor(NewRegistry(BuiltinDefinitionsWithRuntime(nil, runtime)))
for _, input := range []string{"/new", "/session list"} {
res := ex.Execute(context.Background(), Request{
Channel: "whatsapp",
Text: input,
})
if res.Outcome != OutcomePassthrough {
t.Fatalf("text=%q outcome=%v, want=%v", input, res.Outcome, OutcomePassthrough)
}
}
}
func TestSessionHandlers_EmptyScope_Passthrough(t *testing.T) {
runtime := &sessionHandlerFakeRuntime{
channel: "whatsapp",
scope: " ",
ops: &sessionHandlerFakeSessionOps{},
cfg: &config.Config{},
}
ex := NewExecutor(NewRegistry(BuiltinDefinitionsWithRuntime(nil, runtime)))
for _, input := range []string{"/new", "/session list"} {
res := ex.Execute(context.Background(), Request{
Channel: "whatsapp",
Text: input,
})
if res.Outcome != OutcomePassthrough {
t.Fatalf("text=%q outcome=%v, want=%v", input, res.Outcome, OutcomePassthrough)
}
}
}
func TestSessionHandlers_ErrorAndValidationReplies(t *testing.T) {
tests := []struct {
name string
text string
ops *sessionHandlerFakeSessionOps
wantReply string
}{
{
name: "start new error",
text: "/new",
ops: &sessionHandlerFakeSessionOps{startNewErr: errors.New("boom")},
wantReply: "Failed to start new session: boom",
},
{
name: "prune error",
text: "/new",
ops: &sessionHandlerFakeSessionOps{
startNewValue: "scope#2",
pruneErr: errors.New("prune failed"),
},
wantReply: "Started new session (scope#2), but pruning old sessions failed: prune failed",
},
{
name: "list error",
text: "/session list",
ops: &sessionHandlerFakeSessionOps{listErr: errors.New("list failed")},
wantReply: "Failed to list sessions: list failed",
},
{
name: "resume error",
text: "/session resume 2",
ops: &sessionHandlerFakeSessionOps{resumeErr: errors.New("resume failed")},
wantReply: "Failed to resume session 2: resume failed",
},
{
name: "resume missing index",
text: "/session resume",
ops: &sessionHandlerFakeSessionOps{},
wantReply: "Usage: /session resume <index>",
},
{
name: "resume non numeric index",
text: "/session resume abc",
ops: &sessionHandlerFakeSessionOps{},
wantReply: "Usage: /session resume <index>",
},
{
name: "resume zero index",
text: "/session resume 0",
ops: &sessionHandlerFakeSessionOps{},
wantReply: "Usage: /session resume <index>",
},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
runtime := &sessionHandlerFakeRuntime{
channel: "whatsapp",
scope: "scope",
ops: tc.ops,
cfg: &config.Config{},
}
var reply string
ex := NewExecutor(NewRegistry(BuiltinDefinitionsWithRuntime(nil, runtime)))
res := ex.Execute(context.Background(), Request{
Channel: "whatsapp",
Text: tc.text,
Reply: func(text string) error {
reply = text
return nil
},
})
if res.Outcome != OutcomeHandled {
t.Fatalf("outcome=%v, want=%v", res.Outcome, OutcomeHandled)
}
if reply != tc.wantReply {
t.Fatalf("reply=%q, want=%q", reply, tc.wantReply)
}
})
}
}