summaryrefslogtreecommitdiff
path: root/internal/comic/comic_test.go
diff options
context:
space:
mode:
Diffstat (limited to 'internal/comic/comic_test.go')
-rw-r--r--internal/comic/comic_test.go56
1 files changed, 53 insertions, 3 deletions
diff --git a/internal/comic/comic_test.go b/internal/comic/comic_test.go
index 3b89871..319c7b2 100644
--- a/internal/comic/comic_test.go
+++ b/internal/comic/comic_test.go
@@ -143,7 +143,7 @@ func TestArtistAndRunnerEndToEndWithFakes(t *testing.T) {
t.Parallel()
originalLeakValidation := validateImagePromptLeakageFn
- validateImagePromptLeakageFn = func(context.Context, string, string) error { return nil }
+ validateImagePromptLeakageFn = func(context.Context, string, string, string) error { return nil }
t.Cleanup(func() {
validateImagePromptLeakageFn = originalLeakValidation
})
@@ -201,7 +201,7 @@ func TestRunnerPropagatesRenderFailures(t *testing.T) {
originalSleep := sleep
sleep = func(time.Duration) {}
originalLeakValidation := validateImagePromptLeakageFn
- validateImagePromptLeakageFn = func(context.Context, string, string) error { return nil }
+ validateImagePromptLeakageFn = func(context.Context, string, string, string) error { return nil }
t.Cleanup(func() {
sleep = originalSleep
validateImagePromptLeakageFn = originalLeakValidation
@@ -237,7 +237,7 @@ func TestRunnerPropagatesRenderFailures(t *testing.T) {
func TestDrawComicPagesUsesOneStyleAcrossTheWholePDF(t *testing.T) {
originalLeakValidation := validateImagePromptLeakageFn
- validateImagePromptLeakageFn = func(context.Context, string, string) error { return nil }
+ validateImagePromptLeakageFn = func(context.Context, string, string, string) error { return nil }
t.Cleanup(func() {
validateImagePromptLeakageFn = originalLeakValidation
})
@@ -278,6 +278,38 @@ func TestDrawComicPagesUsesOneStyleAcrossTheWholePDF(t *testing.T) {
}
}
+func TestDrawComicPagesChainsReferenceImages(t *testing.T) {
+ originalLeakValidation := validateImagePromptLeakageFn
+ validateImagePromptLeakageFn = func(context.Context, string, string, string) error { return nil }
+ t.Cleanup(func() {
+ validateImagePromptLeakageFn = originalLeakValidation
+ })
+
+ provider := &refTrackingImageProvider{t: t}
+ artist := NewArtist(&ArtistConfig{
+ ImageProvider: provider,
+ Prompts: fakePromptRenderer{},
+ OutputDir: t.TempDir(),
+ UltraRealistic: false,
+ })
+
+ if _, err := artist.DrawComicPages(context.Background(), "история", "библия", "slug", []WordEntry{{Word: "ябълка"}}, nil); err != nil {
+ t.Fatalf("DrawComicPages() error = %v", err)
+ }
+ if len(provider.referenceCounts) < 3 {
+ t.Fatalf("referenceCounts = %v, want multiple page generations", provider.referenceCounts)
+ }
+ if provider.referenceCounts[0] != 0 {
+ t.Fatalf("cover refs = %d, want 0", provider.referenceCounts[0])
+ }
+ if provider.referenceCounts[1] != 1 {
+ t.Fatalf("first story page refs = %d, want 1", provider.referenceCounts[1])
+ }
+ if provider.referenceCounts[2] < 1 {
+ t.Fatalf("second page refs = %d, want chained references", provider.referenceCounts[2])
+ }
+}
+
func TestConvertToStereoFallsBackToCopyWhenFFmpegMissing(t *testing.T) {
originalLookPath := lookPath
lookPath = func(string) (string, error) {
@@ -406,6 +438,24 @@ func (f fakeImageProvider) GenerateImage(_ context.Context, _ string, outputFile
return nil
}
+type refTrackingImageProvider struct {
+ t *testing.T
+ referenceCounts []int
+}
+
+func (p *refTrackingImageProvider) Name() string { return "ref-tracking-image" }
+func (p *refTrackingImageProvider) IsAvailable() error { return nil }
+func (p *refTrackingImageProvider) GenerateImage(ctx context.Context, prompt, outputFile string) error {
+ return p.GenerateImageWithReferences(ctx, prompt, outputFile, nil)
+}
+func (p *refTrackingImageProvider) GenerateImageWithReferences(_ context.Context, _ string, outputFile string, refs [][]byte) error {
+ p.referenceCounts = append(p.referenceCounts, len(refs))
+ if err := os.WriteFile(outputFile, []byte("png"), 0o644); err != nil {
+ p.t.Fatal(err)
+ }
+ return nil
+}
+
type failingImageProvider struct{}
func (failingImageProvider) Name() string { return "failing-image" }