nomadcode/services/core/internal/scheduler/jobs_test.go

489 lines
13 KiB
Go

package scheduler
import (
"context"
"encoding/json"
"errors"
"log/slog"
"strings"
"testing"
"time"
"github.com/riverqueue/river"
"github.com/nomadcode/nomadcode-core/internal/agent"
"github.com/nomadcode/nomadcode-core/internal/model"
"github.com/nomadcode/nomadcode-core/internal/notification"
"github.com/nomadcode/nomadcode-core/internal/storage"
"github.com/nomadcode/nomadcode-core/internal/workflow"
)
func TestRunAgentTaskCompletesFromMessage(t *testing.T) {
worker := &TaskWorker{
Agent: fakeAgentClient{
result: agent.SendMessageResult{
Message: &agent.Message{
Kind: "message",
Role: "agent",
MessageID: "msg-1",
Parts: []agent.Part{{Kind: "text", Text: "done"}},
},
Raw: json.RawMessage(`{"kind":"message","messageId":"msg-1"}`),
},
},
}
raw, message, err := worker.runAgentTask(context.Background(), storage.Task{
ID: "task-1",
Title: "Fix issue",
Source: "manual",
Payload: json.RawMessage(`{"prompt":"please fix it"}`),
})
if err != nil {
t.Fatalf("runAgentTask returned error: %v", err)
}
if message != "done" {
t.Fatalf("unexpected message: %q", message)
}
var output map[string]any
if err := json.Unmarshal(raw, &output); err != nil {
t.Fatalf("decode output: %v", err)
}
if output["mode"] != "a2a" || output["message"] != "done" || output["message_id"] != "msg-1" {
t.Fatalf("unexpected output: %#v", output)
}
}
func TestRunAgentTaskRejectsNonTerminalTask(t *testing.T) {
worker := &TaskWorker{
Agent: fakeAgentClient{
result: agent.SendMessageResult{
Task: &agent.Task{
Kind: "task",
ID: "remote-1",
Status: agent.TaskStatus{State: agent.TaskStateWorking},
},
},
},
}
_, _, err := worker.runAgentTask(context.Background(), storage.Task{
ID: "task-1",
Title: "Fix issue",
Source: "manual",
Payload: json.RawMessage(`{}`),
})
if err == nil || !strings.Contains(err.Error(), "non-terminal") {
t.Fatalf("expected non-terminal state error, got %v", err)
}
}
func TestBuildAgentInputUsesPromptPayload(t *testing.T) {
input := buildAgentInput(storage.Task{
ID: "task-1",
Title: "Fallback title",
Source: "manual",
Payload: json.RawMessage(`{"prompt":"do the thing","instructions":"be concise"}`),
})
if !input.Blocking {
t.Fatal("expected blocking A2A call")
}
if !strings.Contains(input.Text, "be concise") || !strings.Contains(input.Text, "do the thing") {
t.Fatalf("unexpected text: %q", input.Text)
}
if input.Metadata["nomadcode_task_id"] != "task-1" {
t.Fatalf("unexpected metadata: %#v", input.Metadata)
}
}
type fakeAgentClient struct {
result agent.SendMessageResult
err error
}
func (c fakeAgentClient) SendMessage(context.Context, agent.SendMessageInput) (agent.SendMessageResult, error) {
return c.result, c.err
}
func (c fakeAgentClient) GetTask(context.Context, agent.GetTaskInput) (agent.Task, error) {
return agent.Task{}, nil
}
func (c fakeAgentClient) CancelTask(context.Context, agent.CancelTaskInput) (agent.Task, error) {
return agent.Task{}, nil
}
type fakeTaskLifecycle struct {
started []string
completed []string
failed []string
failMessages []string
failInputs []workflow.FailureInput
task storage.Task
}
func (f *fakeTaskLifecycle) StartTask(ctx context.Context, id string) (storage.Task, error) {
if f.task.Status != "pending" && f.task.Status != "queued" && f.task.Status != "failed" && f.task.Status != "" {
return storage.Task{}, errors.New("invalid transition to running")
}
f.started = append(f.started, id)
f.task.Status = "running"
return f.task, nil
}
func (f *fakeTaskLifecycle) CompleteTask(ctx context.Context, id string, result json.RawMessage) (storage.Task, error) {
if f.task.Status != "running" {
return storage.Task{}, errors.New("invalid transition to completed")
}
f.completed = append(f.completed, id)
f.task.Status = "completed"
return f.task, nil
}
func (f *fakeTaskLifecycle) FailTaskWithMetadata(ctx context.Context, id string, input workflow.FailureInput) (storage.Task, error) {
if f.task.Status != "running" {
return storage.Task{}, errors.New("invalid transition to failed")
}
f.failed = append(f.failed, id)
f.failMessages = append(f.failMessages, input.Message)
f.failInputs = append(f.failInputs, input)
f.task.Status = "failed"
return f.task, nil
}
func (f *fakeTaskLifecycle) FailTask(ctx context.Context, id string, message string) (storage.Task, error) {
return f.FailTaskWithMetadata(ctx, id, workflow.FailureInput{
Message: message,
Type: workflow.FailureTypeExecution,
})
}
func TestWorkCompletesThroughLifecycle(t *testing.T) {
fakeLifecycle := &fakeTaskLifecycle{
task: storage.Task{
ID: "task-1",
Title: "Test task",
Source: "test",
Payload: json.RawMessage(`{}`),
},
}
worker := &TaskWorker{
Lifecycle: fakeLifecycle,
}
job := &river.Job[TaskJobArgs]{
Args: TaskJobArgs{TaskID: "task-1"},
}
err := worker.Work(context.Background(), job)
if err != nil {
t.Fatalf("Work returned error: %v", err)
}
if len(fakeLifecycle.started) != 1 || fakeLifecycle.started[0] != "task-1" {
t.Fatalf("expected started to have task-1, got %v", fakeLifecycle.started)
}
if len(fakeLifecycle.completed) != 1 || fakeLifecycle.completed[0] != "task-1" {
t.Fatalf("expected completed to have task-1, got %v", fakeLifecycle.completed)
}
if len(fakeLifecycle.failed) != 0 {
t.Fatalf("expected no failed tasks, got %v", fakeLifecycle.failed)
}
}
func TestWorkFailsThroughLifecycle(t *testing.T) {
fakeLifecycle := &fakeTaskLifecycle{
task: storage.Task{
ID: "task-1",
Title: "Test task",
Source: "test",
Payload: json.RawMessage(`{}`),
},
}
worker := &TaskWorker{
Lifecycle: fakeLifecycle,
Agent: fakeAgentClient{
err: errors.New("agent failed"),
},
}
job := &river.Job[TaskJobArgs]{
Args: TaskJobArgs{TaskID: "task-1"},
}
err := worker.Work(context.Background(), job)
if err == nil {
t.Fatal("expected Work to return error")
}
if len(fakeLifecycle.started) != 1 || fakeLifecycle.started[0] != "task-1" {
t.Fatalf("expected started to have task-1, got %v", fakeLifecycle.started)
}
if len(fakeLifecycle.completed) != 0 {
t.Fatalf("expected no completed tasks, got %v", fakeLifecycle.completed)
}
if len(fakeLifecycle.failed) != 1 || fakeLifecycle.failed[0] != "task-1" {
t.Fatalf("expected failed to have task-1, got %v", fakeLifecycle.failed)
}
}
type blockingModelClient struct{}
func (m blockingModelClient) Generate(ctx context.Context, input model.GenerateInput) (model.GenerateResult, error) {
<-ctx.Done()
return model.GenerateResult{}, ctx.Err()
}
func TestWorkMarksTimeoutFailure(t *testing.T) {
fakeLifecycle := &fakeTaskLifecycle{
task: storage.Task{
ID: "task-1",
Title: "Test task",
Source: "test",
Payload: json.RawMessage(`{}`),
},
}
worker := &TaskWorker{
Lifecycle: fakeLifecycle,
Model: blockingModelClient{},
RunTimeout: 1 * time.Millisecond,
}
job := &river.Job[TaskJobArgs]{
Args: TaskJobArgs{TaskID: "task-1"},
}
err := worker.Work(context.Background(), job)
if err == nil {
t.Fatal("expected Work to return timeout error")
}
if len(fakeLifecycle.started) != 1 || fakeLifecycle.started[0] != "task-1" {
t.Fatalf("expected started to have task-1, got %v", fakeLifecycle.started)
}
if len(fakeLifecycle.completed) != 0 {
t.Fatalf("expected no completed tasks, got %v", fakeLifecycle.completed)
}
if len(fakeLifecycle.failed) != 1 || fakeLifecycle.failed[0] != "task-1" {
t.Fatalf("expected failed to have task-1, got %v", fakeLifecycle.failed)
}
if len(fakeLifecycle.failMessages) != 1 || !strings.Contains(fakeLifecycle.failMessages[0], "timeout") {
t.Fatalf("expected fail message containing timeout, got %v", fakeLifecycle.failMessages)
}
if len(fakeLifecycle.failInputs) != 1 {
t.Fatalf("expected 1 fail input, got %d", len(fakeLifecycle.failInputs))
}
if fakeLifecycle.failInputs[0].Type != workflow.FailureTypeTimeout {
t.Errorf("expected failure type %q, got %q", workflow.FailureTypeTimeout, fakeLifecycle.failInputs[0].Type)
}
}
func TestWorkRetriesAfterFailureState(t *testing.T) {
fakeLifecycle := &fakeTaskLifecycle{
task: storage.Task{
ID: "task-1",
Status: "pending",
Title: "Test task",
Source: "test",
Payload: json.RawMessage(`{}`),
},
}
worker := &TaskWorker{
Lifecycle: fakeLifecycle,
Agent: fakeAgentClient{
err: errors.New("first attempt agent failure"),
},
}
job := &river.Job[TaskJobArgs]{
Args: TaskJobArgs{TaskID: "task-1"},
}
err := worker.Work(context.Background(), job)
if err == nil {
t.Fatal("expected first Work to return error")
}
if fakeLifecycle.task.Status != "failed" {
t.Fatalf("expected task status to be failed, got: %s", fakeLifecycle.task.Status)
}
if len(fakeLifecycle.failInputs) != 1 {
t.Fatalf("expected 1 fail input, got %d", len(fakeLifecycle.failInputs))
}
if fakeLifecycle.failInputs[0].Type != workflow.FailureTypeExecution {
t.Errorf("expected failure type %q, got %q", workflow.FailureTypeExecution, fakeLifecycle.failInputs[0].Type)
}
worker.Agent = nil
err = worker.Work(context.Background(), job)
if err != nil {
t.Fatalf("expected second Work to succeed, got error: %v", err)
}
if fakeLifecycle.task.Status != "completed" {
t.Fatalf("expected task status to be completed, got: %s", fakeLifecycle.task.Status)
}
if len(fakeLifecycle.started) != 2 {
t.Fatalf("expected StartTask to be called 2 times, got: %d", len(fakeLifecycle.started))
}
if len(fakeLifecycle.failed) != 1 {
t.Fatalf("expected FailTask to be called 1 time, got: %d", len(fakeLifecycle.failed))
}
if len(fakeLifecycle.completed) != 1 {
t.Fatalf("expected CompleteTask to be called 1 time, got: %d", len(fakeLifecycle.completed))
}
}
type logRecord struct {
msg string
args map[string]any
}
type spyHandler struct {
records []logRecord
}
func (h *spyHandler) Enabled(ctx context.Context, level slog.Level) bool {
return true
}
func (h *spyHandler) Handle(ctx context.Context, r slog.Record) error {
args := make(map[string]any)
r.Attrs(func(a slog.Attr) bool {
args[a.Key] = a.Value.Any()
return true
})
h.records = append(h.records, logRecord{
msg: r.Message,
args: args,
})
return nil
}
func (h *spyHandler) WithAttrs(attrs []slog.Attr) slog.Handler {
return h
}
func (h *spyHandler) WithGroup(name string) slog.Handler {
return h
}
func TestWorkEmitsRunningAndCompletedEvents(t *testing.T) {
spy := &spyHandler{}
logger := slog.New(spy)
notifService := notification.NewService(nil, logger)
fakeLifecycle := &fakeTaskLifecycle{
task: storage.Task{
ID: "task-1",
Title: "Success Task",
Source: "test",
Status: "pending",
Payload: json.RawMessage(`{}`),
},
}
worker := &TaskWorker{
Lifecycle: fakeLifecycle,
Notifications: notifService,
Logger: logger,
}
job := &river.Job[TaskJobArgs]{
Args: TaskJobArgs{TaskID: "task-1"},
}
err := worker.Work(context.Background(), job)
if err != nil {
t.Fatalf("Work returned error: %v", err)
}
// Verify events
var eventTypes []notification.TaskEventType
for _, rec := range spy.records {
if rec.msg == "task event notification requested" {
if tp, ok := rec.args["type"].(notification.TaskEventType); ok {
eventTypes = append(eventTypes, tp)
}
}
}
if len(eventTypes) != 2 {
t.Fatalf("expected 2 notification events, got %d: %v", len(eventTypes), eventTypes)
}
if eventTypes[0] != notification.TaskEventRunning {
t.Errorf("expected first event to be %s, got %s", notification.TaskEventRunning, eventTypes[0])
}
if eventTypes[1] != notification.TaskEventCompleted {
t.Errorf("expected second event to be %s, got %s", notification.TaskEventCompleted, eventTypes[1])
}
}
func TestWorkEmitsFailedEvent(t *testing.T) {
spy := &spyHandler{}
logger := slog.New(spy)
notifService := notification.NewService(nil, logger)
fakeLifecycle := &fakeTaskLifecycle{
task: storage.Task{
ID: "task-2",
Title: "Failure Task",
Source: "test",
Status: "pending",
Payload: json.RawMessage(`{}`),
},
}
worker := &TaskWorker{
Lifecycle: fakeLifecycle,
Notifications: notifService,
Agent: fakeAgentClient{
err: errors.New("agent failed"),
},
Logger: logger,
}
job := &river.Job[TaskJobArgs]{
Args: TaskJobArgs{TaskID: "task-2"},
}
err := worker.Work(context.Background(), job)
if err == nil {
t.Fatal("expected Work to return error")
}
// Verify events
var eventTypes []notification.TaskEventType
for _, rec := range spy.records {
if rec.msg == "task event notification requested" {
if tp, ok := rec.args["type"].(notification.TaskEventType); ok {
eventTypes = append(eventTypes, tp)
}
}
}
// Running + Failed
if len(eventTypes) != 2 {
t.Fatalf("expected 2 notification events, got %d: %v", len(eventTypes), eventTypes)
}
if eventTypes[0] != notification.TaskEventRunning {
t.Errorf("expected first event to be %s, got %s", notification.TaskEventRunning, eventTypes[0])
}
if eventTypes[1] != notification.TaskEventFailed {
t.Errorf("expected second event to be %s, got %s", notification.TaskEventFailed, eventTypes[1])
}
}