Skip to content

Commit 5e696b3

Browse files
authored
Add job context (#30)
* add no timescale option, add tests for manager * add parent rid functionality, add tests, add context to queuer * fix formatting
1 parent a478d6b commit 5e696b3

13 files changed

Lines changed: 369 additions & 28 deletions

core/runner.go

Lines changed: 36 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -8,15 +8,21 @@ import (
88
"sync"
99
"time"
1010

11+
"github.com/google/uuid"
1112
"github.com/siherrmann/queuer/helper"
1213
"github.com/siherrmann/queuer/model"
1314
vh "github.com/siherrmann/validator/helper"
1415
)
1516

17+
type ContextKey string
18+
19+
const JobRIDContextKey ContextKey = "JobRIDContextKey"
20+
1621
type Runner struct {
1722
cancel context.CancelFunc
1823
cancelMu sync.RWMutex
1924
Options *model.Options
25+
JobRID *uuid.UUID
2026
Task interface{}
2127
Parameters model.Parameters
2228
// Result channel to return results
@@ -31,14 +37,22 @@ func NewRunner(options *model.Options, task interface{}, parameters ...interface
3137
taskInputParameters, err := helper.GetInputParametersFromTask(task)
3238
if err != nil {
3339
return nil, helper.NewError("getting task input parameters", err)
34-
} else if len(taskInputParameters) != len(parameters) {
35-
return nil, fmt.Errorf("task expects %d parameters, got %d", len(taskInputParameters), len(parameters))
40+
}
41+
42+
startIndex := 0
43+
contextType := reflect.TypeOf((*context.Context)(nil)).Elem()
44+
if len(taskInputParameters) > 0 && taskInputParameters[0] == contextType {
45+
startIndex = 1
46+
}
47+
48+
if len(taskInputParameters)-startIndex != len(parameters) {
49+
return nil, fmt.Errorf("task expects %d parameters, got %d", len(taskInputParameters)-startIndex, len(parameters))
3650
}
3751

3852
for i, param := range parameters {
39-
paramConverted, err := vh.AnyToType(param, taskInputParameters[i])
53+
paramConverted, err := vh.AnyToType(param, taskInputParameters[i+startIndex])
4054
if err != nil {
41-
return nil, fmt.Errorf("error converting parameter %d to type %s: %v", i, taskInputParameters[i].Kind(), err)
55+
return nil, fmt.Errorf("error converting parameter %d to type %s: %v", i, taskInputParameters[i+startIndex].Kind(), err)
4256
}
4357
parameters[i] = paramConverted
4458
}
@@ -73,6 +87,7 @@ func NewRunnerFromJob(task *model.Task, job *model.Job) (*Runner, error) {
7387
return nil, fmt.Errorf("error creating runner from job: %v", err)
7488
}
7589

90+
runner.JobRID = &job.RID
7691
return runner, nil
7792
}
7893

@@ -111,7 +126,21 @@ func (r *Runner) Run(ctx context.Context) {
111126
}()
112127

113128
taskFunc := reflect.ValueOf(r.Task)
114-
results := taskFunc.Call(r.Parameters.ToReflectValues())
129+
130+
var callParameters []reflect.Value
131+
taskType := reflect.TypeOf(r.Task)
132+
contextType := reflect.TypeOf((*context.Context)(nil)).Elem()
133+
134+
if taskType.NumIn() > 0 && taskType.In(0) == contextType {
135+
jobCtx := ctx
136+
if r.JobRID != nil {
137+
jobCtx = context.WithValue(ctx, JobRIDContextKey, *r.JobRID)
138+
}
139+
callParameters = append(callParameters, reflect.ValueOf(jobCtx))
140+
}
141+
142+
callParameters = append(callParameters, r.Parameters.ToReflectValues()...)
143+
results := taskFunc.Call(callParameters)
115144
resultValues := []interface{}{}
116145
for _, result := range results {
117146
resultValues = append(resultValues, result.Interface())
@@ -125,7 +154,8 @@ func (r *Runner) Run(ctx context.Context) {
125154

126155
var ok bool
127156
if len(resultValues) > 0 {
128-
if err, ok = resultValues[len(resultValues)-1].(error); ok || (outputParameters[1].String() == "error" && resultValues[len(resultValues)-1] == nil) {
157+
lastOutIdx := len(outputParameters) - 1
158+
if err, ok = resultValues[len(resultValues)-1].(error); ok || (lastOutIdx >= 0 && outputParameters[lastOutIdx].String() == "error" && resultValues[len(resultValues)-1] == nil) {
129159
resultValues = resultValues[:len(resultValues)-1]
130160
}
131161
}

core/runner_test.go

Lines changed: 40 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -692,3 +692,43 @@ func TestNewRunnerWithNestedStruct(t *testing.T) {
692692
})
693693
}
694694
}
695+
696+
func TestRunnerContextInjection(t *testing.T) {
697+
rid := uuid.New()
698+
taskFn := func(ctx context.Context, p string) string {
699+
ctxRid := ctx.Value(JobRIDContextKey)
700+
if ctxRid != nil {
701+
if id, ok := ctxRid.(uuid.UUID); ok {
702+
return id.String() + "-" + p
703+
}
704+
}
705+
return "no-ctx-" + p
706+
}
707+
708+
task, err := model.NewTask(taskFn)
709+
require.NoError(t, err)
710+
711+
job := &model.Job{
712+
RID: rid,
713+
TaskName: task.Name,
714+
Parameters: []interface{}{"test"},
715+
}
716+
717+
runner, err := NewRunnerFromJob(task, job)
718+
require.NoError(t, err)
719+
720+
// Run the runner with a background context
721+
go runner.Run(context.Background())
722+
723+
select {
724+
case results := <-runner.ResultsChannel:
725+
require.Len(t, results, 1)
726+
resStr, ok := results[0].(string)
727+
require.True(t, ok)
728+
assert.Equal(t, rid.String()+"-test", resStr, "Context should contain the JobRID injected by Runner")
729+
case err := <-runner.ErrorChannel:
730+
t.Fatalf("Runner returned error: %v", err)
731+
case <-time.After(1 * time.Second):
732+
t.Fatal("Runner timed out")
733+
}
734+
}

database/dbJob_test.go

Lines changed: 45 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -819,3 +819,48 @@ func TestJobSelectAllJobsFromArchiveBySearch(t *testing.T) {
819819
assert.NoError(t, err, "Expected SelectAllJobsFromArchiveBySearch to not return an error")
820820
assert.Len(t, paginatedJobsBySearchFromArchive, pageLength, "Expected SelectAllJobsFromArchiveBySearch to return 3 archived jobs")
821821
}
822+
823+
func TestAddRetentionArchive(t *testing.T) {
824+
helper.SetTestDatabaseConfigEnvs(t, dbPort)
825+
dbConfig, err := helper.NewDatabaseConfiguration()
826+
if err != nil {
827+
t.Fatalf("failed to create database configuration: %v", err)
828+
}
829+
database := helper.NewTestDatabase(dbConfig)
830+
jobDbHandler, err := NewJobDBHandler(database, dbConfig)
831+
require.NoError(t, err)
832+
833+
err = jobDbHandler.AddRetentionArchive(24 * time.Hour)
834+
_ = err
835+
}
836+
837+
func TestRemoveRetentionArchive(t *testing.T) {
838+
helper.SetTestDatabaseConfigEnvs(t, dbPort)
839+
dbConfig, err := helper.NewDatabaseConfiguration()
840+
if err != nil {
841+
t.Fatalf("failed to create database configuration: %v", err)
842+
}
843+
database := helper.NewTestDatabase(dbConfig)
844+
jobDbHandler, err := NewJobDBHandler(database, dbConfig)
845+
require.NoError(t, err)
846+
847+
err = jobDbHandler.RemoveRetentionArchive()
848+
_ = err
849+
}
850+
851+
func TestBatchInsertJobs(t *testing.T) {
852+
helper.SetTestDatabaseConfigEnvs(t, dbPort)
853+
dbConfig, err := helper.NewDatabaseConfiguration()
854+
if err != nil {
855+
t.Fatalf("failed to create database configuration: %v", err)
856+
}
857+
database := helper.NewTestDatabase(dbConfig)
858+
jobDbHandler, err := NewJobDBHandler(database, dbConfig)
859+
require.NoError(t, err)
860+
861+
job1, _ := model.NewJob("Task1", nil, nil)
862+
job2, _ := model.NewJob("Task2", nil, nil)
863+
864+
err = jobDbHandler.BatchInsertJobs([]*model.Job{job1, job2})
865+
assert.NoError(t, err)
866+
}

helper/database.go

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -354,6 +354,9 @@ func (d *Database) DropFunctionsFromPublicSchema(functionNames []string) error {
354354
}
355355
signatures = append(signatures, signature)
356356
}
357+
if err := rows.Err(); err != nil {
358+
return NewError("rows iteration", err)
359+
}
357360

358361
// Drop each overloaded function by its full signature
359362
for _, signature := range signatures {

helper/database_test.go

Lines changed: 35 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2,6 +2,7 @@ package helper
22

33
import (
44
"context"
5+
"database/sql"
56
"log"
67
"testing"
78

@@ -302,3 +303,37 @@ func TestDropIndex(t *testing.T) {
302303
err = database.DropIndex("test_drop_index", "name")
303304
assert.NoError(t, err, "expected no error when dropping non-existing index")
304305
}
306+
307+
func TestNewDatabaseWithDB(t *testing.T) {
308+
db := NewDatabaseWithDB("test_with_db", nil, nil)
309+
assert.NotNil(t, db)
310+
assert.Equal(t, "test_with_db", db.Name)
311+
assert.Nil(t, db.Instance)
312+
}
313+
314+
func TestDatabaseClose(t *testing.T) {
315+
dbConn, _ := sql.Open("postgres", "user=pqgotest dbname=pqgotest sslmode=verify-full")
316+
db := NewDatabaseWithDB("test", dbConn, nil)
317+
err := db.Close()
318+
assert.NoError(t, err)
319+
}
320+
321+
func TestDropFunctionsFromPublicSchema(t *testing.T) {
322+
dbConn, _ := sql.Open("postgres", "user=pqgotest dbname=pqgotest sslmode=verify-full")
323+
db := NewDatabaseWithDB("test", dbConn, nil)
324+
// it will probably error with connection refused, but it tests the nil guard if any, or it tests the query logic
325+
err := db.DropFunctionsFromPublicSchema([]string{"non_existent_func"})
326+
assert.Error(t, err) // since it's a dummy connection string, it should error
327+
}
328+
329+
func TestMustStartPostgresContainer(t *testing.T) {
330+
teardown, port, err := MustStartPostgresContainer()
331+
assert.NoError(t, err)
332+
assert.NotEmpty(t, port)
333+
assert.NotNil(t, teardown)
334+
335+
if teardown != nil {
336+
err := teardown(context.Background())
337+
assert.NoError(t, err)
338+
}
339+
}

helper/prettyLog_test.go

Lines changed: 58 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,58 @@
1+
package helper
2+
3+
import (
4+
"bytes"
5+
"context"
6+
"log/slog"
7+
"testing"
8+
"time"
9+
10+
"github.com/stretchr/testify/assert"
11+
)
12+
13+
func TestPrettyHandler(t *testing.T) {
14+
var buf bytes.Buffer
15+
handler := NewPrettyHandler(&buf, PrettyHandlerOptions{})
16+
17+
logger := slog.New(handler)
18+
logger.Info("test message", "key", "value")
19+
20+
output := buf.String()
21+
assert.Contains(t, output, "test message")
22+
assert.Contains(t, output, "key")
23+
assert.Contains(t, output, "value")
24+
assert.Contains(t, output, "INFO")
25+
}
26+
27+
func TestPrettyHandlerLevels(t *testing.T) {
28+
var buf bytes.Buffer
29+
handler := NewPrettyHandler(&buf, PrettyHandlerOptions{
30+
SlogOpts: slog.HandlerOptions{
31+
Level: slog.LevelDebug,
32+
},
33+
})
34+
35+
logger := slog.New(handler)
36+
logger.Debug("debug msg")
37+
logger.Warn("warn msg")
38+
logger.Error("error msg")
39+
40+
output := buf.String()
41+
assert.Contains(t, output, "debug msg")
42+
assert.Contains(t, output, "warn msg")
43+
assert.Contains(t, output, "error msg")
44+
assert.Contains(t, output, "DEBUG")
45+
assert.Contains(t, output, "WARN")
46+
assert.Contains(t, output, "ERROR")
47+
}
48+
49+
func TestPrettyHandlerHandleDirect(t *testing.T) {
50+
var buf bytes.Buffer
51+
handler := NewPrettyHandler(&buf, PrettyHandlerOptions{})
52+
53+
r := slog.NewRecord(time.Now(), slog.LevelInfo, "direct handle", 0)
54+
55+
err := handler.Handle(context.Background(), r)
56+
assert.NoError(t, err)
57+
assert.Contains(t, buf.String(), "direct handle")
58+
}

helper/task.go

Lines changed: 14 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,7 @@
11
package helper
22

33
import (
4+
"context"
45
"fmt"
56
"reflect"
67
"runtime"
@@ -23,21 +24,29 @@ func CheckValidTask(task interface{}) error {
2324
}
2425

2526
// CheckValidTaskWithParameters checks if the provided task and parameters are valid.
26-
// It checks if the task is a valid function and if the parameters match the task's input types.
27+
// It checks if the task is a valid function and if the parameters match the task's input types, ignoring an optional context.Context as first argument.
2728
func CheckValidTaskWithParameters(task interface{}, parameters ...interface{}) error {
2829
err := CheckValidTask(task)
2930
if err != nil {
3031
return err
3132
}
3233

3334
taskType := reflect.TypeOf(task)
34-
if taskType.NumIn() != len(parameters) {
35-
return fmt.Errorf("task expects %d parameters, got %d", taskType.NumIn(), len(parameters))
35+
expectedParams := taskType.NumIn()
36+
startIndex := 0
37+
38+
contextType := reflect.TypeOf((*context.Context)(nil)).Elem()
39+
if expectedParams > 0 && taskType.In(0) == contextType {
40+
startIndex = 1
41+
}
42+
43+
if expectedParams-startIndex != len(parameters) {
44+
return fmt.Errorf("task expects %d parameters, got %d", expectedParams-startIndex, len(parameters))
3645
}
3746

3847
for i, param := range parameters {
39-
if !reflect.TypeOf(param).AssignableTo(taskType.In(i)) {
40-
return fmt.Errorf("parameter %d of task must be of type %s, got %s", i, taskType.In(i).Kind(), reflect.TypeOf(param).Kind())
48+
if !reflect.TypeOf(param).AssignableTo(taskType.In(i + startIndex)) {
49+
return fmt.Errorf("parameter %d of task must be of type %s, got %s", i, taskType.In(i+startIndex).Kind(), reflect.TypeOf(param).Kind())
4150
}
4251
}
4352

helper/task_test.go

Lines changed: 6 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -123,8 +123,8 @@ func TestCheckValidTaskWithParameters(t *testing.T) {
123123
err := CheckValidTaskWithParameters(testFuncOneParam, "hello")
124124
assert.NoError(t, err, "should not return error for matching params")
125125

126-
err = CheckValidTaskWithParameters(testFuncMultiParamsReturn, context.Background(), 10, 3.14)
127-
assert.NoError(t, err, "should not return error for multiple matching params")
126+
err = CheckValidTaskWithParameters(testFuncMultiParamsReturn, 10, 3.14)
127+
assert.NoError(t, err, "should not return error for multiple matching params, ignoring context")
128128
})
129129

130130
t.Run("Not a function", func(t *testing.T) {
@@ -152,15 +152,15 @@ func TestCheckValidTaskWithParameters(t *testing.T) {
152152
})
153153

154154
t.Run("Function with wrong parameter type (multiple params - first wrong)", func(t *testing.T) {
155-
err := CheckValidTaskWithParameters(testFuncMultiParamsReturn, "wrong", 10, 3.14)
155+
err := CheckValidTaskWithParameters(testFuncMultiParamsReturn, "wrong", 3.14)
156156
require.Error(t, err)
157-
assert.Contains(t, err.Error(), "parameter 0 of task must be of type interface, got string", "should fail on first wrong type")
157+
assert.Contains(t, err.Error(), "parameter 0 of task must be of type int, got string", "should fail on first wrong type")
158158
})
159159

160160
t.Run("Function with wrong parameter type (multiple params - second wrong)", func(t *testing.T) {
161-
err := CheckValidTaskWithParameters(testFuncMultiParamsReturn, context.Background(), "wrong", 3.14)
161+
err := CheckValidTaskWithParameters(testFuncMultiParamsReturn, 10, "wrong")
162162
require.Error(t, err)
163-
assert.Contains(t, err.Error(), "parameter 1 of task must be of type int, got string", "should fail on second wrong type")
163+
assert.Contains(t, err.Error(), "parameter 1 of task must be of type float64, got string", "should fail on second wrong type")
164164
})
165165

166166
t.Run("Function with different parameter type (time.Duration vs int)", func(t *testing.T) {

model/options.go

Lines changed: 5 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -5,12 +5,15 @@ import (
55
"encoding/json"
66
"errors"
77

8+
"github.com/google/uuid"
89
"github.com/siherrmann/queuer/helper"
910
)
1011

1112
type Options struct {
12-
OnError *OnError `json:"on_error,omitempty"`
13-
Schedule *Schedule `json:"schedule,omitempty"`
13+
OnError *OnError `json:"on_error,omitempty"`
14+
Schedule *Schedule `json:"schedule,omitempty"`
15+
Metadata map[string]interface{} `json:"metadata,omitempty"`
16+
ParentRID *uuid.UUID `json:"parent_rid,omitempty"`
1417
}
1518

1619
func (c *Options) IsValid() error {

0 commit comments

Comments
 (0)