| 1 | package imageinput |
| 2 | |
| 3 | import ( |
| 4 | "context" |
| 5 | "errors" |
| 6 | "sync/atomic" |
| 7 | "testing" |
| 8 | |
| 9 | "reasonix/internal/event" |
| 10 | "reasonix/internal/provider" |
| 11 | ) |
| 12 | |
| 13 | type fakeProvider struct { |
| 14 | calls atomic.Int32 |
| 15 | entered chan struct{} |
| 16 | release chan struct{} |
| 17 | fail bool |
| 18 | seen chan provider.Request |
| 19 | } |
| 20 | |
| 21 | func (*fakeProvider) Name() string { return "vision" } |
| 22 | func (p *fakeProvider) Stream(ctx context.Context, r provider.Request) (<-chan provider.Chunk, error) { |
| 23 | p.calls.Add(1) |
| 24 | if p.seen != nil { |
| 25 | p.seen <- r |
| 26 | } |
| 27 | if p.entered != nil { |
| 28 | p.entered <- struct{}{} |
| 29 | } |
| 30 | if p.fail { |
| 31 | return nil, errors.New("test failure") |
| 32 | } |
| 33 | out := make(chan provider.Chunk, 2) |
| 34 | go func() { |
| 35 | defer close(out) |
| 36 | if p.release != nil { |
| 37 | select { |
| 38 | case <-ctx.Done(): |
| 39 | return |
| 40 | case <-p.release: |
| 41 | } |
| 42 | } |
| 43 | out <- provider.Chunk{Type: provider.ChunkText, Text: "left red, right blue, OCR Z7"} |
| 44 | out <- provider.Chunk{Type: provider.ChunkUsage, Usage: &provider.Usage{}} |
| 45 | }() |
| 46 | return out, nil |
| 47 | } |
| 48 | func service(p *fakeProvider) *Service { |
| 49 | return New(Config{Model: "vision/model", Resolve: func(string) (provider.Provider, error) { return p, nil }}) |
| 50 | } |
| 51 | func TestSerializedCacheAndCancellation(t *testing.T) { |
| 52 | p := &fakeProvider{entered: make(chan struct{}, 2), release: make(chan struct{})} |
| 53 | s := service(p) |
| 54 | refs := []string{"data:image/png;base64,QUFB"} |
| 55 | done := make(chan error, 2) |
| 56 | go func() { |
| 57 | _, err := s.Understand(context.Background(), "text/model", refs, nil, event.Discard) |
| 58 | done <- err |
| 59 | }() |
| 60 | <-p.entered |
| 61 | ctx, cancel := context.WithCancel(context.Background()) |
| 62 | cancel() |
| 63 | if _, err := s.Understand(ctx, "text/model", refs, nil, event.Discard); !errors.Is(err, context.Canceled) { |
| 64 | t.Fatalf("canceled queue: %v", err) |
| 65 | } |
| 66 | go func() { |
| 67 | _, err := s.Understand(context.Background(), "text/model", refs, nil, event.Discard) |
| 68 | done <- err |
| 69 | }() |
| 70 | close(p.release) |
| 71 | for range 2 { |
| 72 | if err := <-done; err != nil { |
| 73 | t.Fatal(err) |
| 74 | } |
| 75 | } |
| 76 | if p.calls.Load() != 1 { |
| 77 | t.Fatalf("calls %d", p.calls.Load()) |
| 78 | } |
| 79 | _, err := s.Understand(context.Background(), "text/model", []string{"data:image/png;base64,QkJC"}, nil, event.Discard) |
| 80 | if err != nil { |
| 81 | t.Fatal(err) |
| 82 | } |
| 83 | if p.calls.Load() != 2 { |
| 84 | t.Fatal("changed bytes reused stale summary") |
| 85 | } |
| 86 | } |
| 87 | func TestRestoredCacheAndMutableURL(t *testing.T) { |
| 88 | p := &fakeProvider{} |
| 89 | s := service(p) |
| 90 | refs := []string{"data:image/png;base64,QUFB"} |
| 91 | v, err := s.Understand(context.Background(), "text/model", refs, nil, event.Discard) |
| 92 | if err != nil { |
| 93 | t.Fatal(err) |
| 94 | } |
| 95 | restored := service(p) |
| 96 | got, err := restored.Understand(context.Background(), "text/model", refs, func() []provider.Message { return []provider.Message{{Role: provider.RoleTool, VisionSummary: v}} }, event.Discard) |
| 97 | if err != nil || got.Summary != v.Summary || p.calls.Load() != 1 { |
| 98 | t.Fatalf("restored: %v %v", got, err) |
| 99 | } |
| 100 | for range 2 { |
| 101 | _, err = restored.Understand(context.Background(), "text/model", []string{"https://example.test/mutable.png"}, nil, event.Discard) |
| 102 | if err != nil { |
| 103 | t.Fatal(err) |
| 104 | } |
| 105 | } |
| 106 | if p.calls.Load() != 3 { |
| 107 | t.Fatal("URL cache incorrectly reused") |
| 108 | } |
| 109 | } |
| 110 | func TestSelectionFailuresAndProviderFileIsolation(t *testing.T) { |
| 111 | p := &fakeProvider{} |
| 112 | cfg := Config{Model: "auto", Select: func(current, mode string) (string, bool) { |
| 113 | if current != "text/model" { |
| 114 | t.Fatal(current) |
| 115 | } |
| 116 | return "vision/model", true |
| 117 | }, Resolve: func(string) (provider.Provider, error) { return p, nil }} |
| 118 | s := New(cfg) |
| 119 | if _, err := s.Understand(context.Background(), "text/model", []string{"file-api-abc"}, nil, event.Discard); err == nil { |
| 120 | t.Fatal("cross-provider file ID accepted") |
| 121 | } |
| 122 | cfg.Select = func(string, string) (string, bool) { return "", false } |
| 123 | if _, err := New(cfg).Understand(context.Background(), "text/model", []string{"data:image/png;base64,QUFB"}, nil, event.Discard); err == nil { |
| 124 | t.Fatal("missing auto candidate accepted") |
| 125 | } |
| 126 | if _, err := New(Config{}).Understand(context.Background(), "text/model", nil, nil, event.Discard); err == nil { |
| 127 | t.Fatal("disabled accepted") |
| 128 | } |
| 129 | } |
| 130 | func TestCancelInFlightDoesNotCache(t *testing.T) { |
| 131 | p := &fakeProvider{entered: make(chan struct{}, 2), release: make(chan struct{})} |
| 132 | s := service(p) |
| 133 | ctx, cancel := context.WithCancel(context.Background()) |
| 134 | done := make(chan error, 1) |
| 135 | go func() { |
| 136 | _, err := s.Understand(ctx, "text/model", []string{"data:image/png;base64,QUFB"}, nil, event.Discard) |
| 137 | done <- err |
| 138 | }() |
| 139 | <-p.entered |
| 140 | cancel() |
| 141 | if err := <-done; !errors.Is(err, context.Canceled) { |
| 142 | t.Fatal(err) |
| 143 | } |
| 144 | if s.cached != nil { |
| 145 | t.Fatal("canceled request cached") |
| 146 | } |
| 147 | } |
| 148 | |
| 149 | func TestCachedSummaryDoesNotEmitAdditionalUsage(t *testing.T) { |
| 150 | p := &fakeProvider{} |
| 151 | s := service(p) |
| 152 | var usage atomic.Int32 |
| 153 | sink := event.FuncSink(func(e event.Event) { |
| 154 | if e.Kind == event.Usage { |
| 155 | if e.ModelRef != "vision/model" || e.UsageSource != event.UsageSourceClassifier { |
| 156 | t.Errorf("usage attribution: %+v", e) |
| 157 | } |
| 158 | usage.Add(1) |
| 159 | } |
| 160 | }) |
| 161 | for range 2 { |
| 162 | if _, err := s.Understand(context.Background(), "text/model", []string{"data:image/png;base64,QUFB"}, nil, sink); err != nil { |
| 163 | t.Fatal(err) |
| 164 | } |
| 165 | } |
| 166 | if usage.Load() != 1 { |
| 167 | t.Fatalf("usage events=%d", usage.Load()) |
| 168 | } |
| 169 | } |
| 170 |