nomadcode/services/core/internal/protosocket/tasks_test.go
toki 8bdbc75b4b feat(core): provider registry 조립을 정리한다
외부 provider 추가가 handler 시그니처 고정에 묶이지 않도록 provider 목록 기반 조립으로 정리한다.

Mattermost adapter 완료 리뷰 로그와 외부 통합 마일스톤 진행 근거도 함께 보존한다.
2026-06-05 11:52:28 +09:00

486 lines
15 KiB
Go

package protosocket
import (
"context"
"encoding/base64"
"io"
"log/slog"
"net/http"
"net/http/httptest"
"strings"
"testing"
"time"
toki "git.toki-labs.com/toki/proto-socket/go"
"github.com/jackc/pgx/v5"
"google.golang.org/protobuf/types/known/structpb"
"nhooyr.io/websocket"
apphttp "github.com/nomadcode/nomadcode-core/internal/http"
"github.com/nomadcode/nomadcode-core/internal/notification"
"github.com/nomadcode/nomadcode-core/internal/storage"
"github.com/nomadcode/nomadcode-core/internal/workflow"
)
type fakeTaskService struct {
created workflow.CreateTaskInput
createOut storage.Task
createErr error
listLimit int32
listOut []storage.Task
listErr error
getID string
getOut storage.Task
getErr error
enqueueID string
enqueueOut storage.Task
enqueueErr error
}
func (f *fakeTaskService) CreateTask(_ context.Context, in workflow.CreateTaskInput) (storage.Task, error) {
f.created = in
return f.createOut, f.createErr
}
func (f *fakeTaskService) ListTasks(_ context.Context, limit int32) ([]storage.Task, error) {
f.listLimit = limit
return f.listOut, f.listErr
}
func (f *fakeTaskService) GetTask(_ context.Context, id string) (storage.Task, error) {
f.getID = id
return f.getOut, f.getErr
}
func (f *fakeTaskService) EnqueueTask(_ context.Context, id string) (storage.Task, error) {
f.enqueueID = id
return f.enqueueOut, f.enqueueErr
}
func dispatchTo(fake TaskService, req Envelope) Envelope {
d := NewDispatcher()
NewTaskChannels(fake).Register(d)
return d.Dispatch(context.Background(), req)
}
func TestTaskCreateActionCallsWorkflowAndReturnsResponse(t *testing.T) {
fake := &fakeTaskService{
createOut: storage.Task{ID: "task-1", Status: "pending", Title: "My Task", Source: "client"},
}
res := dispatchTo(fake, Envelope{
ID: "r1",
Type: "request",
Channel: "task",
Action: "task.create",
Payload: map[string]any{
"title": "My Task",
"source": "client",
"payload": map[string]any{},
},
})
if fake.created.Title != "My Task" || fake.created.Source != "client" {
t.Fatalf("workflow not called with expected input: %+v", fake.created)
}
if res.Type != "response" {
t.Fatalf("expected response type, got %q (err=%+v)", res.Type, res.Error)
}
if res.CorrelationID != "r1" {
t.Errorf("expected correlation_id r1, got %q", res.CorrelationID)
}
if res.Payload["id"] != "task-1" || res.Payload["status"] != "pending" {
t.Errorf("unexpected response payload: %+v", res.Payload)
}
if _, ok := res.Payload["task"].(map[string]any); !ok {
t.Errorf("expected nested task map in response payload: %+v", res.Payload)
}
}
func TestTaskListActionForwardsLimitToService(t *testing.T) {
fake := &fakeTaskService{
listOut: []storage.Task{{ID: "task-1"}, {ID: "task-2"}},
}
res := dispatchTo(fake, Envelope{
ID: "r2",
Type: "request",
Channel: "task",
Action: "task.list",
Payload: map[string]any{"limit": float64(5)},
})
if fake.listLimit != 5 {
t.Errorf("expected limit 5 forwarded to service, got %d", fake.listLimit)
}
if res.Type != "response" {
t.Fatalf("expected response type, got %q", res.Type)
}
tasks, ok := res.Payload["tasks"].([]any)
if !ok || len(tasks) != 2 {
t.Errorf("expected 2 tasks in payload, got %+v", res.Payload["tasks"])
}
// Omitted limit forwards 0; the workflow service owns default/bounding.
resDefault := dispatchTo(fake, Envelope{ID: "r2b", Action: "task.list", Channel: "task", Payload: map[string]any{}})
if fake.listLimit != 0 {
t.Errorf("expected omitted limit to forward 0, got %d", fake.listLimit)
}
if resDefault.Type != "response" {
t.Errorf("expected response type for default limit, got %q", resDefault.Type)
}
}
func TestTaskGetActionMapsNotFoundError(t *testing.T) {
fake := &fakeTaskService{getErr: pgx.ErrNoRows}
res := dispatchTo(fake, Envelope{
ID: "r3",
Action: "task.get",
Channel: "task",
Payload: map[string]any{"id": "missing"},
})
if fake.getID != "missing" {
t.Errorf("expected get called with 'missing', got %q", fake.getID)
}
if res.Type != "error" {
t.Fatalf("expected error type, got %q", res.Type)
}
if res.Error == nil || res.Error.Code != "task.not_found" {
t.Errorf("expected task.not_found, got %+v", res.Error)
}
if res.Error.Retryable {
t.Error("expected not_found to be non-retryable")
}
}
func TestTaskEnqueueActionMapsConflictError(t *testing.T) {
fake := &fakeTaskService{enqueueErr: workflow.ErrTaskCannotBeEnqueued}
res := dispatchTo(fake, Envelope{
ID: "r4",
Action: "task.enqueue",
Channel: "task",
Payload: map[string]any{"id": "task-9"},
})
if res.Type != "error" {
t.Fatalf("expected error type, got %q", res.Type)
}
if res.Error == nil || res.Error.Code != "task.conflict" {
t.Errorf("expected task.conflict, got %+v", res.Error)
}
}
func TestTaskActionRejectsInvalidPayload(t *testing.T) {
fake := &fakeTaskService{}
res := dispatchTo(fake, Envelope{
ID: "r5",
Action: "task.get",
Channel: "task",
Payload: map[string]any{"id": float64(123)}, // wrong type for string id
})
if res.Type != "error" {
t.Fatalf("expected error type, got %q", res.Type)
}
if res.Error == nil || res.Error.Code != "task.invalid_payload" {
t.Errorf("expected task.invalid_payload, got %+v", res.Error)
}
if fake.getID != "" {
t.Errorf("service should not be called on invalid payload, got %q", fake.getID)
}
}
type captureBroadcaster struct {
env Envelope
err error
}
func (c *captureBroadcaster) BroadcastEnvelope(_ context.Context, env Envelope) error {
c.env = env
return c.err
}
func TestHandleTaskEventBroadcastsStatusChangedEnvelope(t *testing.T) {
cap := &captureBroadcaster{}
sink := NewTaskEventBroadcaster(cap)
now := time.Now().UTC()
err := sink.HandleTaskEvent(context.Background(), notification.TaskEvent{
Type: notification.TaskEventRunning,
TaskID: "task-7",
Title: "Run me",
Status: "running",
Message: "started",
Attempt: 2,
Reason: "retry-on-failure",
OccurredAt: now,
})
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if cap.env.Type != "event" || cap.env.Channel != "task" || cap.env.Action != "task.status.changed" {
t.Fatalf("unexpected envelope: type=%q channel=%q action=%q", cap.env.Type, cap.env.Channel, cap.env.Action)
}
if cap.env.Payload["id"] != "task-7" || cap.env.Payload["status"] != "running" {
t.Errorf("unexpected event payload: %+v", cap.env.Payload)
}
if cap.env.Payload["type"] != string(notification.TaskEventRunning) {
t.Errorf("expected type %q, got %v", notification.TaskEventRunning, cap.env.Payload["type"])
}
if cap.env.Payload["attempt"] != 2 {
t.Errorf("expected attempt 2, got %v", cap.env.Payload["attempt"])
}
if cap.env.Payload["reason"] != "retry-on-failure" {
t.Errorf("expected reason 'retry-on-failure', got %v", cap.env.Payload["reason"])
}
if cap.env.Payload["occurred_at"] != now.Format(time.RFC3339Nano) {
t.Errorf("expected occurred_at %q, got %v", now.Format(time.RFC3339Nano), cap.env.Payload["occurred_at"])
}
// Test omission of empty reason and zero occurred_at
err = sink.HandleTaskEvent(context.Background(), notification.TaskEvent{
Type: notification.TaskEventRunning,
TaskID: "task-8",
Title: "Run me again",
Status: "running",
Message: "started",
Attempt: 1,
})
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if _, ok := cap.env.Payload["reason"]; ok {
t.Error("expected empty reason to be omitted from payload, but it was present")
}
if _, ok := cap.env.Payload["occurred_at"]; ok {
t.Error("expected zero occurred_at to be omitted from payload, but it was present")
}
}
// --- route-level contract tests over the real router + proto-socket server ---
func newContractRouter(t *testing.T, fake TaskService) (*httptest.Server, func()) {
t.Helper()
logger := slog.New(slog.NewTextHandler(io.Discard, nil))
srv := NewServer(Config{HeartbeatIntervalSec: 0, HeartbeatWaitSec: 0}, logger)
NewTaskChannels(fake).Register(srv.Dispatcher())
router := apphttp.NewRouter(
apphttp.NewHandler(nil, nil, logger),
logger,
apphttp.AuthConfig{Username: "nomadcode", Password: "secret"},
srv,
"/proto-socket",
)
ts := httptest.NewServer(router)
return ts, func() {
ts.Close()
_ = srv.Close()
}
}
func wsURL(ts *httptest.Server) string {
return strings.Replace(ts.URL, "http", "ws", 1) + "/proto-socket"
}
func authDialOptions() *websocket.DialOptions {
header := http.Header{}
header.Set("Authorization", "Basic "+base64.StdEncoding.EncodeToString([]byte("nomadcode:secret")))
return &websocket.DialOptions{HTTPHeader: header}
}
func TestRouterProtoSocketTaskListRequiresAuth(t *testing.T) {
fake := &fakeTaskService{listOut: []storage.Task{{ID: "task-1"}}}
ts, cleanup := newContractRouter(t, fake)
defer cleanup()
ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
defer cancel()
// Unauthenticated dial must fail at the route.
if conn, _, err := websocket.Dial(ctx, wsURL(ts), nil); err == nil {
conn.Close(websocket.StatusNormalClosure, "")
t.Fatal("expected unauthenticated websocket dial to fail")
}
// Authenticated dial succeeds and task.list returns a response envelope.
conn, _, err := websocket.Dial(ctx, wsURL(ts), authDialOptions())
if err != nil {
t.Fatalf("authenticated dial failed: %v", err)
}
defer conn.Close(websocket.StatusNormalClosure, "")
client := toki.NewWsClient(conn, 0, 0, ParserMap())
defer client.Close()
reqStruct, err := Envelope{
ProtocolVersion: ProtocolVersion,
ID: "req-list",
Type: "request",
Channel: "task",
Action: "task.list",
Payload: map[string]any{"limit": float64(20)},
}.ToStruct()
if err != nil {
t.Fatalf("failed to build request: %v", err)
}
resStruct, err := toki.SendRequestTyped[*structpb.Struct, *structpb.Struct](&client.Communicator, reqStruct, 2*time.Second)
if err != nil {
t.Fatalf("send request failed: %v", err)
}
res, err := EnvelopeFromStruct(resStruct)
if err != nil {
t.Fatalf("parse response failed: %v", err)
}
if res.Type != "response" || res.CorrelationID != "req-list" {
t.Fatalf("unexpected response: type=%q correlation=%q", res.Type, res.CorrelationID)
}
if _, ok := res.Payload["tasks"]; !ok {
t.Errorf("expected tasks payload, got %+v", res.Payload)
}
requireDiagnosticsMeta(t, res, "task", "task.list", "")
}
func TestProtoSocketTaskActionsMatchRestSemantics(t *testing.T) {
fake := &fakeTaskService{
createOut: storage.Task{ID: "task-c", Status: "pending"},
getOut: storage.Task{ID: "task-g", Status: "running"},
enqueueOut: storage.Task{ID: "task-e", Status: "queued"},
}
ts, cleanup := newContractRouter(t, fake)
defer cleanup()
ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
defer cancel()
conn, _, err := websocket.Dial(ctx, wsURL(ts), authDialOptions())
if err != nil {
t.Fatalf("dial failed: %v", err)
}
defer conn.Close(websocket.StatusNormalClosure, "")
client := toki.NewWsClient(conn, 0, 0, ParserMap())
defer client.Close()
send := func(action string, payload map[string]any) Envelope {
reqStruct, err := Envelope{
ProtocolVersion: ProtocolVersion,
ID: "req-" + action,
Type: "request",
Channel: "task",
Action: action,
Payload: payload,
}.ToStruct()
if err != nil {
t.Fatalf("build request %s: %v", action, err)
}
resStruct, err := toki.SendRequestTyped[*structpb.Struct, *structpb.Struct](&client.Communicator, reqStruct, 2*time.Second)
if err != nil {
t.Fatalf("send %s: %v", action, err)
}
res, err := EnvelopeFromStruct(resStruct)
if err != nil {
t.Fatalf("parse %s: %v", action, err)
}
return res
}
create := send("task.create", map[string]any{"title": "T", "source": "client", "payload": map[string]any{}})
if create.Type != "response" || create.Payload["id"] != "task-c" || create.Payload["status"] != "pending" {
t.Errorf("unexpected create response: %+v (err=%+v)", create.Payload, create.Error)
}
requireDiagnosticsMeta(t, create, "task", "task.create", "")
get := send("task.get", map[string]any{"id": "task-g"})
if get.Type != "response" {
t.Errorf("unexpected get response type: %q", get.Type)
}
if task, ok := get.Payload["task"].(map[string]any); !ok || task["status"] != "running" {
t.Errorf("unexpected get task payload: %+v", get.Payload)
}
requireDiagnosticsMeta(t, get, "task", "task.get", "")
enqueue := send("task.enqueue", map[string]any{"id": "task-e"})
if enqueue.Type != "response" || enqueue.Payload["status"] != "queued" {
t.Errorf("unexpected enqueue response: %+v", enqueue.Payload)
}
requireDiagnosticsMeta(t, enqueue, "task", "task.enqueue", "")
}
func TestProtoSocketTaskErrorEnvelopeCodes(t *testing.T) {
fake := &fakeTaskService{
getErr: pgx.ErrNoRows,
enqueueErr: workflow.ErrTaskCannotBeEnqueued,
createErr: workflow.ErrInvalidTaskInput,
}
ts, cleanup := newContractRouter(t, fake)
defer cleanup()
ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
defer cancel()
conn, _, err := websocket.Dial(ctx, wsURL(ts), authDialOptions())
if err != nil {
t.Fatalf("dial failed: %v", err)
}
defer conn.Close(websocket.StatusNormalClosure, "")
client := toki.NewWsClient(conn, 0, 0, ParserMap())
defer client.Close()
send := func(action string, payload map[string]any) Envelope {
reqStruct, _ := Envelope{ProtocolVersion: ProtocolVersion, ID: "e-" + action, Type: "request", Channel: "task", Action: action, Payload: payload}.ToStruct()
resStruct, err := toki.SendRequestTyped[*structpb.Struct, *structpb.Struct](&client.Communicator, reqStruct, 2*time.Second)
if err != nil {
t.Fatalf("send %s: %v", action, err)
}
res, _ := EnvelopeFromStruct(resStruct)
return res
}
cases := []struct {
action string
payload map[string]any
code string
}{
{"task.get", map[string]any{"id": "x"}, "task.not_found"},
{"task.enqueue", map[string]any{"id": "x"}, "task.conflict"},
{"task.create", map[string]any{"title": "T", "source": "client"}, "task.invalid_input"},
{"task.get", map[string]any{"id": float64(1)}, "task.invalid_payload"},
{"task.unknown", map[string]any{}, "UNSUPPORTED_ACTION"},
}
for _, tc := range cases {
res := send(tc.action, tc.payload)
if res.Type != "error" {
t.Errorf("%s: expected error type, got %q", tc.action, res.Type)
continue
}
if res.Error == nil || res.Error.Code != tc.code {
t.Errorf("%s: expected code %q, got %+v", tc.action, tc.code, res.Error)
}
requireDiagnosticsMeta(t, res, "task", tc.action, tc.code)
}
}
func TestProtoSocketServerUsesConfiguredHeartbeat(t *testing.T) {
logger := slog.New(slog.NewTextHandler(io.Discard, nil))
srv := NewServer(Config{HeartbeatIntervalSec: 1, HeartbeatWaitSec: 1}, logger)
defer srv.Close()
ts := httptest.NewServer(srv)
defer ts.Close()
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
defer cancel()
conn, _, err := websocket.Dial(ctx, strings.Replace(ts.URL, "http", "ws", 1), nil)
if err != nil {
t.Fatalf("dial failed: %v", err)
}
conn.Close(websocket.StatusNormalClosure, "")
}