217 lines
6.5 KiB
Go
217 lines
6.5 KiB
Go
package batch
|
|
|
|
import (
|
|
"context"
|
|
"io"
|
|
"os"
|
|
"path/filepath"
|
|
"strings"
|
|
"sync/atomic"
|
|
"testing"
|
|
"time"
|
|
|
|
"github.com/example/ollama-fair-gateway/internal/config"
|
|
)
|
|
|
|
func batchTestConfig() config.BatchJobsConfig {
|
|
return config.BatchJobsConfig{Enabled: true, Retention: config.Duration(time.Hour), MaxJobs: 100, MaxConcurrent: 1, MaxInputBytes: 1 << 20}
|
|
}
|
|
|
|
func waitState(t *testing.T, m *Manager, id, state string) Job {
|
|
t.Helper()
|
|
deadline := time.Now().Add(2 * time.Second)
|
|
for time.Now().Before(deadline) {
|
|
j, ok := m.Get(id, "t", "u", false)
|
|
if ok && j.State == state {
|
|
return j
|
|
}
|
|
time.Sleep(5 * time.Millisecond)
|
|
}
|
|
j, _ := m.Get(id, "t", "u", false)
|
|
t.Fatalf("job %s did not reach %s; got %#v", id, state, j)
|
|
return Job{}
|
|
}
|
|
|
|
func TestCreateRunPersistAndOutput(t *testing.T) {
|
|
dir := t.TempDir()
|
|
meta := filepath.Join(dir, "batch-jobs.json")
|
|
spool := filepath.Join(dir, "batch")
|
|
m, err := New(batchTestConfig(), meta, spool)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
ctx, cancel := context.WithCancel(context.Background())
|
|
defer cancel()
|
|
m.Start(ctx, func(ctx context.Context, j Job, in io.Reader, out io.Writer) RunResult {
|
|
b, _ := io.ReadAll(in)
|
|
if string(b) != `{"model":"m","input":"x"}` {
|
|
t.Errorf("input=%s", b)
|
|
}
|
|
io.WriteString(out, `{"id":"ok"}`)
|
|
return RunResult{HTTPStatus: 200, ResponseContentType: "application/json", RequestID: "req-1"}
|
|
})
|
|
j, err := m.Create(IdentitySnapshot{Tenant: "t", Subject: "u", Actor: "u", AuthType: "oidc"}, "/v1/responses", "m", []byte(`{"model":"m","input":"x"}`))
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
j = waitState(t, m, j.ID, StateCompleted)
|
|
if j.OutputRef == "" || j.ExecutionRequestID != "req-1" || j.HTTPStatus != 200 {
|
|
t.Fatalf("job=%#v", j)
|
|
}
|
|
f, _, err := m.OpenOutput(j.ID, "t", "u", false)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
b, _ := io.ReadAll(f)
|
|
f.Close()
|
|
if string(b) != `{"id":"ok"}` {
|
|
t.Fatalf("output=%s", b)
|
|
}
|
|
|
|
m2, err := New(batchTestConfig(), meta, spool)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
got, ok := m2.Get(j.ID, "t", "u", false)
|
|
if !ok || got.State != StateCompleted {
|
|
t.Fatalf("reloaded=%#v ok=%v", got, ok)
|
|
}
|
|
if _, ok := m2.Get(j.ID, "t", "other", false); ok {
|
|
t.Fatal("cross-actor job lookup must be hidden")
|
|
}
|
|
}
|
|
|
|
func TestPauseResumeAndCancel(t *testing.T) {
|
|
dir := t.TempDir()
|
|
m, err := New(batchTestConfig(), filepath.Join(dir, "jobs.json"), filepath.Join(dir, "spool"))
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
ctx, cancel := context.WithCancel(context.Background())
|
|
defer cancel()
|
|
var attempts atomic.Int32
|
|
m.Start(ctx, func(ctx context.Context, j Job, in io.Reader, out io.Writer) RunResult {
|
|
n := attempts.Add(1)
|
|
if n == 1 {
|
|
<-ctx.Done()
|
|
return RunResult{HTTPStatus: 499, Error: context.Cause(ctx).Error()}
|
|
}
|
|
io.WriteString(out, "done")
|
|
return RunResult{HTTPStatus: 200, RequestID: "req-done"}
|
|
})
|
|
j, err := m.Create(IdentitySnapshot{Tenant: "t", Subject: "u", Actor: "u", AuthType: "oidc"}, "/api/chat", "m", []byte(`{"model":"m"}`))
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
waitState(t, m, j.ID, StateRunning)
|
|
if _, err := m.Pause(j.ID, "t", "u", false); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
waitState(t, m, j.ID, StatePaused)
|
|
if _, err := m.Resume(j.ID, "t", "u", false); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
waitState(t, m, j.ID, StateCompleted)
|
|
if attempts.Load() != 2 {
|
|
t.Fatalf("attempts=%d", attempts.Load())
|
|
}
|
|
|
|
j2, err := m.Create(IdentitySnapshot{Tenant: "t", Subject: "u", Actor: "u", AuthType: "oidc"}, "/api/chat", "m", []byte(`{"model":"m"}`))
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
// The runner completes quickly on later attempts, so cancel while queued by
|
|
// first pausing it synchronously.
|
|
if _, err := m.Pause(j2.ID, "t", "u", false); err != nil && err != ErrInvalidState {
|
|
t.Fatal(err)
|
|
}
|
|
cur, _ := m.Get(j2.ID, "t", "u", false)
|
|
if cur.State == StatePaused {
|
|
if _, err := m.Cancel(j2.ID, "t", "u", false); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
waitState(t, m, j2.ID, StateCancelled)
|
|
}
|
|
}
|
|
|
|
func TestRetentionDeletesContentFiles(t *testing.T) {
|
|
dir := t.TempDir()
|
|
cfg := batchTestConfig()
|
|
cfg.Retention = config.Duration(time.Minute)
|
|
m, err := New(cfg, filepath.Join(dir, "jobs.json"), filepath.Join(dir, "spool"))
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
now := time.Date(2026, 9, 8, 8, 0, 0, 0, time.UTC)
|
|
m.now = func() time.Time { return now }
|
|
j, err := m.Create(IdentitySnapshot{Tenant: "t", Subject: "u", Actor: "u", AuthType: "oidc"}, "/api/chat", "m", []byte(`{"model":"m"}`))
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if _, err := m.Cancel(j.ID, "t", "u", false); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
job, _ := m.Get(j.ID, "t", "u", false)
|
|
inputPath := m.refPath(job.InputRef)
|
|
if _, err := os.Stat(inputPath); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
now = now.Add(2 * time.Minute)
|
|
if err := m.Compact(); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if _, ok := m.Get(j.ID, "t", "u", false); ok {
|
|
t.Fatal("expired job retained")
|
|
}
|
|
if _, err := os.Stat(inputPath); !os.IsNotExist(err) {
|
|
t.Fatalf("input still exists: %v", err)
|
|
}
|
|
}
|
|
|
|
func TestInputLimit(t *testing.T) {
|
|
dir := t.TempDir()
|
|
cfg := batchTestConfig()
|
|
cfg.MaxInputBytes = 4
|
|
m, _ := New(cfg, filepath.Join(dir, "jobs.json"), filepath.Join(dir, "spool"))
|
|
_, err := m.Create(IdentitySnapshot{Tenant: "t", Actor: "u"}, "/api/chat", "m", []byte(strings.Repeat("x", 5)))
|
|
if err != ErrInputTooLarge {
|
|
t.Fatalf("err=%v", err)
|
|
}
|
|
}
|
|
|
|
func TestWaitPersistsRestartSafeStateAfterShutdown(t *testing.T) {
|
|
dir := t.TempDir()
|
|
m, err := New(batchTestConfig(), filepath.Join(dir, "jobs.json"), filepath.Join(dir, "spool"))
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
root, cancel := context.WithCancel(context.Background())
|
|
m.Start(root, func(ctx context.Context, j Job, in io.Reader, out io.Writer) RunResult {
|
|
<-ctx.Done()
|
|
return RunResult{HTTPStatus: 499, Error: context.Cause(ctx).Error()}
|
|
})
|
|
j, err := m.Create(IdentitySnapshot{Tenant: "t", Subject: "u", Actor: "u", AuthType: "oidc"}, "/api/chat", "m", []byte(`{"model":"m"}`))
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
waitState(t, m, j.ID, StateRunning)
|
|
cancel()
|
|
wctx, wcancel := context.WithTimeout(context.Background(), time.Second)
|
|
defer wcancel()
|
|
if err := m.Wait(wctx); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
got, ok := m.Get(j.ID, "t", "u", false)
|
|
if !ok || got.State != StateQueued || !strings.Contains(got.Error, "shutdown") {
|
|
t.Fatalf("restart-safe state not persisted: %#v ok=%v", got, ok)
|
|
}
|
|
m2, err := New(batchTestConfig(), filepath.Join(dir, "jobs.json"), filepath.Join(dir, "spool"))
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
reloaded, ok := m2.Get(j.ID, "t", "u", false)
|
|
if !ok || reloaded.State != StateQueued {
|
|
t.Fatalf("reloaded state=%#v ok=%v", reloaded, ok)
|
|
}
|
|
}
|