feat: add WebSocket channel tests and stubs (TDD red phase)
Test-first scaffolding for outbound WebSocket channel that connects picoclaw agents to an external orchestration service. Tests compile and run but fail where implementation is needed (Start, Send, listen, reconnect). Includes a smoke-test orchestrator mock in examples/. Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
parent
946af6b53d
commit
d49723e380
5 changed files with 1015 additions and 0 deletions
224
examples/ws-orchestrator/main.go
Normal file
224
examples/ws-orchestrator/main.go
Normal file
|
|
@ -0,0 +1,224 @@
|
|||
// ws-orchestrator is a minimal WebSocket server that mimics an orchestration
|
||||
// service for smoke-testing the picoclaw WebSocket channel.
|
||||
//
|
||||
// Usage:
|
||||
//
|
||||
// go run ./examples/ws-orchestrator -addr :9090 -token secret
|
||||
//
|
||||
// Then configure picoclaw with:
|
||||
//
|
||||
// "websocket": {
|
||||
// "enabled": true,
|
||||
// "ws_url": "ws://localhost:9090/ws",
|
||||
// "access_token": "secret",
|
||||
// "agent_id": "researcher",
|
||||
// "reconnect_interval": 5
|
||||
// }
|
||||
//
|
||||
// The orchestrator will:
|
||||
// 1. Validate the auth handshake (agent_id + token)
|
||||
// 2. Send a test message after successful auth
|
||||
// 3. Print any messages received from the agent
|
||||
// 4. Read lines from stdin and send them as inbound messages
|
||||
package main
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"encoding/json"
|
||||
"flag"
|
||||
"fmt"
|
||||
"log"
|
||||
"net/http"
|
||||
"os"
|
||||
"sync"
|
||||
|
||||
"github.com/gorilla/websocket"
|
||||
)
|
||||
|
||||
type WSEnvelope struct {
|
||||
Type string `json:"type"`
|
||||
Data json.RawMessage `json:"data"`
|
||||
}
|
||||
|
||||
type WSAuthData struct {
|
||||
AgentID string `json:"agent_id"`
|
||||
Token string `json:"token"`
|
||||
}
|
||||
|
||||
type WSInboundMessage struct {
|
||||
SenderID string `json:"sender_id"`
|
||||
ChatID string `json:"chat_id"`
|
||||
Content string `json:"content"`
|
||||
SessionKey string `json:"session_key,omitempty"`
|
||||
MessageID string `json:"message_id,omitempty"`
|
||||
Metadata map[string]string `json:"metadata,omitempty"`
|
||||
}
|
||||
|
||||
type WSOutboundMessage struct {
|
||||
Channel string `json:"channel"`
|
||||
ChatID string `json:"chat_id"`
|
||||
Content string `json:"content"`
|
||||
ContentType string `json:"content_type"`
|
||||
}
|
||||
|
||||
var (
|
||||
addr = flag.String("addr", ":9090", "listen address")
|
||||
token = flag.String("token", "secret", "expected access token")
|
||||
)
|
||||
|
||||
var upgrader = websocket.Upgrader{CheckOrigin: func(r *http.Request) bool { return true }}
|
||||
|
||||
// connMu guards activeConn so stdin sender and HTTP handler don't race.
|
||||
var (
|
||||
connMu sync.Mutex
|
||||
activeConn *websocket.Conn
|
||||
)
|
||||
|
||||
func main() {
|
||||
flag.Parse()
|
||||
|
||||
http.HandleFunc("/ws", handleWS)
|
||||
|
||||
// Read stdin in background for interactive message sending.
|
||||
go stdinLoop()
|
||||
|
||||
log.Printf("ws-orchestrator listening on %s/ws (token=%q)", *addr, *token)
|
||||
log.Fatal(http.ListenAndServe(*addr, nil))
|
||||
}
|
||||
|
||||
func handleWS(w http.ResponseWriter, r *http.Request) {
|
||||
conn, err := upgrader.Upgrade(w, r, nil)
|
||||
if err != nil {
|
||||
log.Printf("upgrade error: %v", err)
|
||||
return
|
||||
}
|
||||
defer func() {
|
||||
connMu.Lock()
|
||||
if activeConn == conn {
|
||||
activeConn = nil
|
||||
}
|
||||
connMu.Unlock()
|
||||
conn.Close()
|
||||
}()
|
||||
|
||||
log.Println("agent connected")
|
||||
|
||||
// 1. Read auth envelope
|
||||
_, data, err := conn.ReadMessage()
|
||||
if err != nil {
|
||||
log.Printf("read auth: %v", err)
|
||||
return
|
||||
}
|
||||
|
||||
var env WSEnvelope
|
||||
if err := json.Unmarshal(data, &env); err != nil {
|
||||
log.Printf("bad envelope: %v", err)
|
||||
return
|
||||
}
|
||||
if env.Type != "auth" {
|
||||
log.Printf("expected auth, got %q", env.Type)
|
||||
conn.WriteMessage(websocket.CloseMessage,
|
||||
websocket.FormatCloseMessage(websocket.ClosePolicyViolation, "expected auth"))
|
||||
return
|
||||
}
|
||||
|
||||
var auth WSAuthData
|
||||
if err := json.Unmarshal(env.Data, &auth); err != nil {
|
||||
log.Printf("bad auth data: %v", err)
|
||||
return
|
||||
}
|
||||
if auth.Token != *token {
|
||||
log.Printf("bad token from %s: %q", auth.AgentID, auth.Token)
|
||||
conn.WriteMessage(websocket.CloseMessage,
|
||||
websocket.FormatCloseMessage(websocket.ClosePolicyViolation, "bad token"))
|
||||
return
|
||||
}
|
||||
|
||||
log.Printf("auth OK: agent_id=%s", auth.AgentID)
|
||||
|
||||
connMu.Lock()
|
||||
activeConn = conn
|
||||
connMu.Unlock()
|
||||
|
||||
// 2. Send a test message
|
||||
sendInbound(conn, WSInboundMessage{
|
||||
SenderID: "orchestrator",
|
||||
ChatID: "smoke-test",
|
||||
Content: fmt.Sprintf("Hello %s, this is a smoke test from the orchestrator.", auth.AgentID),
|
||||
SessionKey: "smoke-001",
|
||||
MessageID: "msg-001",
|
||||
})
|
||||
|
||||
// 3. Read loop — print messages from agent
|
||||
for {
|
||||
_, data, err := conn.ReadMessage()
|
||||
if err != nil {
|
||||
log.Printf("agent disconnected: %v", err)
|
||||
return
|
||||
}
|
||||
var env WSEnvelope
|
||||
if err := json.Unmarshal(data, &env); err != nil {
|
||||
log.Printf("bad envelope from agent: %v", err)
|
||||
continue
|
||||
}
|
||||
switch env.Type {
|
||||
case "message":
|
||||
var out WSOutboundMessage
|
||||
if err := json.Unmarshal(env.Data, &out); err != nil {
|
||||
log.Printf("bad outbound message: %v", err)
|
||||
continue
|
||||
}
|
||||
log.Printf("[%s] chat=%s type=%s content=%s",
|
||||
out.Channel, out.ChatID, out.ContentType, out.Content)
|
||||
default:
|
||||
log.Printf("unknown envelope type %q: %s", env.Type, string(env.Data))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func sendInbound(conn *websocket.Conn, msg WSInboundMessage) {
|
||||
data, err := json.Marshal(msg)
|
||||
if err != nil {
|
||||
log.Printf("marshal inbound: %v", err)
|
||||
return
|
||||
}
|
||||
env := WSEnvelope{Type: "message", Data: json.RawMessage(data)}
|
||||
envData, err := json.Marshal(env)
|
||||
if err != nil {
|
||||
log.Printf("marshal envelope: %v", err)
|
||||
return
|
||||
}
|
||||
if err := conn.WriteMessage(websocket.TextMessage, envData); err != nil {
|
||||
log.Printf("write inbound: %v", err)
|
||||
return
|
||||
}
|
||||
log.Printf("→ sent: chat=%s content=%s", msg.ChatID, msg.Content)
|
||||
}
|
||||
|
||||
// stdinLoop reads lines from stdin and sends them as inbound messages.
|
||||
func stdinLoop() {
|
||||
scanner := bufio.NewScanner(os.Stdin)
|
||||
fmt.Println("Type a message and press Enter to send to the agent (chat_id=interactive):")
|
||||
msgID := 0
|
||||
for scanner.Scan() {
|
||||
line := scanner.Text()
|
||||
if line == "" {
|
||||
continue
|
||||
}
|
||||
connMu.Lock()
|
||||
conn := activeConn
|
||||
connMu.Unlock()
|
||||
if conn == nil {
|
||||
fmt.Println("(no agent connected)")
|
||||
continue
|
||||
}
|
||||
msgID++
|
||||
sendInbound(conn, WSInboundMessage{
|
||||
SenderID: "orchestrator",
|
||||
ChatID: "interactive",
|
||||
Content: line,
|
||||
SessionKey: "interactive",
|
||||
MessageID: fmt.Sprintf("stdin-%d", msgID),
|
||||
})
|
||||
}
|
||||
}
|
||||
13
pkg/channels/websocket/init.go
Normal file
13
pkg/channels/websocket/init.go
Normal file
|
|
@ -0,0 +1,13 @@
|
|||
package websocket
|
||||
|
||||
import (
|
||||
"github.com/sipeed/picoclaw/pkg/bus"
|
||||
"github.com/sipeed/picoclaw/pkg/channels"
|
||||
"github.com/sipeed/picoclaw/pkg/config"
|
||||
)
|
||||
|
||||
func init() {
|
||||
channels.RegisterFactory("websocket", func(cfg *config.Config, b *bus.MessageBus) (channels.Channel, error) {
|
||||
return NewWebSocketChannel(cfg.Channels.WebSocket, b)
|
||||
})
|
||||
}
|
||||
100
pkg/channels/websocket/websocket.go
Normal file
100
pkg/channels/websocket/websocket.go
Normal file
|
|
@ -0,0 +1,100 @@
|
|||
package websocket
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"sync"
|
||||
|
||||
"github.com/gorilla/websocket"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/bus"
|
||||
"github.com/sipeed/picoclaw/pkg/channels"
|
||||
"github.com/sipeed/picoclaw/pkg/config"
|
||||
)
|
||||
|
||||
// WSEnvelope is the typed envelope for all WebSocket messages.
|
||||
type WSEnvelope struct {
|
||||
Type string `json:"type"`
|
||||
Data json.RawMessage `json:"data"`
|
||||
}
|
||||
|
||||
// WSAuthData is the auth handshake payload sent on connect.
|
||||
type WSAuthData struct {
|
||||
AgentID string `json:"agent_id"`
|
||||
Token string `json:"token"`
|
||||
}
|
||||
|
||||
// WSInboundMessage is the inbound message payload from the orchestrator.
|
||||
type WSInboundMessage struct {
|
||||
SenderID string `json:"sender_id"`
|
||||
ChatID string `json:"chat_id"`
|
||||
Content string `json:"content"`
|
||||
SessionKey string `json:"session_key,omitempty"`
|
||||
MessageID string `json:"message_id,omitempty"`
|
||||
Metadata map[string]string `json:"metadata,omitempty"`
|
||||
}
|
||||
|
||||
// WSOutboundMessage wraps bus.OutboundMessage with content_type for the wire.
|
||||
type WSOutboundMessage struct {
|
||||
Channel string `json:"channel"`
|
||||
ChatID string `json:"chat_id"`
|
||||
Content string `json:"content"`
|
||||
ContentType string `json:"content_type"`
|
||||
}
|
||||
|
||||
// WebSocketChannel connects outward to an orchestrator via WebSocket.
|
||||
type WebSocketChannel struct {
|
||||
*channels.BaseChannel
|
||||
config config.WebSocketConfig
|
||||
conn *websocket.Conn
|
||||
ctx context.Context
|
||||
cancel context.CancelFunc
|
||||
mu sync.Mutex // guards conn
|
||||
writeMu sync.Mutex // guards writes
|
||||
}
|
||||
|
||||
// NewWebSocketChannel creates a new WebSocket channel (stub).
|
||||
func NewWebSocketChannel(cfg config.WebSocketConfig, mb *bus.MessageBus) (*WebSocketChannel, error) {
|
||||
base := channels.NewBaseChannel(
|
||||
"websocket",
|
||||
cfg,
|
||||
mb,
|
||||
[]string(cfg.AllowFrom),
|
||||
channels.WithReasoningChannelID(cfg.ReasoningChannelID),
|
||||
)
|
||||
ch := &WebSocketChannel{
|
||||
BaseChannel: base,
|
||||
config: cfg,
|
||||
}
|
||||
return ch, nil
|
||||
}
|
||||
|
||||
func (c *WebSocketChannel) Start(ctx context.Context) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *WebSocketChannel) Stop(ctx context.Context) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *WebSocketChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
|
||||
if !c.IsRunning() {
|
||||
return channels.ErrNotRunning
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// DetectContentType returns "json" if content is valid JSON, otherwise "text".
|
||||
func DetectContentType(content string) string {
|
||||
if json.Valid([]byte(content)) {
|
||||
return "json"
|
||||
}
|
||||
return "text"
|
||||
}
|
||||
|
||||
// parseEnvelope parses raw bytes into a WSEnvelope.
|
||||
func parseEnvelope(data []byte) (WSEnvelope, error) {
|
||||
var env WSEnvelope
|
||||
err := json.Unmarshal(data, &env)
|
||||
return env, err
|
||||
}
|
||||
667
pkg/channels/websocket/websocket_test.go
Normal file
667
pkg/channels/websocket/websocket_test.go
Normal file
|
|
@ -0,0 +1,667 @@
|
|||
package websocket
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/gorilla/websocket"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/bus"
|
||||
"github.com/sipeed/picoclaw/pkg/channels"
|
||||
"github.com/sipeed/picoclaw/pkg/config"
|
||||
)
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// helpers
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
var upgrader = websocket.Upgrader{CheckOrigin: func(r *http.Request) bool { return true }}
|
||||
|
||||
// newTestChannel creates a WebSocketChannel with the given config and a fresh bus.
|
||||
func newTestChannel(t *testing.T, cfg config.WebSocketConfig) (*WebSocketChannel, *bus.MessageBus) {
|
||||
t.Helper()
|
||||
mb := bus.NewMessageBus()
|
||||
t.Cleanup(func() { mb.Close() })
|
||||
ch, err := NewWebSocketChannel(cfg, mb)
|
||||
if err != nil {
|
||||
t.Fatalf("NewWebSocketChannel: %v", err)
|
||||
}
|
||||
return ch, mb
|
||||
}
|
||||
|
||||
// mockServer starts an httptest server that upgrades to WebSocket and calls handler
|
||||
// with each connection. Returns the server and its ws:// URL.
|
||||
func mockServer(t *testing.T, handler func(*websocket.Conn)) (*httptest.Server, string) {
|
||||
t.Helper()
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
conn, err := upgrader.Upgrade(w, r, nil)
|
||||
if err != nil {
|
||||
t.Logf("upgrade error: %v", err)
|
||||
return
|
||||
}
|
||||
handler(conn)
|
||||
}))
|
||||
t.Cleanup(srv.Close)
|
||||
wsURL := "ws" + strings.TrimPrefix(srv.URL, "http")
|
||||
return srv, wsURL
|
||||
}
|
||||
|
||||
// readEnvelope reads one WSEnvelope from the connection with a timeout.
|
||||
func readEnvelope(t *testing.T, conn *websocket.Conn, timeout time.Duration) WSEnvelope {
|
||||
t.Helper()
|
||||
conn.SetReadDeadline(time.Now().Add(timeout))
|
||||
_, data, err := conn.ReadMessage()
|
||||
if err != nil {
|
||||
t.Fatalf("readEnvelope: %v", err)
|
||||
}
|
||||
var env WSEnvelope
|
||||
if err := json.Unmarshal(data, &env); err != nil {
|
||||
t.Fatalf("readEnvelope unmarshal: %v", err)
|
||||
}
|
||||
return env
|
||||
}
|
||||
|
||||
// writeEnvelope writes a WSEnvelope to the connection.
|
||||
func writeEnvelope(t *testing.T, conn *websocket.Conn, env WSEnvelope) {
|
||||
t.Helper()
|
||||
data, err := json.Marshal(env)
|
||||
if err != nil {
|
||||
t.Fatalf("writeEnvelope marshal: %v", err)
|
||||
}
|
||||
if err := conn.WriteMessage(websocket.TextMessage, data); err != nil {
|
||||
t.Fatalf("writeEnvelope write: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// makeMessageEnvelope builds a WSEnvelope with type "message" from the given inbound fields.
|
||||
func makeMessageEnvelope(t *testing.T, msg WSInboundMessage) WSEnvelope {
|
||||
t.Helper()
|
||||
data, err := json.Marshal(msg)
|
||||
if err != nil {
|
||||
t.Fatalf("makeMessageEnvelope: %v", err)
|
||||
}
|
||||
return WSEnvelope{Type: "message", Data: json.RawMessage(data)}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Unit Tests
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestNewWebSocketChannel(t *testing.T) {
|
||||
cfg := config.WebSocketConfig{
|
||||
Enabled: true,
|
||||
WSUrl: "ws://localhost:9999",
|
||||
AccessToken: "tok",
|
||||
AgentID: "researcher",
|
||||
ReconnectInterval: 5,
|
||||
AllowFrom: config.FlexibleStringSlice{"orchestrator"},
|
||||
ReasoningChannelID: "reason-chan",
|
||||
}
|
||||
ch, _ := newTestChannel(t, cfg)
|
||||
|
||||
if ch.Name() != "websocket" {
|
||||
t.Errorf("Name() = %q, want %q", ch.Name(), "websocket")
|
||||
}
|
||||
if ch.ReasoningChannelID() != "reason-chan" {
|
||||
t.Errorf("ReasoningChannelID() = %q, want %q", ch.ReasoningChannelID(), "reason-chan")
|
||||
}
|
||||
if ch.IsRunning() {
|
||||
t.Error("expected IsRunning() == false before Start")
|
||||
}
|
||||
if ch.config.AgentID != "researcher" {
|
||||
t.Errorf("config.AgentID = %q, want %q", ch.config.AgentID, "researcher")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSend_NotRunning(t *testing.T) {
|
||||
cfg := config.WebSocketConfig{Enabled: true}
|
||||
ch, _ := newTestChannel(t, cfg)
|
||||
|
||||
err := ch.Send(context.Background(), bus.OutboundMessage{
|
||||
Channel: "websocket",
|
||||
ChatID: "chat-1",
|
||||
Content: "hello",
|
||||
})
|
||||
if err != channels.ErrNotRunning {
|
||||
t.Errorf("Send() error = %v, want %v", err, channels.ErrNotRunning)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSend_WritesToConn(t *testing.T) {
|
||||
received := make(chan WSEnvelope, 1)
|
||||
_, wsURL := mockServer(t, func(conn *websocket.Conn) {
|
||||
// skip auth envelope
|
||||
conn.SetReadDeadline(time.Now().Add(2 * time.Second))
|
||||
conn.ReadMessage()
|
||||
// read the actual message
|
||||
env := readEnvelope(t, conn, 2*time.Second)
|
||||
received <- env
|
||||
conn.Close()
|
||||
})
|
||||
|
||||
cfg := config.WebSocketConfig{
|
||||
Enabled: true,
|
||||
WSUrl: wsURL,
|
||||
AccessToken: "tok",
|
||||
AgentID: "a1",
|
||||
}
|
||||
ch, _ := newTestChannel(t, cfg)
|
||||
ctx := context.Background()
|
||||
|
||||
if err := ch.Start(ctx); err != nil {
|
||||
t.Fatalf("Start: %v", err)
|
||||
}
|
||||
defer ch.Stop(ctx)
|
||||
|
||||
err := ch.Send(ctx, bus.OutboundMessage{
|
||||
Channel: "websocket",
|
||||
ChatID: "task-1",
|
||||
Content: `{"summary": "done"}`,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Send: %v", err)
|
||||
}
|
||||
|
||||
select {
|
||||
case env := <-received:
|
||||
if env.Type != "message" {
|
||||
t.Errorf("envelope type = %q, want %q", env.Type, "message")
|
||||
}
|
||||
var out WSOutboundMessage
|
||||
if err := json.Unmarshal(env.Data, &out); err != nil {
|
||||
t.Fatalf("unmarshal outbound: %v", err)
|
||||
}
|
||||
if out.ContentType != "json" {
|
||||
t.Errorf("content_type = %q, want %q", out.ContentType, "json")
|
||||
}
|
||||
if out.ChatID != "task-1" {
|
||||
t.Errorf("chat_id = %q, want %q", out.ChatID, "task-1")
|
||||
}
|
||||
case <-time.After(3 * time.Second):
|
||||
t.Fatal("timed out waiting for message on mock server")
|
||||
}
|
||||
}
|
||||
|
||||
func TestContentType_AutoDetect(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
content string
|
||||
want string
|
||||
}{
|
||||
{"plain text", "hello world", "text"},
|
||||
{"valid JSON object", `{"key": "value"}`, "json"},
|
||||
{"valid JSON array", `[1, 2, 3]`, "json"},
|
||||
{"empty string", "", "text"},
|
||||
{"invalid JSON", `{"broken":`, "text"},
|
||||
{"JSON string literal", `"just a string"`, "json"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
got := DetectContentType(tt.content)
|
||||
if got != tt.want {
|
||||
t.Errorf("DetectContentType(%q) = %q, want %q", tt.content, got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestSessionKey_Passthrough(t *testing.T) {
|
||||
_, wsURL := mockServer(t, func(conn *websocket.Conn) {
|
||||
// skip auth
|
||||
conn.SetReadDeadline(time.Now().Add(2 * time.Second))
|
||||
conn.ReadMessage()
|
||||
|
||||
msg := WSInboundMessage{
|
||||
SenderID: "orchestrator",
|
||||
ChatID: "task-1",
|
||||
Content: "analyze this",
|
||||
SessionKey: "task-042",
|
||||
MessageID: "msg-1",
|
||||
}
|
||||
data, _ := json.Marshal(msg)
|
||||
env := WSEnvelope{Type: "message", Data: json.RawMessage(data)}
|
||||
envData, _ := json.Marshal(env)
|
||||
conn.WriteMessage(websocket.TextMessage, envData)
|
||||
// keep connection open briefly
|
||||
time.Sleep(500 * time.Millisecond)
|
||||
conn.Close()
|
||||
})
|
||||
|
||||
cfg := config.WebSocketConfig{
|
||||
Enabled: true,
|
||||
WSUrl: wsURL,
|
||||
AccessToken: "tok",
|
||||
AgentID: "a1",
|
||||
}
|
||||
ch, mb := newTestChannel(t, cfg)
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer cancel()
|
||||
|
||||
if err := ch.Start(ctx); err != nil {
|
||||
t.Fatalf("Start: %v", err)
|
||||
}
|
||||
defer ch.Stop(ctx)
|
||||
|
||||
msg, ok := mb.ConsumeInbound(ctx)
|
||||
if !ok {
|
||||
t.Fatal("ConsumeInbound returned false")
|
||||
}
|
||||
if msg.SessionKey != "task-042" {
|
||||
t.Errorf("SessionKey = %q, want %q", msg.SessionKey, "task-042")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSessionKey_Fallback(t *testing.T) {
|
||||
_, wsURL := mockServer(t, func(conn *websocket.Conn) {
|
||||
conn.SetReadDeadline(time.Now().Add(2 * time.Second))
|
||||
conn.ReadMessage() // skip auth
|
||||
|
||||
msg := WSInboundMessage{
|
||||
SenderID: "orchestrator",
|
||||
ChatID: "task-2",
|
||||
Content: "do something",
|
||||
// no session_key
|
||||
}
|
||||
data, _ := json.Marshal(msg)
|
||||
env := WSEnvelope{Type: "message", Data: json.RawMessage(data)}
|
||||
envData, _ := json.Marshal(env)
|
||||
conn.WriteMessage(websocket.TextMessage, envData)
|
||||
time.Sleep(500 * time.Millisecond)
|
||||
conn.Close()
|
||||
})
|
||||
|
||||
cfg := config.WebSocketConfig{
|
||||
Enabled: true,
|
||||
WSUrl: wsURL,
|
||||
AccessToken: "tok",
|
||||
AgentID: "a1",
|
||||
}
|
||||
ch, mb := newTestChannel(t, cfg)
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer cancel()
|
||||
|
||||
if err := ch.Start(ctx); err != nil {
|
||||
t.Fatalf("Start: %v", err)
|
||||
}
|
||||
defer ch.Stop(ctx)
|
||||
|
||||
msg, ok := mb.ConsumeInbound(ctx)
|
||||
if !ok {
|
||||
t.Fatal("ConsumeInbound returned false")
|
||||
}
|
||||
if msg.SessionKey != "" {
|
||||
t.Errorf("SessionKey = %q, want empty string", msg.SessionKey)
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseEnvelope_InvalidJSON(t *testing.T) {
|
||||
_, err := parseEnvelope([]byte(`not json`))
|
||||
if err == nil {
|
||||
t.Error("expected error for invalid JSON, got nil")
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseEnvelope_UnknownType(t *testing.T) {
|
||||
env, err := parseEnvelope([]byte(`{"type":"unknown","data":{}}`))
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if env.Type != "unknown" {
|
||||
t.Errorf("Type = %q, want %q", env.Type, "unknown")
|
||||
}
|
||||
// Unknown types should be parseable — the channel is responsible for ignoring them.
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Integration Tests (mock WebSocket server via httptest)
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestConnectAndReceive(t *testing.T) {
|
||||
_, wsURL := mockServer(t, func(conn *websocket.Conn) {
|
||||
conn.SetReadDeadline(time.Now().Add(2 * time.Second))
|
||||
conn.ReadMessage() // skip auth
|
||||
|
||||
msg := makeMessageEnvelope(t, WSInboundMessage{
|
||||
SenderID: "orchestrator",
|
||||
ChatID: "task-1",
|
||||
Content: "Analyze this document",
|
||||
MessageID: "msg-100",
|
||||
})
|
||||
writeEnvelope(t, conn, msg)
|
||||
time.Sleep(500 * time.Millisecond)
|
||||
conn.Close()
|
||||
})
|
||||
|
||||
cfg := config.WebSocketConfig{
|
||||
Enabled: true,
|
||||
WSUrl: wsURL,
|
||||
AccessToken: "tok",
|
||||
AgentID: "researcher",
|
||||
}
|
||||
ch, mb := newTestChannel(t, cfg)
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer cancel()
|
||||
|
||||
if err := ch.Start(ctx); err != nil {
|
||||
t.Fatalf("Start: %v", err)
|
||||
}
|
||||
defer ch.Stop(ctx)
|
||||
|
||||
inbound, ok := mb.ConsumeInbound(ctx)
|
||||
if !ok {
|
||||
t.Fatal("ConsumeInbound returned false")
|
||||
}
|
||||
if inbound.Channel != "websocket" {
|
||||
t.Errorf("Channel = %q, want %q", inbound.Channel, "websocket")
|
||||
}
|
||||
if inbound.SenderID != "orchestrator" {
|
||||
t.Errorf("SenderID = %q, want %q", inbound.SenderID, "orchestrator")
|
||||
}
|
||||
if inbound.ChatID != "task-1" {
|
||||
t.Errorf("ChatID = %q, want %q", inbound.ChatID, "task-1")
|
||||
}
|
||||
if inbound.Content != "Analyze this document" {
|
||||
t.Errorf("Content = %q, want %q", inbound.Content, "Analyze this document")
|
||||
}
|
||||
if inbound.MessageID != "msg-100" {
|
||||
t.Errorf("MessageID = %q, want %q", inbound.MessageID, "msg-100")
|
||||
}
|
||||
}
|
||||
|
||||
func TestConnectAndSend(t *testing.T) {
|
||||
received := make(chan WSEnvelope, 1)
|
||||
_, wsURL := mockServer(t, func(conn *websocket.Conn) {
|
||||
conn.SetReadDeadline(time.Now().Add(2 * time.Second))
|
||||
conn.ReadMessage() // skip auth
|
||||
env := readEnvelope(t, conn, 2*time.Second)
|
||||
received <- env
|
||||
conn.Close()
|
||||
})
|
||||
|
||||
cfg := config.WebSocketConfig{
|
||||
Enabled: true,
|
||||
WSUrl: wsURL,
|
||||
AccessToken: "tok",
|
||||
AgentID: "writer",
|
||||
}
|
||||
ch, _ := newTestChannel(t, cfg)
|
||||
ctx := context.Background()
|
||||
|
||||
if err := ch.Start(ctx); err != nil {
|
||||
t.Fatalf("Start: %v", err)
|
||||
}
|
||||
defer ch.Stop(ctx)
|
||||
|
||||
err := ch.Send(ctx, bus.OutboundMessage{
|
||||
Channel: "websocket",
|
||||
ChatID: "task-1",
|
||||
Content: "plain text response",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Send: %v", err)
|
||||
}
|
||||
|
||||
select {
|
||||
case env := <-received:
|
||||
if env.Type != "message" {
|
||||
t.Errorf("type = %q, want %q", env.Type, "message")
|
||||
}
|
||||
var out WSOutboundMessage
|
||||
if err := json.Unmarshal(env.Data, &out); err != nil {
|
||||
t.Fatalf("unmarshal: %v", err)
|
||||
}
|
||||
if out.ContentType != "text" {
|
||||
t.Errorf("content_type = %q, want %q", out.ContentType, "text")
|
||||
}
|
||||
if out.Content != "plain text response" {
|
||||
t.Errorf("content = %q, want %q", out.Content, "plain text response")
|
||||
}
|
||||
case <-time.After(3 * time.Second):
|
||||
t.Fatal("timed out waiting for message on mock server")
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuth_Handshake(t *testing.T) {
|
||||
authReceived := make(chan WSAuthData, 1)
|
||||
_, wsURL := mockServer(t, func(conn *websocket.Conn) {
|
||||
env := readEnvelope(t, conn, 2*time.Second)
|
||||
if env.Type != "auth" {
|
||||
t.Errorf("first envelope type = %q, want %q", env.Type, "auth")
|
||||
}
|
||||
var auth WSAuthData
|
||||
if err := json.Unmarshal(env.Data, &auth); err != nil {
|
||||
t.Fatalf("unmarshal auth: %v", err)
|
||||
}
|
||||
authReceived <- auth
|
||||
time.Sleep(500 * time.Millisecond)
|
||||
conn.Close()
|
||||
})
|
||||
|
||||
cfg := config.WebSocketConfig{
|
||||
Enabled: true,
|
||||
WSUrl: wsURL,
|
||||
AccessToken: "secret-token",
|
||||
AgentID: "researcher",
|
||||
}
|
||||
ch, _ := newTestChannel(t, cfg)
|
||||
ctx := context.Background()
|
||||
|
||||
if err := ch.Start(ctx); err != nil {
|
||||
t.Fatalf("Start: %v", err)
|
||||
}
|
||||
defer ch.Stop(ctx)
|
||||
|
||||
select {
|
||||
case auth := <-authReceived:
|
||||
if auth.AgentID != "researcher" {
|
||||
t.Errorf("agent_id = %q, want %q", auth.AgentID, "researcher")
|
||||
}
|
||||
if auth.Token != "secret-token" {
|
||||
t.Errorf("token = %q, want %q", auth.Token, "secret-token")
|
||||
}
|
||||
case <-time.After(3 * time.Second):
|
||||
t.Fatal("timed out waiting for auth handshake")
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuth_Rejected(t *testing.T) {
|
||||
_, wsURL := mockServer(t, func(conn *websocket.Conn) {
|
||||
// Read auth and immediately close with an error
|
||||
conn.SetReadDeadline(time.Now().Add(2 * time.Second))
|
||||
conn.ReadMessage()
|
||||
conn.WriteMessage(websocket.CloseMessage,
|
||||
websocket.FormatCloseMessage(websocket.ClosePolicyViolation, "bad token"))
|
||||
conn.Close()
|
||||
})
|
||||
|
||||
cfg := config.WebSocketConfig{
|
||||
Enabled: true,
|
||||
WSUrl: wsURL,
|
||||
AccessToken: "wrong-token",
|
||||
AgentID: "a1",
|
||||
ReconnectInterval: 0, // no reconnect
|
||||
}
|
||||
ch, _ := newTestChannel(t, cfg)
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
|
||||
defer cancel()
|
||||
|
||||
// Start should either return an error or handle gracefully
|
||||
// (the channel should not panic or hang)
|
||||
_ = ch.Start(ctx)
|
||||
defer ch.Stop(ctx)
|
||||
|
||||
// Give it a moment to process the close
|
||||
time.Sleep(200 * time.Millisecond)
|
||||
// Channel should not be running after auth rejection (or should have handled it)
|
||||
}
|
||||
|
||||
func TestReconnect(t *testing.T) {
|
||||
connCount := make(chan int, 5)
|
||||
count := 0
|
||||
var countMu sync.Mutex
|
||||
|
||||
_, wsURL := mockServer(t, func(conn *websocket.Conn) {
|
||||
countMu.Lock()
|
||||
count++
|
||||
current := count
|
||||
countMu.Unlock()
|
||||
connCount <- current
|
||||
|
||||
// Read auth
|
||||
conn.SetReadDeadline(time.Now().Add(2 * time.Second))
|
||||
conn.ReadMessage()
|
||||
|
||||
if current == 1 {
|
||||
// First connection: close immediately to trigger reconnect
|
||||
conn.Close()
|
||||
return
|
||||
}
|
||||
// Second connection: stay open a bit
|
||||
time.Sleep(2 * time.Second)
|
||||
conn.Close()
|
||||
})
|
||||
|
||||
cfg := config.WebSocketConfig{
|
||||
Enabled: true,
|
||||
WSUrl: wsURL,
|
||||
AccessToken: "tok",
|
||||
AgentID: "a1",
|
||||
ReconnectInterval: 1, // 1 second
|
||||
}
|
||||
ch, _ := newTestChannel(t, cfg)
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||
defer cancel()
|
||||
|
||||
if err := ch.Start(ctx); err != nil {
|
||||
t.Fatalf("Start: %v", err)
|
||||
}
|
||||
defer ch.Stop(ctx)
|
||||
|
||||
// Wait for at least 2 connections (initial + reconnect)
|
||||
seen := 0
|
||||
for seen < 2 {
|
||||
select {
|
||||
case <-connCount:
|
||||
seen++
|
||||
case <-ctx.Done():
|
||||
t.Fatalf("timed out waiting for reconnect, saw %d connections", seen)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestStop_CleanShutdown(t *testing.T) {
|
||||
_, wsURL := mockServer(t, func(conn *websocket.Conn) {
|
||||
conn.SetReadDeadline(time.Now().Add(5 * time.Second))
|
||||
conn.ReadMessage() // auth
|
||||
// Hold connection open
|
||||
time.Sleep(5 * time.Second)
|
||||
conn.Close()
|
||||
})
|
||||
|
||||
cfg := config.WebSocketConfig{
|
||||
Enabled: true,
|
||||
WSUrl: wsURL,
|
||||
AccessToken: "tok",
|
||||
AgentID: "a1",
|
||||
}
|
||||
ch, _ := newTestChannel(t, cfg)
|
||||
ctx := context.Background()
|
||||
|
||||
if err := ch.Start(ctx); err != nil {
|
||||
t.Fatalf("Start: %v", err)
|
||||
}
|
||||
|
||||
// Give it a moment to connect
|
||||
time.Sleep(200 * time.Millisecond)
|
||||
|
||||
if err := ch.Stop(ctx); err != nil {
|
||||
t.Fatalf("Stop: %v", err)
|
||||
}
|
||||
if ch.IsRunning() {
|
||||
t.Error("expected IsRunning() == false after Stop")
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Concurrency Tests
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestConcurrentSends(t *testing.T) {
|
||||
received := make(chan struct{}, 100)
|
||||
_, wsURL := mockServer(t, func(conn *websocket.Conn) {
|
||||
conn.SetReadDeadline(time.Now().Add(5 * time.Second))
|
||||
conn.ReadMessage() // auth
|
||||
|
||||
// Read all messages until error
|
||||
for {
|
||||
conn.SetReadDeadline(time.Now().Add(3 * time.Second))
|
||||
_, _, err := conn.ReadMessage()
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
received <- struct{}{}
|
||||
}
|
||||
})
|
||||
|
||||
cfg := config.WebSocketConfig{
|
||||
Enabled: true,
|
||||
WSUrl: wsURL,
|
||||
AccessToken: "tok",
|
||||
AgentID: "a1",
|
||||
}
|
||||
ch, _ := newTestChannel(t, cfg)
|
||||
ctx := context.Background()
|
||||
|
||||
if err := ch.Start(ctx); err != nil {
|
||||
t.Fatalf("Start: %v", err)
|
||||
}
|
||||
defer ch.Stop(ctx)
|
||||
|
||||
// Give it a moment to connect
|
||||
time.Sleep(200 * time.Millisecond)
|
||||
|
||||
const numGoroutines = 10
|
||||
var wg sync.WaitGroup
|
||||
wg.Add(numGoroutines)
|
||||
errs := make(chan error, numGoroutines)
|
||||
|
||||
for i := 0; i < numGoroutines; i++ {
|
||||
go func(i int) {
|
||||
defer wg.Done()
|
||||
err := ch.Send(ctx, bus.OutboundMessage{
|
||||
Channel: "websocket",
|
||||
ChatID: "task-1",
|
||||
Content: "concurrent message",
|
||||
})
|
||||
if err != nil {
|
||||
errs <- err
|
||||
}
|
||||
}(i)
|
||||
}
|
||||
wg.Wait()
|
||||
close(errs)
|
||||
|
||||
for err := range errs {
|
||||
t.Errorf("concurrent Send error: %v", err)
|
||||
}
|
||||
|
||||
// Verify server received messages (allow some time for delivery)
|
||||
timer := time.NewTimer(3 * time.Second)
|
||||
defer timer.Stop()
|
||||
count := 0
|
||||
for count < numGoroutines {
|
||||
select {
|
||||
case <-received:
|
||||
count++
|
||||
case <-timer.C:
|
||||
t.Errorf("received %d/%d messages", count, numGoroutines)
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -206,6 +206,7 @@ type ChannelsConfig struct {
|
|||
WeComApp WeComAppConfig `json:"wecom_app"`
|
||||
WeComAIBot WeComAIBotConfig `json:"wecom_aibot"`
|
||||
Pico PicoConfig `json:"pico"`
|
||||
WebSocket WebSocketConfig `json:"websocket"`
|
||||
}
|
||||
|
||||
// GroupTriggerConfig controls when the bot responds in group chats.
|
||||
|
|
@ -386,6 +387,16 @@ type PicoConfig struct {
|
|||
Placeholder PlaceholderConfig `json:"placeholder,omitempty"`
|
||||
}
|
||||
|
||||
type WebSocketConfig struct {
|
||||
Enabled bool `json:"enabled" env:"PICOCLAW_CHANNELS_WEBSOCKET_ENABLED"`
|
||||
WSUrl string `json:"ws_url" env:"PICOCLAW_CHANNELS_WEBSOCKET_WS_URL"`
|
||||
AccessToken string `json:"access_token" env:"PICOCLAW_CHANNELS_WEBSOCKET_ACCESS_TOKEN"`
|
||||
AgentID string `json:"agent_id" env:"PICOCLAW_CHANNELS_WEBSOCKET_AGENT_ID"`
|
||||
ReconnectInterval int `json:"reconnect_interval" env:"PICOCLAW_CHANNELS_WEBSOCKET_RECONNECT_INTERVAL"`
|
||||
AllowFrom FlexibleStringSlice `json:"allow_from" env:"PICOCLAW_CHANNELS_WEBSOCKET_ALLOW_FROM"`
|
||||
ReasoningChannelID string `json:"reasoning_channel_id" env:"PICOCLAW_CHANNELS_WEBSOCKET_REASONING_CHANNEL_ID"`
|
||||
}
|
||||
|
||||
type HeartbeatConfig struct {
|
||||
Enabled bool `json:"enabled" env:"PICOCLAW_HEARTBEAT_ENABLED"`
|
||||
Interval int `json:"interval" env:"PICOCLAW_HEARTBEAT_INTERVAL"` // minutes, min 5
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue