From 7109df3ae03661d750c263b89b2d64d476377541 Mon Sep 17 00:00:00 2001 From: Paul Buetow Date: Sun, 19 Apr 2026 22:20:51 +0300 Subject: v4: add provider-neutral image layer --- internal/image/doc.go | 2 + internal/image/download.go | 223 ++++++++++++++ internal/image/download_test.go | 144 +++++++++ internal/image/gemini.go | 599 ++++++++++++++++++++++++++++++++++++ internal/image/gemini_test.go | 333 ++++++++++++++++++++ internal/image/prompt.go | 167 ++++++++++ internal/image/prompt_test.go | 60 ++++ internal/image/registry.go | 74 +++++ internal/image/styles.go | 100 ++++++ internal/image/styles_test.go | 49 +++ internal/image/test_helpers_test.go | 38 +++ internal/image/types.go | 113 +++++++ internal/image/types_test.go | 101 ++++++ 13 files changed, 2003 insertions(+) create mode 100644 internal/image/doc.go create mode 100644 internal/image/download.go create mode 100644 internal/image/download_test.go create mode 100644 internal/image/gemini.go create mode 100644 internal/image/gemini_test.go create mode 100644 internal/image/prompt.go create mode 100644 internal/image/prompt_test.go create mode 100644 internal/image/registry.go create mode 100644 internal/image/styles.go create mode 100644 internal/image/styles_test.go create mode 100644 internal/image/test_helpers_test.go create mode 100644 internal/image/types.go create mode 100644 internal/image/types_test.go diff --git a/internal/image/doc.go b/internal/image/doc.go new file mode 100644 index 0000000..ae5d586 --- /dev/null +++ b/internal/image/doc.go @@ -0,0 +1,2 @@ +// Package image provides provider-neutral image generation and download helpers. +package image diff --git a/internal/image/download.go b/internal/image/download.go new file mode 100644 index 0000000..7a196fa --- /dev/null +++ b/internal/image/download.go @@ -0,0 +1,223 @@ +package image + +import ( + "context" + "fmt" + "io" + "os" + "path/filepath" + "strings" +) + +// DownloadOptions configures image download behavior. +type DownloadOptions struct { + OutputDir string + OverwriteExisting bool + CreateDir bool + FileNamePattern string + MaxSizeBytes int64 +} + +// DefaultDownloadOptions returns sensible defaults for image downloads. +func DefaultDownloadOptions() *DownloadOptions { + return &DownloadOptions{ + OutputDir: "./images", + OverwriteExisting: false, + CreateDir: true, + FileNamePattern: "{word}_{source}", + MaxSizeBytes: 10 * 1024 * 1024, + } +} + +// Downloader handles image downloads from search results. +type Downloader struct { + provider ImageProvider + options *DownloadOptions +} + +// NewDownloader creates a new image downloader. +func NewDownloader(provider ImageProvider, options *DownloadOptions) *Downloader { + if options == nil { + options = DefaultDownloadOptions() + } + return &Downloader{ + provider: provider, + options: options, + } +} + +// DownloadImage downloads a single image to the specified path. +func (d *Downloader) DownloadImage(ctx context.Context, result *SearchResult, outputPath string) (err error) { + if d == nil || d.provider == nil { + return fmt.Errorf("image provider is required") + } + if result == nil { + return fmt.Errorf("search result is required") + } + + dir := filepath.Dir(outputPath) + if d.options != nil && d.options.CreateDir && dir != "" && dir != "." { + if err := os.MkdirAll(dir, 0o755); err != nil { + return fmt.Errorf("create output dir %q: %w", dir, err) + } + } + + if d.options != nil && !d.options.OverwriteExisting { + if _, err := os.Stat(outputPath); err == nil { + return fmt.Errorf("output file exists: %s", outputPath) + } + } + + reader, err := d.provider.Download(ctx, result.URL) + if err != nil { + return fmt.Errorf("download %q: %w", result.URL, err) + } + defer func() { + _ = reader.Close() + }() + + file, err := os.Create(outputPath) + if err != nil { + return fmt.Errorf("create output file %q: %w", outputPath, err) + } + defer func() { + if closeErr := file.Close(); err == nil && closeErr != nil { + err = fmt.Errorf("close output file %q: %w", outputPath, closeErr) + } + }() + + if d.options != nil && d.options.MaxSizeBytes > 0 { + written, copyErr := io.CopyN(file, reader, d.options.MaxSizeBytes) + if copyErr != nil && copyErr != io.EOF { + _ = os.Remove(outputPath) + return fmt.Errorf("write output file %q: %w", outputPath, copyErr) + } + + if written == d.options.MaxSizeBytes { + var probe [1]byte + if n, probeErr := reader.Read(probe[:]); n > 0 || probeErr != io.EOF { + _ = os.Remove(outputPath) + return fmt.Errorf("image exceeds max size %d bytes", d.options.MaxSizeBytes) + } + } + } else { + if _, err = io.Copy(file, reader); err != nil { + _ = os.Remove(outputPath) + return fmt.Errorf("write output file %q: %w", outputPath, err) + } + } + + if err := file.Sync(); err != nil { + return fmt.Errorf("sync output file %q: %w", outputPath, err) + } + + if attribution := d.provider.GetAttribution(result); attribution != "" { + attrPath := strings.TrimSuffix(outputPath, filepath.Ext(outputPath)) + "_attribution.txt" + if err := os.WriteFile(attrPath, []byte(attribution), 0o644); err != nil { + fmt.Fprintf(os.Stderr, "Warning: failed to save attribution: %v\n", err) + } + } + + return nil +} + +// DownloadBestMatch downloads the best matching image for a query. +func (d *Downloader) DownloadBestMatch(ctx context.Context, query string) (*SearchResult, string, error) { + opts := DefaultSearchOptions(query) + opts.PerPage = 5 + return d.DownloadBestMatchWithOptions(ctx, opts) +} + +// DownloadBestMatchWithOptions downloads the best matching image for given search options. +func (d *Downloader) DownloadBestMatchWithOptions(ctx context.Context, opts *SearchOptions) (*SearchResult, string, error) { + if d == nil || d.provider == nil { + return nil, "", fmt.Errorf("image provider is required") + } + if opts == nil { + return nil, "", fmt.Errorf("search options are required") + } + + searchOpts := *opts + searchOpts.PerPage = 5 + + results, err := d.provider.Search(ctx, &searchOpts) + if err != nil { + return nil, "", fmt.Errorf("search images: %w", err) + } + if len(results) == 0 { + return nil, "", fmt.Errorf("no images found for %q", opts.Query) + } + + for i, result := range results { + filename := d.generateFileName(opts.Query, &result, i) + outputDir := "./images" + if d.options != nil && d.options.OutputDir != "" { + outputDir = d.options.OutputDir + } + outputPath := filepath.Join(outputDir, filename) + + if err := d.DownloadImage(ctx, &result, outputPath); err == nil { + return &result, outputPath, nil + } else { + fmt.Fprintf(os.Stderr, "Warning: failed to download image %d: %v\n", i+1, err) + } + } + + return nil, "", fmt.Errorf("no downloadable images found for %q", opts.Query) +} + +func (d *Downloader) generateFileName(word string, result *SearchResult, index int) string { + filename := "" + if d != nil && d.options != nil { + filename = d.options.FileNamePattern + } + if filename == "" { + filename = "{word}_{source}" + } + + filename = strings.ReplaceAll(filename, "{word}", sanitizeFileName(word)) + if result != nil { + filename = strings.ReplaceAll(filename, "{source}", result.Source) + filename = strings.ReplaceAll(filename, "{id}", result.ID) + } + filename = strings.ReplaceAll(filename, "{index}", fmt.Sprintf("%d", index)) + + ext := "" + if result != nil { + ext = filepath.Ext(result.URL) + if strings.HasPrefix(result.URL, geminiDataPrefix) { + ext = ".png" + } else if ext == "" || len(ext) > 5 { + ext = ".jpg" + } + } + + if filepath.Ext(filename) == "" { + filename += ext + } + + return filename +} + +func sanitizeFileName(name string) string { + replacer := strings.NewReplacer( + "/", "_", + "\\", "_", + ":", "_", + "*", "_", + "?", "_", + "\"", "_", + "<", "_", + ">", "_", + "|", "_", + " ", "_", + ".", "_", + ) + + sanitized := replacer.Replace(name) + if len(sanitized) > 50 { + sanitized = sanitized[:50] + } + + return sanitized +} diff --git a/internal/image/download_test.go b/internal/image/download_test.go new file mode 100644 index 0000000..588d875 --- /dev/null +++ b/internal/image/download_test.go @@ -0,0 +1,144 @@ +package image + +import ( + "context" + "io" + "os" + "path/filepath" + "strings" + "testing" +) + +type mockDownloaderProvider struct { + results []SearchResult + searchErr error + payload string + attribution string + searchQueries []string + downloadURLs []string +} + +func (m *mockDownloaderProvider) Name() string { return "mock" } + +func (m *mockDownloaderProvider) Search(_ context.Context, opts *SearchOptions) ([]SearchResult, error) { + if opts != nil { + m.searchQueries = append(m.searchQueries, opts.Query) + } + if m.searchErr != nil { + return nil, m.searchErr + } + return append([]SearchResult(nil), m.results...), nil +} + +func (m *mockDownloaderProvider) Download(_ context.Context, url string) (io.ReadCloser, error) { + m.downloadURLs = append(m.downloadURLs, url) + return io.NopCloser(strings.NewReader(m.payload)), nil +} + +func (m *mockDownloaderProvider) GetAttribution(*SearchResult) string { + return m.attribution +} + +func TestDownloaderGenerateFileName_DataURIUsesPNG(t *testing.T) { + t.Parallel() + + d := NewDownloader(&mockDownloaderProvider{}, &DownloadOptions{ + FileNamePattern: "{word}_{source}", + }) + result := &SearchResult{ + URL: "data:image/png;base64,AAAA", + Source: Gemini, + } + + if got := d.generateFileName("ябълка", result, 0); got != "ябълка_gemini.png" { + t.Fatalf("generateFileName() = %q, want %q", got, "ябълка_gemini.png") + } +} + +func TestDownloadImageWritesAttribution(t *testing.T) { + t.Parallel() + + provider := &mockDownloaderProvider{ + payload: "image-bytes", + attribution: "attribution text", + } + d := NewDownloader(provider, &DownloadOptions{ + OutputDir: t.TempDir(), + CreateDir: true, + OverwriteExisting: false, + FileNamePattern: "{word}_{source}", + MaxSizeBytes: 10 * 1024 * 1024, + }) + + outputPath := filepath.Join(d.options.OutputDir, "ябълка_gemini.png") + if err := d.DownloadImage(context.Background(), &SearchResult{ + URL: "https://example.com/image.png", + Source: Gemini, + ID: "1", + }, outputPath); err != nil { + t.Fatalf("DownloadImage() error = %v", err) + } + + data, err := os.ReadFile(outputPath) + if err != nil { + t.Fatalf("ReadFile() error = %v", err) + } + if string(data) != "image-bytes" { + t.Fatalf("downloaded file = %q, want %q", string(data), "image-bytes") + } + + attrPath := strings.TrimSuffix(outputPath, filepath.Ext(outputPath)) + "_attribution.txt" + attr, err := os.ReadFile(attrPath) + if err != nil { + t.Fatalf("ReadFile(attribution) error = %v", err) + } + if string(attr) != "attribution text" { + t.Fatalf("attribution = %q, want %q", string(attr), "attribution text") + } +} + +func TestDownloadBestMatchWithOptions(t *testing.T) { + t.Parallel() + + provider := &mockDownloaderProvider{ + results: []SearchResult{ + { + ID: "1", + URL: "https://example.com/image1.jpg", + Source: Gemini, + }, + }, + payload: "image-bytes", + attribution: "attribution text", + } + d := NewDownloader(provider, &DownloadOptions{ + OutputDir: t.TempDir(), + CreateDir: true, + OverwriteExisting: true, + FileNamePattern: "{word}_{source}", + MaxSizeBytes: 10 * 1024 * 1024, + }) + + result, path, err := d.DownloadBestMatchWithOptions(context.Background(), &SearchOptions{Query: "ябълка"}) + if err != nil { + t.Fatalf("DownloadBestMatchWithOptions() error = %v", err) + } + if result == nil || result.ID != "1" { + t.Fatalf("DownloadBestMatchWithOptions() result = %+v, want ID 1", result) + } + if !strings.HasSuffix(path, ".jpg") { + t.Fatalf("DownloadBestMatchWithOptions() path = %q, want jpg suffix", path) + } + if _, err := os.Stat(path); err != nil { + t.Fatalf("downloaded file missing: %v", err) + } +} + +func TestDownloadImageRejectsNilResult(t *testing.T) { + t.Parallel() + + d := NewDownloader(&mockDownloaderProvider{}, nil) + if err := d.DownloadImage(context.Background(), nil, filepath.Join(t.TempDir(), "out.png")); err == nil { + t.Fatal("expected error for nil result") + } +} diff --git a/internal/image/gemini.go b/internal/image/gemini.go new file mode 100644 index 0000000..2ccd6fa --- /dev/null +++ b/internal/image/gemini.go @@ -0,0 +1,599 @@ +package image + +import ( + "bytes" + "context" + "crypto/md5" + "encoding/base64" + "encoding/hex" + "fmt" + "image" + _ "image/jpeg" + "image/png" + "io" + "net/http" + "strings" + "time" + + "google.golang.org/genai" + + "codeberg.org/snonux/comicforge/internal/apicircuit" + "codeberg.org/snonux/comicforge/internal/httpctx" +) + +const ( + // DefaultGeminiImageModel is the Gemini image model used for image generation. + DefaultGeminiImageModel = "gemini-3.1-flash-image-preview" + + // DefaultGeminiTextModel is the Gemini text model used for scene generation. + DefaultGeminiTextModel = "gemini-2.5-flash" + + geminiAspectRatio = "4:3" + geminiDataPrefix = "data:image/png;base64," + geminiSource = Gemini + maxCustomPrompt = 4000 +) + +// GeminiConfig holds the settings needed to build a Gemini-backed image provider. +type GeminiConfig struct { + APIKey string + Model string + TextModel string +} + +// GeminiProvider implements ImageProvider for Google Gemini image generation. +type GeminiProvider struct { + client *genai.Client + initErr error + config *GeminiConfig + lastPrompt string + + // PromptCallback runs after the prompt is generated and before image creation. + PromptCallback func(prompt string) +} + +var _ ImageProvider = (*GeminiProvider)(nil) + +var imageHTTPClient = httpctx.ImageDownloadHTTPClient() +var newGeminiClient = httpctx.NewGenAIClient +var geminiGenerateText = func(ctx context.Context, c *GeminiProvider, model, systemPrompt, userPrompt string, temperature float32, maxOutputTokens int32) (string, error) { + return c.generateText(ctx, model, systemPrompt, userPrompt, temperature, maxOutputTokens) +} +var geminiGenerateImage = func(ctx context.Context, c *GeminiProvider, prompt, aspectRatio string) ([]byte, string, error) { + return c.generateImage(ctx, prompt, aspectRatio) +} + +// NewGeminiProvider creates a new Gemini image provider. +func NewGeminiProvider(config *GeminiConfig) *GeminiProvider { + normalized := normalizeGeminiConfig(config) + client := &GeminiProvider{config: normalized} + + if normalized.APIKey == "" { + return client + } + + genaiClient, err := newGeminiClient(context.Background(), &genai.ClientConfig{ + APIKey: normalized.APIKey, + Backend: genai.BackendGeminiAPI, + }) + if err != nil { + client.initErr = err + return client + } + + client.client = genaiClient + return client +} + +// Search generates an educational image for the requested word or phrase. +func (c *GeminiProvider) Search(ctx context.Context, opts *SearchOptions) ([]SearchResult, error) { + ctx, cancel := httpctx.WithTimeoutUnlessSet(ctx, httpctx.OperationTimeoutDefault) + defer cancel() + + if err := c.ensureReady(); err != nil { + return nil, err + } + if opts == nil { + return nil, &SearchError{ + Provider: geminiSource, + Code: "INVALID_OPTIONS", + Message: "search options are required", + } + } + + prompt, translatedWord, err := c.buildPrompt(ctx, opts) + if err != nil { + return nil, err + } + + c.lastPrompt = prompt + if c.PromptCallback != nil { + c.PromptCallback(prompt) + } + + aspectRatio := geminiAspectRatio + if opts.AspectRatio != "" { + aspectRatio = opts.AspectRatio + } + + fmt.Printf("Gemini Image Generation Prompt (%d chars): %s\n", len(prompt), prompt) + fmt.Printf("Gemini Image Generation: Using model %q with aspect ratio %q\n", c.modelName(), aspectRatio) + + var imageBytes []byte + var mimeType string + if len(opts.ReferenceImages) > 0 { + imageBytes, mimeType, err = c.generateImageWithRefs(ctx, prompt, aspectRatio, opts.ReferenceImages) + } else { + imageBytes, mimeType, err = geminiGenerateImage(ctx, c, prompt, aspectRatio) + } + if err != nil { + if searchErr, ok := err.(*SearchError); ok { + return nil, searchErr + } + return nil, &SearchError{ + Provider: geminiSource, + Code: "API_ERROR", + Message: fmt.Sprintf("failed to generate image: %v", err), + } + } + + dataURL, err := encodeDataURL(imageBytes, mimeType) + if err != nil { + return nil, err + } + width, height, err := decodedImageDimensions(imageBytes) + if err != nil { + return nil, err + } + + description := fmt.Sprintf("Generated educational image for %s", opts.Query) + if translatedWord != "" { + description = fmt.Sprintf("%s (%s)", description, translatedWord) + } + + result := SearchResult{ + ID: c.generateImageID(opts.Query), + URL: dataURL, + ThumbnailURL: dataURL, + Width: width, + Height: height, + Description: description, + Attribution: "Generated by Google Gemini Nano Banana", + Source: geminiSource, + } + + return []SearchResult{result}, nil +} + +// Download returns the image bytes for a data URI or a remote URL. +func (c *GeminiProvider) Download(ctx context.Context, url string) (io.ReadCloser, error) { + if strings.HasPrefix(url, geminiDataPrefix) { + return decodeDataURL(url) + } + + req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil) + if err != nil { + return nil, err + } + + resp, err := imageHTTPClient.Do(req) + if err != nil { + return nil, err + } + + if resp.StatusCode != http.StatusOK { + if closeErr := resp.Body.Close(); closeErr != nil { + return nil, fmt.Errorf("HTTP %d: %s (failed to close response body: %v)", resp.StatusCode, resp.Status, closeErr) + } + return nil, fmt.Errorf("HTTP %d: %s", resp.StatusCode, resp.Status) + } + + return resp.Body, nil +} + +// GetAttribution returns attribution text for the generated image. +func (c *GeminiProvider) GetAttribution(result *SearchResult) string { + width, height := 0, 0 + if result != nil { + width = result.Width + height = result.Height + } + + var attribution strings.Builder + attribution.WriteString("Image generated by Google Gemini Nano Banana\n\n") + fmt.Fprintf(&attribution, "Model: %s\n", c.modelName()) + fmt.Fprintf(&attribution, "Text model: %s\n", c.textModelName()) + fmt.Fprintf(&attribution, "Aspect ratio: %s\n", geminiAspectRatio) + fmt.Fprintf(&attribution, "Size: %dx%d\n", width, height) + if result != nil && result.Description != "" { + fmt.Fprintf(&attribution, "Result: %s\n", result.Description) + } + fmt.Fprintf(&attribution, "\nPrompt used:\n%s\n", c.lastPrompt) + fmt.Fprintf(&attribution, "\nGenerated at: %s\n", time.Now().Format("2006-01-02 15:04:05")) + return attribution.String() +} + +// Name returns the provider name. +func (c *GeminiProvider) Name() string { + return geminiSource +} + +// LastPrompt returns the most recent image prompt. +func (c *GeminiProvider) LastPrompt() string { + return c.lastPrompt +} + +// SetPromptCallback registers a callback that runs after prompt generation. +func (c *GeminiProvider) SetPromptCallback(callback func(prompt string)) { + c.PromptCallback = callback +} + +func (c *GeminiProvider) ensureReady() error { + if c == nil || c.config == nil { + return &SearchError{ + Provider: geminiSource, + Code: "NO_CONFIG", + Message: "Gemini client not initialized", + } + } + if c.config.APIKey == "" { + return &SearchError{ + Provider: geminiSource, + Code: "NO_API_KEY", + Message: "Google API key not configured", + } + } + if c.initErr != nil { + return &SearchError{ + Provider: geminiSource, + Code: "CLIENT_INIT_FAILED", + Message: fmt.Sprintf("failed to initialize client: %v", c.initErr), + } + } + if c.client == nil { + return &SearchError{ + Provider: geminiSource, + Code: "CLIENT_NOT_READY", + Message: "Gemini client not initialized", + } + } + + return nil +} + +func (c *GeminiProvider) resolveTranslation(_ context.Context, opts *SearchOptions, translation string) (string, error) { + if translation != "" { + fmt.Printf("Using provided translation: %s -> %s\n", opts.Query, translation) + return translation, nil + } + + return opts.Query, nil +} + +func (c *GeminiProvider) resolvePrompt(ctx context.Context, opts *SearchOptions, translatedWord string) (string, error) { + if customPrompt := strings.TrimSpace(opts.CustomPrompt); customPrompt != "" { + if len(customPrompt) > maxCustomPrompt { + customPrompt = customPrompt[:maxCustomPrompt-3] + "..." + } + fmt.Printf("Using custom prompt: %s\n", customPrompt) + return customPrompt, nil + } + + return c.createEducationalPrompt(ctx, opts.Query, translatedWord), nil +} + +func (c *GeminiProvider) buildPrompt(ctx context.Context, opts *SearchOptions) (string, string, error) { + if opts == nil { + return "", "", &SearchError{ + Provider: geminiSource, + Code: "INVALID_OPTIONS", + Message: "search options are required", + } + } + + translation := strings.TrimSpace(opts.Translation) + if customPrompt := strings.TrimSpace(opts.CustomPrompt); customPrompt != "" { + if len(customPrompt) > maxCustomPrompt { + customPrompt = customPrompt[:maxCustomPrompt-3] + "..." + } + fmt.Printf("Using custom prompt: %s\n", customPrompt) + return customPrompt, translation, nil + } + + translatedWord, err := c.resolveTranslation(ctx, opts, translation) + if err != nil { + return "", "", err + } + + prompt, err := c.resolvePrompt(ctx, opts, translatedWord) + if err != nil { + return "", "", err + } + + return prompt, translatedWord, nil +} + +// createEducationalPrompt generates a prompt optimized for image generation. +func (c *GeminiProvider) createEducationalPrompt(ctx context.Context, query, translation string) string { + subject := promptSubject(translation, query) + + scene, err := c.generateSceneDescription(ctx, query, translation) + if err != nil { + fmt.Printf(" Failed to generate scene: %v, using basic prompt\n", err) + scene = "" + } + if scene != "" { + scene = sanitizeSceneDescription(scene) + if !usableSceneDescription(scene) { + fmt.Printf(" Scene response was too short or generic, using basic prompt\n") + scene = "" + } + } + + selectedStyle := chooseArtisticStyle() + if selectedStyle == defaultArtisticStyle { + fmt.Printf(" No artistic styles available, using generic prompt\n") + } + fmt.Printf(" Using image style: %s\n", selectedStyle) + + return buildEducationalPrompt(selectedStyle, scene, subject) +} + +func (c *GeminiProvider) generateSceneDescription(ctx context.Context, query, translation string) (string, error) { + fmt.Printf("Gemini Scene Generation: Creating scene for %q (%s)\n", query, translation) + + scene, err := geminiGenerateText( + ctx, + c, + c.textModelName(), + "You are helping create educational flashcards for language learning. Generate a brief, vivid scene description that incorporates the given English word in a memorable, contextual way. The scene should be visually interesting and help with memory retention. Keep it to 1-2 sentences, focusing on visual elements that can be illustrated. The subject (the English word) should be the clear focal point of the image, prominent and centered.", + fmt.Sprintf("Create a scene description for the English word %q that would make a memorable flashcard image. Make sure %q is the main focus and most prominent element in the scene.", translation, translation), + 0.7, + 100, + ) + if err != nil { + return "", fmt.Errorf("scene generation failed: %w", err) + } + scene = sanitizeSceneDescription(scene) + if !usableSceneDescription(scene) { + return "", fmt.Errorf("scene generation returned unusable content") + } + + fmt.Printf("Generated scene: %s\n", scene) + return scene, nil +} + +func (c *GeminiProvider) generateText(ctx context.Context, model, systemPrompt, userPrompt string, temperature float32, maxOutputTokens int32) (string, error) { + temp := temperature + resp, err := apicircuit.Execute(nil, c.Name(), apicircuit.CapabilityText, func() (*genai.GenerateContentResponse, error) { + return c.client.Models.GenerateContent(ctx, model, []*genai.Content{ + genai.NewContentFromText(userPrompt, genai.RoleUser), + }, &genai.GenerateContentConfig{ + SystemInstruction: genai.NewContentFromText(systemPrompt, genai.RoleUser), + Temperature: &temp, + MaxOutputTokens: maxOutputTokens, + }) + }) + if err != nil { + return "", fmt.Errorf("gemini API error: %w", err) + } + + text := strings.TrimSpace(resp.Text()) + if text == "" { + return "", fmt.Errorf("no response received") + } + + return text, nil +} + +func (c *GeminiProvider) generateImage(ctx context.Context, prompt, aspectRatio string) ([]byte, string, error) { + if aspectRatio == "" { + aspectRatio = geminiAspectRatio + } + + cfg := &genai.GenerateContentConfig{ + ResponseModalities: []string{string(genai.ModalityImage)}, + ImageConfig: &genai.ImageConfig{ + AspectRatio: aspectRatio, + }, + } + + resp, err := apicircuit.Execute(nil, c.Name(), apicircuit.CapabilityImage, func() (*genai.GenerateContentResponse, error) { + return c.client.Models.GenerateContent(ctx, c.modelName(), []*genai.Content{ + genai.NewContentFromText(prompt, genai.RoleUser), + }, cfg) + }) + if err != nil { + return nil, "", &SearchError{ + Provider: geminiSource, + Code: "API_ERROR", + Message: fmt.Sprintf("failed to generate image: %v", err), + } + } + + imageBytes, mimeType, err := extractGeneratedImage(resp) + if err != nil { + return nil, "", &SearchError{ + Provider: geminiSource, + Code: "NO_RESULTS", + Message: err.Error(), + } + } + + return imageBytes, mimeType, nil +} + +func (c *GeminiProvider) generateImageWithRefs(ctx context.Context, prompt, aspectRatio string, refs [][]byte) ([]byte, string, error) { + if aspectRatio == "" { + aspectRatio = geminiAspectRatio + } + + cfg := &genai.GenerateContentConfig{ + ResponseModalities: []string{string(genai.ModalityImage)}, + ImageConfig: &genai.ImageConfig{AspectRatio: aspectRatio}, + } + + parts := make([]*genai.Part, 0, len(refs)+1) + for _, ref := range refs { + if len(ref) > 0 { + parts = append(parts, &genai.Part{ + InlineData: &genai.Blob{MIMEType: "image/png", Data: ref}, + }) + } + } + refNote := fmt.Sprintf( + "The %d reference image(s) above show the exact character appearance that must be preserved. "+ + "Every character, animal, or object must look identical in the new image. Now generate:\n\n", + len(refs), + ) + parts = append(parts, &genai.Part{Text: refNote + prompt}) + + resp, err := apicircuit.Execute(nil, c.Name(), apicircuit.CapabilityImage, func() (*genai.GenerateContentResponse, error) { + return c.client.Models.GenerateContent(ctx, c.modelName(), []*genai.Content{ + { + Role: string(genai.RoleUser), + Parts: parts, + }, + }, cfg) + }) + if err != nil { + return nil, "", &SearchError{ + Provider: geminiSource, + Code: "API_ERROR", + Message: fmt.Sprintf("failed to generate image with refs: %v", err), + } + } + + imageBytes, mimeType, err := extractGeneratedImage(resp) + if err != nil { + return nil, "", &SearchError{ + Provider: geminiSource, + Code: "NO_RESULTS", + Message: err.Error(), + } + } + + return imageBytes, mimeType, nil +} + +func extractGeneratedImage(response *genai.GenerateContentResponse) ([]byte, string, error) { + if response == nil { + return nil, "", fmt.Errorf("no response from Gemini") + } + + for _, candidate := range response.Candidates { + if candidate == nil || candidate.Content == nil { + continue + } + + for _, part := range candidate.Content.Parts { + if part == nil || part.InlineData == nil || len(part.InlineData.Data) == 0 { + continue + } + + mimeType := part.InlineData.MIMEType + if mimeType == "" { + mimeType = "image/png" + } + + return append([]byte(nil), part.InlineData.Data...), mimeType, nil + } + } + + return nil, "", fmt.Errorf("no image data returned from Gemini") +} + +func encodeDataURL(imageBytes []byte, mimeType string) (string, error) { + if len(imageBytes) == 0 { + return "", fmt.Errorf("no image bytes returned") + } + + normalizedBytes, err := normalizePNG(imageBytes, mimeType) + if err != nil { + return "", err + } + + return geminiDataPrefix + base64.StdEncoding.EncodeToString(normalizedBytes), nil +} + +func decodeDataURL(url string) (io.ReadCloser, error) { + header, payload, ok := strings.Cut(url, ",") + if !ok || !strings.HasPrefix(header, "data:") || !strings.Contains(header, ";base64") { + return nil, fmt.Errorf("unsupported data URI: %s", url) + } + + data, err := base64.StdEncoding.DecodeString(payload) + if err != nil { + return nil, fmt.Errorf("decode data URI: %w", err) + } + + return io.NopCloser(bytes.NewReader(data)), nil +} + +func normalizePNG(imageBytes []byte, mimeType string) ([]byte, error) { + if strings.EqualFold(strings.TrimSpace(mimeType), "image/png") { + return append([]byte(nil), imageBytes...), nil + } + + img, _, err := image.Decode(bytes.NewReader(imageBytes)) + if err != nil { + return nil, fmt.Errorf("decode generated image: %w", err) + } + + var buffer bytes.Buffer + if err := png.Encode(&buffer, img); err != nil { + return nil, fmt.Errorf("encode generated image as png: %w", err) + } + + return buffer.Bytes(), nil +} + +func decodedImageDimensions(imageBytes []byte) (int, int, error) { + cfg, _, err := image.DecodeConfig(bytes.NewReader(imageBytes)) + if err != nil { + return 0, 0, fmt.Errorf("decode generated image dimensions: %w", err) + } + + return cfg.Width, cfg.Height, nil +} + +func (c *GeminiProvider) generateImageID(word string) string { + hash := md5.Sum([]byte(word)) + return hex.EncodeToString(hash[:])[:8] +} + +func (c *GeminiProvider) modelName() string { + if c == nil || c.config == nil || strings.TrimSpace(c.config.Model) == "" { + return DefaultGeminiImageModel + } + + return c.config.Model +} + +func (c *GeminiProvider) textModelName() string { + if c == nil || c.config == nil || strings.TrimSpace(c.config.TextModel) == "" { + return DefaultGeminiTextModel + } + + return c.config.TextModel +} + +func normalizeGeminiConfig(config *GeminiConfig) *GeminiConfig { + normalized := &GeminiConfig{} + if config != nil { + *normalized = *config + } + + normalized.APIKey = strings.TrimSpace(normalized.APIKey) + normalized.Model = strings.TrimSpace(normalized.Model) + normalized.TextModel = strings.TrimSpace(normalized.TextModel) + + if normalized.Model == "" { + normalized.Model = DefaultGeminiImageModel + } + if normalized.TextModel == "" { + normalized.TextModel = DefaultGeminiTextModel + } + + return normalized +} diff --git a/internal/image/gemini_test.go b/internal/image/gemini_test.go new file mode 100644 index 0000000..8c34c6e --- /dev/null +++ b/internal/image/gemini_test.go @@ -0,0 +1,333 @@ +package image + +import ( + "bytes" + "context" + "encoding/base64" + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "google.golang.org/genai" +) + +func TestNewGeminiProvider(t *testing.T) { + t.Parallel() + + client := NewGeminiProvider(&GeminiConfig{APIKey: "test-key"}) + if client == nil { + t.Fatal("expected client") + } + if client.config == nil { + t.Fatal("expected normalized config") + } + if client.config.Model != DefaultGeminiImageModel { + t.Fatalf("expected default model %q, got %q", DefaultGeminiImageModel, client.config.Model) + } + if client.config.TextModel != DefaultGeminiTextModel { + t.Fatalf("expected default text model %q, got %q", DefaultGeminiTextModel, client.config.TextModel) + } + if client.Name() != Gemini { + t.Fatalf("Name() = %q, want %q", client.Name(), Gemini) + } +} + +func TestGeminiProvider_NoAPIKey(t *testing.T) { + client := NewGeminiProvider(&GeminiConfig{}) + + _, err := client.Search(context.Background(), DefaultSearchOptions("ябълка")) + if err == nil { + t.Fatal("expected error for missing API key") + } + + searchErr, ok := err.(*SearchError) + if !ok { + t.Fatalf("expected SearchError, got %T", err) + } + if searchErr.Code != "NO_API_KEY" { + t.Fatalf("expected NO_API_KEY error, got %s", searchErr.Code) + } +} + +func TestGeminiProvider_Search_CustomPromptSkipsTextGeneration(t *testing.T) { + originalText := geminiGenerateText + originalImage := geminiGenerateImage + t.Cleanup(func() { + geminiGenerateText = originalText + geminiGenerateImage = originalImage + }) + + geminiGenerateText = func(context.Context, *GeminiProvider, string, string, string, float32, int32) (string, error) { + t.Fatal("unexpected text generation for custom prompt") + return "", nil + } + + var gotPrompt string + geminiGenerateImage = func(_ context.Context, _ *GeminiProvider, prompt, _ string) ([]byte, string, error) { + gotPrompt = prompt + return mustJPEGBytes(t), "image/jpeg", nil + } + + client := NewGeminiProvider(&GeminiConfig{APIKey: "test-key"}) + callbackCalled := false + client.SetPromptCallback(func(prompt string) { + callbackCalled = true + if prompt != "custom flashcard prompt" { + t.Fatalf("callback prompt = %q, want %q", prompt, "custom flashcard prompt") + } + }) + + results, err := client.Search(context.Background(), &SearchOptions{ + Query: "ябълка", + Translation: "banana", + CustomPrompt: " custom flashcard prompt ", + }) + if err != nil { + t.Fatalf("Search() unexpected error: %v", err) + } + if gotPrompt != "custom flashcard prompt" { + t.Fatalf("image prompt = %q, want %q", gotPrompt, "custom flashcard prompt") + } + if !callbackCalled { + t.Fatal("expected prompt callback to be called") + } + if client.LastPrompt() != "custom flashcard prompt" { + t.Fatalf("LastPrompt() = %q, want %q", client.LastPrompt(), "custom flashcard prompt") + } + if len(results) != 1 { + t.Fatalf("expected 1 result, got %d", len(results)) + } + if results[0].Source != Gemini { + t.Fatalf("Source = %q, want %q", results[0].Source, Gemini) + } + if !strings.HasPrefix(results[0].URL, geminiDataPrefix) { + t.Fatalf("expected PNG data URI, got %q", results[0].URL) + } +} + +func TestGeminiProvider_Search_GeneratedPromptFlow(t *testing.T) { + originalText := geminiGenerateText + originalImage := geminiGenerateImage + originalStyles := append([]string(nil), ArtisticStyles...) + t.Cleanup(func() { + geminiGenerateText = originalText + geminiGenerateImage = originalImage + ArtisticStyles = originalStyles + }) + + ArtisticStyles = []string{"Photorealism"} + + var sceneCalls int + var gotPrompt string + var callbackPrompt string + geminiGenerateText = func(_ context.Context, _ *GeminiProvider, _, systemPrompt, userPrompt string, temperature float32, maxOutputTokens int32) (string, error) { + if strings.Contains(systemPrompt, "educational flashcards for language learning") { + sceneCalls++ + if temperature != 0.7 || maxOutputTokens != 100 { + t.Fatalf("scene params = %v/%d, want 0.7/100", temperature, maxOutputTokens) + } + if !strings.Contains(userPrompt, "apple") { + t.Fatalf("scene prompt = %q, want English translation", userPrompt) + } + return "A bright apple sits centered on a wooden table.", nil + } + t.Fatalf("unexpected system prompt: %q", systemPrompt) + return "", nil + } + + geminiGenerateImage = func(_ context.Context, _ *GeminiProvider, prompt, _ string) ([]byte, string, error) { + gotPrompt = prompt + return mustJPEGBytes(t), "image/jpeg", nil + } + + client := NewGeminiProvider(&GeminiConfig{APIKey: "test-key"}) + callbackCalled := false + client.SetPromptCallback(func(prompt string) { + callbackCalled = true + callbackPrompt = prompt + }) + + results, err := client.Search(context.Background(), &SearchOptions{ + Query: "ябълка", + Translation: "apple", + }) + if err != nil { + t.Fatalf("Search() unexpected error: %v", err) + } + if sceneCalls != 1 { + t.Fatalf("sceneCalls = %d, want 1", sceneCalls) + } + if !callbackCalled { + t.Fatal("expected prompt callback to be called") + } + if callbackPrompt != gotPrompt { + t.Fatalf("prompt callback = %q, want %q", callbackPrompt, gotPrompt) + } + if client.LastPrompt() != gotPrompt { + t.Fatalf("LastPrompt() = %q, want %q", client.LastPrompt(), gotPrompt) + } + if len(results) != 1 { + t.Fatalf("expected 1 result, got %d", len(results)) + } + + result := results[0] + if result.Source != Gemini { + t.Fatalf("Source = %q, want %q", result.Source, Gemini) + } + if result.Width != 1 || result.Height != 1 { + t.Fatalf("Size = %dx%d, want %dx%d", result.Width, result.Height, 1, 1) + } + if !strings.Contains(result.Description, "apple") { + t.Fatalf("Description = %q, want translated word", result.Description) + } + if !strings.Contains(gotPrompt, "Generate a Photorealism educational flashcard image illustrating \"apple\".") { + t.Fatalf("Prompt = %q, want translated subject in generated prompt", gotPrompt) + } + if !strings.Contains(gotPrompt, "Scene: A bright apple sits centered on a wooden table.") { + t.Fatalf("Prompt = %q, want generated scene in prompt", gotPrompt) + } + + reader, err := client.Download(context.Background(), result.URL) + if err != nil { + t.Fatalf("Download() unexpected error: %v", err) + } + t.Cleanup(func() { + _ = reader.Close() + }) + + data, err := io.ReadAll(reader) + if err != nil { + t.Fatalf("ReadAll() unexpected error: %v", err) + } + if !bytes.HasPrefix(data, []byte("\x89PNG\r\n\x1a\n")) { + t.Fatalf("Search output was not normalized to PNG") + } +} + +func TestGeminiProvider_Search_InvalidOptions(t *testing.T) { + t.Parallel() + + client := NewGeminiProvider(&GeminiConfig{APIKey: "test-key"}) + _, err := client.Search(context.Background(), nil) + if err == nil { + t.Fatal("expected error for nil options") + } + + searchErr, ok := err.(*SearchError) + if !ok { + t.Fatalf("expected SearchError, got %T", err) + } + if searchErr.Code != "INVALID_OPTIONS" { + t.Fatalf("expected INVALID_OPTIONS, got %s", searchErr.Code) + } +} + +func TestGeminiProvider_DownloadDataURI(t *testing.T) { + client := &GeminiProvider{} + payload := mustPNGBytes(t) + url := geminiDataPrefix + base64.StdEncoding.EncodeToString(payload) + + reader, err := client.Download(context.Background(), url) + if err != nil { + t.Fatalf("Download() unexpected error: %v", err) + } + t.Cleanup(func() { + _ = reader.Close() + }) + + data, err := io.ReadAll(reader) + if err != nil { + t.Fatalf("ReadAll() unexpected error: %v", err) + } + if !bytes.Equal(data, payload) { + t.Fatalf("Download() = %v, want %v", data, payload) + } +} + +func TestGeminiProvider_DownloadHTTPFallback(t *testing.T) { + client := &GeminiProvider{} + payload := []byte("fallback image bytes") + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodGet { + t.Fatalf("request method = %s, want GET", r.Method) + } + _, _ = w.Write(payload) + })) + t.Cleanup(server.Close) + + reader, err := client.Download(context.Background(), server.URL) + if err != nil { + t.Fatalf("Download() unexpected error: %v", err) + } + t.Cleanup(func() { + _ = reader.Close() + }) + + data, err := io.ReadAll(reader) + if err != nil { + t.Fatalf("ReadAll() unexpected error: %v", err) + } + if !bytes.Equal(data, payload) { + t.Fatalf("Download() = %v, want %v", data, payload) + } +} + +func TestExtractGeneratedImage(t *testing.T) { + t.Parallel() + + want := mustPNGBytes(t) + resp := &genai.GenerateContentResponse{ + Candidates: []*genai.Candidate{ + { + Content: &genai.Content{ + Parts: []*genai.Part{ + {InlineData: &genai.Blob{Data: want, MIMEType: "image/png"}}, + }, + }, + }, + }, + } + + got, mimeType, err := extractGeneratedImage(resp) + if err != nil { + t.Fatalf("extractGeneratedImage() error = %v", err) + } + if mimeType != "image/png" { + t.Fatalf("mimeType = %q, want %q", mimeType, "image/png") + } + if !bytes.Equal(got, want) { + t.Fatalf("extractGeneratedImage() = %v, want %v", got, want) + } +} + +func TestEncodeAndNormalizePNG(t *testing.T) { + t.Parallel() + + jpegBytes := mustJPEGBytes(t) + pngBytes, err := normalizePNG(jpegBytes, "image/jpeg") + if err != nil { + t.Fatalf("normalizePNG() unexpected error: %v", err) + } + if !bytes.HasPrefix(pngBytes, []byte("\x89PNG\r\n\x1a\n")) { + t.Fatalf("normalizePNG() did not return PNG data") + } + + dataURL, err := encodeDataURL(jpegBytes, "image/jpeg") + if err != nil { + t.Fatalf("encodeDataURL() unexpected error: %v", err) + } + if !strings.HasPrefix(dataURL, geminiDataPrefix) { + t.Fatalf("encodeDataURL() = %q, want PNG data URI", dataURL) + } +} + +func TestDecodeDataURIErrors(t *testing.T) { + t.Parallel() + + if _, err := decodeDataURL("not-a-data-uri"); err == nil { + t.Fatal("expected error for invalid data URI") + } +} diff --git a/internal/image/prompt.go b/internal/image/prompt.go new file mode 100644 index 0000000..addf936 --- /dev/null +++ b/internal/image/prompt.go @@ -0,0 +1,167 @@ +package image + +import ( + "fmt" + "strings" +) + +const maxImagePromptChars = 1000 + +func promptSubject(englishTranslation, fallback string) string { + subject := normalizePromptText(englishTranslation) + if subject != "" { + return subject + } + + subject = normalizePromptText(fallback) + if subject != "" { + return subject + } + + return "the requested term" +} + +func normalizePromptText(text string) string { + text = trimMarkdownFence(text) + text = strings.TrimSpace(text) + text = strings.Trim(text, "`\"'") + text = strings.Join(strings.Fields(text), " ") + return strings.TrimSpace(text) +} + +func sanitizeSceneDescription(scene string) string { + scene = normalizePromptText(scene) + lower := strings.ToLower(scene) + + for _, prefix := range []string{"scene description:", "scene:", "description:", "image prompt:", "prompt:"} { + if strings.HasPrefix(lower, prefix) { + scene = strings.TrimSpace(scene[len(prefix):]) + break + } + } + + return strings.TrimSpace(strings.Trim(scene, ".")) +} + +func usableSceneDescription(scene string) bool { + scene = sanitizeSceneDescription(scene) + if scene == "" { + return false + } + if len(scene) < 24 { + return false + } + if len(strings.Fields(scene)) < 3 { + return false + } + return true +} + +func trimMarkdownFence(text string) string { + text = strings.TrimSpace(text) + if !strings.HasPrefix(text, "```") { + return text + } + + lines := strings.Split(text, "\n") + if len(lines) < 3 { + return text + } + if !strings.HasPrefix(strings.TrimSpace(lines[0]), "```") { + return text + } + if strings.TrimSpace(lines[len(lines)-1]) != "```" { + return text + } + + return strings.Join(lines[1:len(lines)-1], "\n") +} + +func withTerminalPunctuation(text string) string { + text = strings.TrimSpace(text) + if text == "" { + return "" + } + + switch { + case strings.HasSuffix(text, "."): + return text + case strings.HasSuffix(text, "!"): + return text + case strings.HasSuffix(text, "?"): + return text + default: + return text + "." + } +} + +// buildEducationalPrompt assembles the final image-generation prompt from a +// pre-chosen artistic style, an optional scene description, and the word subject. +func buildEducationalPrompt(style, scene, subject string) string { + var prompt string + + if scene != "" { + fullPrompt := fmt.Sprintf( + "Generate a %s educational flashcard image illustrating \"%s\". Scene: %s "+ + "The image should be educational and suitable for language learning flashcards. "+ + "Requirements: The main subject or concept must be clearly visible, easily recognizable, and prominent in the image. It should occupy the central area with sharp focus and proper lighting. Ensure the scene makes \"%s\" immediately identifiable. "+ + "IMPORTANT: No text whatsoever. Do not include any words, letters, typography, labels, captions, or writing of any kind. Image only, without any text elements.", + style, subject, withTerminalPunctuation(scene), subject, + ) + + if len(fullPrompt) <= maxImagePromptChars { + prompt = fullPrompt + } else { + prompt = fmt.Sprintf( + "Generate a %s flashcard image illustrating \"%s\". Scene: %s "+ + "The image should be educational and suitable for language learning flashcards. "+ + "Requirements: The main subject or concept must be clearly visible, centered, well lit, and easy to identify.", + style, subject, withTerminalPunctuation(scene), + ) + + if len(prompt) > maxImagePromptChars { + template := fmt.Sprintf( + "Generate a %s flashcard image illustrating \"%s\". Scene: "+ + "The image should be educational and suitable for language learning flashcards. "+ + "Requirements: The main subject or concept must be clearly visible, centered, well lit, and easy to identify.", + style, subject, + ) + maxSceneLen := maxImagePromptChars - len(template) + if maxSceneLen > 3 && len(scene) > maxSceneLen { + scene = scene[:maxSceneLen] + "..." + } + prompt = fmt.Sprintf( + "Generate a %s flashcard image illustrating \"%s\". Scene: %s "+ + "The image should be educational and suitable for language learning flashcards. "+ + "Requirements: The main subject or concept must be clearly visible, centered, well lit, and easy to identify.", + style, subject, withTerminalPunctuation(scene), + ) + } + } + } else { + prompt = fmt.Sprintf( + "Generate a %s educational flashcard image illustrating \"%s\". %s "+ + "The image should be educational and suitable for language learning flashcards. "+ + "Requirements: The main subject or concept must be clearly visible, easily recognizable, and prominent in the image. Show it prominently centered with excellent lighting and sharp focus. "+ + "IMPORTANT: No text whatsoever. Do not include any words, letters, typography, labels, captions, or writing of any kind. Image only, without any text elements.", + style, subject, fallbackVisualDirection(subject), + ) + } + + if len(prompt) > maxImagePromptChars { + prompt = prompt[:997] + "..." + } + + return prompt +} + +func fallbackVisualDirection(subject string) string { + subject = normalizePromptText(subject) + lower := strings.ToLower(subject) + + if strings.HasPrefix(lower, "to ") || len(strings.Fields(subject)) > 1 { + return "Show a realistic everyday scene with people, actions, facial expressions, and surrounding objects that make the meaning of \"" + subject + "\" obvious without any text." + } + + return "Show a single " + subject + " as the clear focal point, prominently centered and immediately recognizable." +} diff --git a/internal/image/prompt_test.go b/internal/image/prompt_test.go new file mode 100644 index 0000000..73eb30e --- /dev/null +++ b/internal/image/prompt_test.go @@ -0,0 +1,60 @@ +package image + +import ( + "strings" + "testing" +) + +func TestPromptSubjectUsesTranslationFirst(t *testing.T) { + t.Parallel() + + if got := promptSubject(" apple ", "ябълка"); got != "apple" { + t.Fatalf("promptSubject() = %q, want %q", got, "apple") + } +} + +func TestPromptSubjectFallsBackToOriginalWord(t *testing.T) { + t.Parallel() + + if got := promptSubject(" ", "ябълка"); got != "ябълка" { + t.Fatalf("promptSubject() = %q, want %q", got, "ябълка") + } +} + +func TestSanitizeSceneDescriptionRemovesLabelsAndFences(t *testing.T) { + t.Parallel() + + scene := "```text\nScene: A bright apple sits centered on a wooden table.\n```" + if got := sanitizeSceneDescription(scene); got != "A bright apple sits centered on a wooden table" { + t.Fatalf("sanitizeSceneDescription() = %q", got) + } +} + +func TestUsableSceneDescriptionRejectsTrivialContent(t *testing.T) { + t.Parallel() + + if usableSceneDescription("A") { + t.Fatal("usableSceneDescription() unexpectedly accepted trivial scene") + } + if usableSceneDescription("A single, perfectly") { + t.Fatal("usableSceneDescription() unexpectedly accepted incomplete scene fragment") + } +} + +func TestFallbackVisualDirectionForPhrase(t *testing.T) { + t.Parallel() + + got := fallbackVisualDirection("to indulge someone") + if got == "" || !usableSceneDescription(got) { + t.Fatalf("fallbackVisualDirection() returned unusable phrase direction: %q", got) + } +} + +func TestFallbackVisualDirectionForSingleWord(t *testing.T) { + t.Parallel() + + got := fallbackVisualDirection("apple") + if got == "" || !strings.Contains(got, "single apple") { + t.Fatalf("fallbackVisualDirection() = %q", got) + } +} diff --git a/internal/image/registry.go b/internal/image/registry.go new file mode 100644 index 0000000..4ddb40b --- /dev/null +++ b/internal/image/registry.go @@ -0,0 +1,74 @@ +package image + +import ( + "fmt" + "sync" +) + +// Factory builds a provider instance. +type Factory func() (ImageProvider, error) + +// Config exposes the configured image provider name. +type Config interface { + ImageProviderName() string +} + +// Registry resolves provider names to factories. +type Registry struct { + mu sync.RWMutex + factories map[string]Factory +} + +// NewRegistry creates an empty provider registry. +func NewRegistry() *Registry { + return &Registry{factories: make(map[string]Factory)} +} + +// Register associates name with a factory. Later registrations replace earlier ones. +func (r *Registry) Register(name string, factory Factory) { + if r == nil || factory == nil { + return + } + + normalized := NormalizeName(name) + if normalized == "" { + return + } + + r.mu.Lock() + defer r.mu.Unlock() + if r.factories == nil { + r.factories = make(map[string]Factory) + } + r.factories[normalized] = factory +} + +// Resolve returns the factory registered for name. +func (r *Registry) Resolve(name string) (Factory, bool) { + if r == nil { + return nil, false + } + + r.mu.RLock() + defer r.mu.RUnlock() + factory, ok := r.factories[NormalizeName(name)] + return factory, ok +} + +// New constructs a provider for name. +func (r *Registry) New(name string) (ImageProvider, error) { + var zero ImageProvider + factory, ok := r.Resolve(name) + if !ok { + return zero, fmt.Errorf("%w: %s", ErrUnknownProvider, name) + } + return factory() +} + +// NewFromConfig resolves the provider name from cfg and constructs it. +func (r *Registry) NewFromConfig(cfg Config) (ImageProvider, error) { + if cfg == nil { + return nil, fmt.Errorf("image config is required") + } + return r.New(cfg.ImageProviderName()) +} diff --git a/internal/image/styles.go b/internal/image/styles.go new file mode 100644 index 0000000..387760b --- /dev/null +++ b/internal/image/styles.go @@ -0,0 +1,100 @@ +package image + +import "math/rand" + +const defaultArtisticStyle = "simple illustration" + +// ArtisticStyles contains the shared pool of artistic styles used for image prompts. +var ArtisticStyles = []string{ + "Photorealism", "Hyperrealism", "Surrealism", "Impressionism", + "Minimalism", "Pop Art", "Art Nouveau", "Digital Art", + "Watercolor", "Oil Painting", "Pencil Sketch", "Ink Drawing", + "3D Rendering", "Low Poly Art", "Pixel Art", "Vector Art", + "Collage", "Mixed Media", "Contemporary Art", "Abstract Expressionism", + "Cubism", "Pointillism", "Fauvism", "Art Deco", + "Baroque", "Renaissance", "Romanticism", "Realism", + "Post-Impressionism", "Expressionism", "Constructivism", "Suprematism", + "Dadaism", "Futurism", "Op Art", "Kinetic Art", + "Street Art", "Graffiti Art", "Installation Art", "Land Art", + "Conceptual Art", "Performance Art", "Video Art", "Net Art", + "Generative Art", "Algorithmic Art", "Fractal Art", "Glitch Art", + "Vaporwave", "Synthwave", "Cyberpunk", "Steampunk", + "Fantasy Art", "Science Fiction Art", "Horror Art", "Gothic Art", + "Anime", "Manga", "Comic Book Art", "Cartoon", + "Caricature", "Editorial Illustration", "Children's Book Illustration", "Fashion Illustration", + "Architectural Rendering", "Technical Illustration", "Scientific Illustration", "Medical Illustration", + "Botanical Illustration", "Zoological Illustration", "Astronomical Art", "Paleoart", + "Infographic", "Data Visualization", "Typography Art", "Calligraphy", + "Mosaic", "Stained Glass", "Tapestry", "Embroidery", + "Sculpture", "Ceramics", "Pottery", "Glass Art", + "Metalwork", "Jewelry Design", "Woodcarving", "Paper Art", + "Origami", "Kirigami", "Quilling", "Book Art", + "Photography", "Documentary Photography", "Portrait Photography", "Landscape Photography", + "Macro Photography", "Aerial Photography", "Underwater Photography", "Astrophotography", + "Film Noir", "Vintage Photography", "Polaroid", "Double Exposure", + "HDR Photography", "Long Exposure", "Tilt-Shift", "Infrared Photography", + "Monochrome", "Sepia Tone", "Cross-Processing", "Cyanotype", + "Folk Art", "Outsider Art", "Naive Art", "Aboriginal Art", + "African Art", "Asian Art", "Islamic Art", "Celtic Art", + "Byzantine Art", "Medieval Art", "Pre-Columbian Art", "Ancient Egyptian Art", + "Ancient Greek Art", "Ancient Roman Art", "Cave Painting", "Petroglyphs", + "Bauhaus", "De Stijl", "Vienna Secession", "Arts and Crafts Movement", + "Prairie School", "International Style", "Brutalism", "Deconstructivism", + "Parametric Design", "Biomimicry", "Sustainable Design", "Universal Design", + "Retro Futurism", "Dieselpunk", "Atompunk", "Biopunk", + "Afrofuturism", "Solarpunk", "Post-Apocalyptic", "Dystopian Art", + "Psychedelic Art", "Visionary Art", "Lowbrow Art", "Outsider Art", + "Trompe-l'oeil", "Anamorphic Art", "Optical Illusion", "Impossible Objects", + "Sacred Geometry", "Mandala", "Yantra", "Celtic Knots", + "Stippling", "Hatching", "Cross-Hatching", "Scumbling", + "Impasto", "Glazing", "Scumbling", "Sgraffito", + "Encaustic", "Fresco", "Tempera", "Gouache", + "Pastel", "Charcoal", "Conte", "Silverpoint", + "Linocut", "Woodcut", "Etching", "Lithography", + "Screen Printing", "Monotype", "Collagraph", "Digital Print", + "Augmented Reality Art", "Virtual Reality Art", "Interactive Art", "Projection Mapping", + "Light Art", "Neon Art", "Holographic Art", "Laser Art", + "Sound Art", "Bio Art", "Eco Art", "Social Practice Art", + "Relational Aesthetics", "Participatory Art", "Community Art", "Activist Art", + "Feminist Art", "Queer Art", "Postcolonial Art", "Decolonial Art", + "Metamodernism", "Post-Internet Art", "Post-Digital Art", "New Aesthetic", + "Speculative Design", "Critical Design", "Design Fiction", "Adversarial Design", + "Transitional Design", "Transformation Design", "Service Design", "Experience Design", + "Slow Design", "Emotional Design", "Inclusive Design", "Regenerative Design", + "Biophilic Design", "Cradle to Cradle", "Circular Design", "Zero Waste Design", + "Modular Design", "Open Design", "Co-Design", "Participatory Design", + "Flat Design", "Material Design", "Neumorphism", "Glassmorphism", + "Maximalism", "Eclecticism", "Kitsch", "Camp", + "Wabi-Sabi", "Hygge", "Lagom", "Ikigai", + "Feng Shui", "Vastu Shastra", "Sacred Architecture", "Organic Architecture", + "Vernacular Architecture", "Adaptive Reuse", "Green Architecture", "Living Architecture", + "Kinetic Architecture", "Responsive Architecture", "Parametric Architecture", "Algorithmic Architecture", + "Blob Architecture", "Deconstructivist Architecture", "High-Tech Architecture", "Neo-Futurism", + "Critical Regionalism", "Tropical Modernism", "Desert Modernism", "Scandinavian Design", + "Japanese Design", "Italian Design", "German Design", "Dutch Design", + "Memphis Group", "Radical Design", "Anti-Design", "Superstudio", + "Archigram", "Metabolism", "Structuralism", "Postmodernism", + "Minimalist Photography", "Conceptual Photography", "Staged Photography", "Candid Photography", +} + +func pickArtisticStyle() string { + if len(ArtisticStyles) == 0 { + return "" + } + + styles := append([]string(nil), ArtisticStyles...) + rand.Shuffle(len(styles), func(i, j int) { + styles[i], styles[j] = styles[j], styles[i] + }) + + return styles[0] +} + +func chooseArtisticStyle() string { + style := pickArtisticStyle() + if style == "" { + return defaultArtisticStyle + } + + return style +} diff --git a/internal/image/styles_test.go b/internal/image/styles_test.go new file mode 100644 index 0000000..4e189c1 --- /dev/null +++ b/internal/image/styles_test.go @@ -0,0 +1,49 @@ +package image + +import ( + "reflect" + "testing" +) + +func TestArtisticStyles(t *testing.T) { + if len(ArtisticStyles) == 0 { + t.Fatal("ArtisticStyles should not be empty") + } + if !hasStyle(ArtisticStyles, "Photorealism") { + t.Fatal(`ArtisticStyles should include "Photorealism"`) + } + if !hasStyle(ArtisticStyles, "Candid Photography") { + t.Fatal(`ArtisticStyles should include "Candid Photography"`) + } +} + +func TestChooseArtisticStyle_EmptyPool(t *testing.T) { + original := ArtisticStyles + ArtisticStyles = nil + t.Cleanup(func() { ArtisticStyles = original }) + + if got := chooseArtisticStyle(); got != defaultArtisticStyle { + t.Fatalf("chooseArtisticStyle() = %q, want %q", got, defaultArtisticStyle) + } +} + +func TestPickArtisticStyle_DoesNotMutateSharedPool(t *testing.T) { + original := append([]string(nil), ArtisticStyles...) + t.Cleanup(func() { ArtisticStyles = original }) + + ArtisticStyles = []string{"Photorealism", "Surrealism", "Impressionism"} + _ = pickArtisticStyle() + + if !reflect.DeepEqual(ArtisticStyles, []string{"Photorealism", "Surrealism", "Impressionism"}) { + t.Fatalf("pickArtisticStyle() mutated shared pool: got %v", ArtisticStyles) + } +} + +func hasStyle(styles []string, want string) bool { + for _, style := range styles { + if style == want { + return true + } + } + return false +} diff --git a/internal/image/test_helpers_test.go b/internal/image/test_helpers_test.go new file mode 100644 index 0000000..31bf30c --- /dev/null +++ b/internal/image/test_helpers_test.go @@ -0,0 +1,38 @@ +package image + +import ( + "bytes" + "image" + "image/color" + "image/jpeg" + "image/png" + "testing" +) + +func mustPNGBytes(t *testing.T) []byte { + t.Helper() + + img := image.NewRGBA(image.Rect(0, 0, 1, 1)) + img.Set(0, 0, color.RGBA{R: 255, A: 255}) + + var buf bytes.Buffer + if err := png.Encode(&buf, img); err != nil { + t.Fatalf("png.Encode() error = %v", err) + } + + return buf.Bytes() +} + +func mustJPEGBytes(t *testing.T) []byte { + t.Helper() + + img := image.NewRGBA(image.Rect(0, 0, 1, 1)) + img.Set(0, 0, color.RGBA{G: 255, A: 255}) + + var buf bytes.Buffer + if err := jpeg.Encode(&buf, img, &jpeg.Options{Quality: 90}); err != nil { + t.Fatalf("jpeg.Encode() error = %v", err) + } + + return buf.Bytes() +} diff --git a/internal/image/types.go b/internal/image/types.go new file mode 100644 index 0000000..e839308 --- /dev/null +++ b/internal/image/types.go @@ -0,0 +1,113 @@ +package image + +import ( + "context" + "errors" + "io" + "strings" +) + +const ( + // Gemini is the canonical provider name for Google's Gemini backend. + Gemini = "gemini" + + // OpenAI is the canonical provider name for OpenAI backends. + OpenAI = "openai" +) + +const ( + defaultSearchLanguage = "bg" + defaultSearchPerPage = 10 + defaultSearchPage = 1 + defaultSearchImageType = "photo" + defaultSearchOrientation = "all" +) + +// SearchResult represents a single image result. +type SearchResult struct { + ID string + URL string + ThumbnailURL string + Width int + Height int + Description string + Attribution string + Source string +} + +// SearchOptions configures image search and generation. +type SearchOptions struct { + Query string + Translation string + Language string + SafeSearch bool + PerPage int + Page int + ImageType string + Orientation string + CustomPrompt string + AspectRatio string + ReferenceImages [][]byte +} + +// DefaultSearchOptions returns sensible defaults for language-learning queries. +func DefaultSearchOptions(query string) *SearchOptions { + return &SearchOptions{ + Query: query, + Language: defaultSearchLanguage, + SafeSearch: true, + PerPage: defaultSearchPerPage, + Page: defaultSearchPage, + ImageType: defaultSearchImageType, + Orientation: defaultSearchOrientation, + } +} + +// ImageProvider generates, downloads, and describes images. +type ImageProvider interface { + Name() string + Search(ctx context.Context, opts *SearchOptions) ([]SearchResult, error) + Download(ctx context.Context, url string) (io.ReadCloser, error) + GetAttribution(result *SearchResult) string +} + +// ImageClient is kept as a compatibility alias for ImageProvider. +type ImageClient = ImageProvider + +// SearchError represents a provider failure. +type SearchError struct { + Provider string + Code string + Message string +} + +func (e *SearchError) Error() string { + switch { + case e == nil: + return "" + case e.Provider == "": + return e.Message + case e.Message == "": + return e.Provider + default: + return e.Provider + ": " + e.Message + } +} + +// ErrUnknownProvider indicates that no provider is registered for a name. +var ErrUnknownProvider = errors.New("unknown image provider") + +// NormalizeName returns a canonical lower-case provider name. +func NormalizeName(name string) string { + return strings.ToLower(strings.TrimSpace(name)) +} + +// IsKnownName reports whether the name matches a supported provider family. +func IsKnownName(name string) bool { + switch NormalizeName(name) { + case Gemini, OpenAI: + return true + default: + return false + } +} diff --git a/internal/image/types_test.go b/internal/image/types_test.go new file mode 100644 index 0000000..c4e04f9 --- /dev/null +++ b/internal/image/types_test.go @@ -0,0 +1,101 @@ +package image + +import ( + "context" + "errors" + "io" + "strings" + "testing" +) + +func TestNormalizeName(t *testing.T) { + t.Parallel() + + if got, want := NormalizeName(" Gemini "), Gemini; got != want { + t.Fatalf("NormalizeName() = %q, want %q", got, want) + } +} + +func TestIsKnownName(t *testing.T) { + t.Parallel() + + if !IsKnownName("gemini") { + t.Fatal("expected gemini to be recognized") + } + if IsKnownName("bogus") { + t.Fatal("expected bogus provider to be rejected") + } +} + +func TestDefaultSearchOptions(t *testing.T) { + t.Parallel() + + opts := DefaultSearchOptions("ябълка") + if opts.Query != "ябълка" { + t.Fatalf("Query = %q, want %q", opts.Query, "ябълка") + } + if opts.Language != defaultSearchLanguage || !opts.SafeSearch || opts.PerPage != defaultSearchPerPage || opts.Page != defaultSearchPage { + t.Fatalf("DefaultSearchOptions() = %+v", opts) + } +} + +func TestSearchError(t *testing.T) { + t.Parallel() + + err := &SearchError{Provider: "test", Code: "404", Message: "Not found"} + if got, want := err.Error(), "test: Not found"; got != want { + t.Fatalf("SearchError.Error() = %q, want %q", got, want) + } +} + +func TestRegistryNewFromConfig(t *testing.T) { + t.Parallel() + + registry := NewRegistry() + registry.Register(Gemini, func() (ImageProvider, error) { + return fakeProvider{name: Gemini}, nil + }) + + provider, err := registry.NewFromConfig(fakeConfig{name: Gemini}) + if err != nil { + t.Fatalf("NewFromConfig() error = %v", err) + } + if got, want := provider.Name(), Gemini; got != want { + t.Fatalf("provider.Name() = %q, want %q", got, want) + } +} + +func TestRegistryUnknownProvider(t *testing.T) { + t.Parallel() + + registry := NewRegistry() + _, err := registry.New("missing") + if err == nil { + t.Fatal("expected unknown provider error") + } + if !errors.Is(err, ErrUnknownProvider) { + t.Fatalf("error = %v, want ErrUnknownProvider", err) + } +} + +type fakeConfig struct { + name string +} + +func (f fakeConfig) ImageProviderName() string { return f.name } + +type fakeProvider struct { + name string +} + +func (f fakeProvider) Name() string { return f.name } + +func (fakeProvider) Search(context.Context, *SearchOptions) ([]SearchResult, error) { + return nil, nil +} + +func (fakeProvider) Download(context.Context, string) (io.ReadCloser, error) { + return io.NopCloser(strings.NewReader("")), nil +} + +func (fakeProvider) GetAttribution(*SearchResult) string { return "" } -- cgit v1.2.3