Files
2026-09-11 06:14:38 +02:00

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)
}
}