-
This commit is contained in:
@@ -0,0 +1,216 @@
|
||||
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)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user