From b006582e4c7323821eb836080a9bc9ec24243a47 Mon Sep 17 00:00:00 2001 From: Paul Buetow Date: Wed, 1 Apr 2026 13:00:40 +0300 Subject: zr: add Gemini-backed translation provider --- internal/cli/command.go | 11 ++ internal/cli/command_test.go | 65 ++++++++++ internal/processor/processor.go | 7 +- internal/processor/processor_test.go | 21 ++-- internal/translation/doc.go | 6 +- internal/translation/translator.go | 213 +++++++++++++++++++++++++------- internal/translation/translator_test.go | 129 +++++++++---------- 7 files changed, 323 insertions(+), 129 deletions(-) (limited to 'internal') diff --git a/internal/cli/command.go b/internal/cli/command.go index b07bd26..a5a8c13 100644 --- a/internal/cli/command.go +++ b/internal/cli/command.go @@ -155,3 +155,14 @@ func GetOpenAIKey() string { // Then check config file return viper.GetString("audio.openai_key") } + +// GetGoogleAPIKey retrieves the Google API key from environment or config. +func GetGoogleAPIKey() string { + // First check environment variable + if key := os.Getenv("GOOGLE_API_KEY"); key != "" { + return key + } + + // Then check config file + return viper.GetString("google.api_key") +} diff --git a/internal/cli/command_test.go b/internal/cli/command_test.go index c6fa3d8..2361cd2 100644 --- a/internal/cli/command_test.go +++ b/internal/cli/command_test.go @@ -237,6 +237,71 @@ func TestGetOpenAIKey(t *testing.T) { } } +func TestGetGoogleAPIKey(t *testing.T) { + // Save original viper state + originalConfig := viper.New() + *originalConfig = *viper.GetViper() + defer func() { + *viper.GetViper() = *originalConfig + }() + + tests := []struct { + name string + envKey string + configKey string + expected string + }{ + { + name: "from environment", + envKey: "env-google-key", + configKey: "config-google-key", + expected: "env-google-key", + }, + { + name: "from config when no env", + envKey: "", + configKey: "config-google-key", + expected: "config-google-key", + }, + { + name: "empty when neither set", + envKey: "", + configKey: "", + expected: "", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + viper.Reset() + + if tt.envKey != "" { + if err := os.Setenv("GOOGLE_API_KEY", tt.envKey); err != nil { + t.Fatalf("Failed to set GOOGLE_API_KEY: %v", err) + } + defer func() { + if err := os.Unsetenv("GOOGLE_API_KEY"); err != nil { + t.Errorf("Failed to unset GOOGLE_API_KEY: %v", err) + } + }() + } else { + if err := os.Unsetenv("GOOGLE_API_KEY"); err != nil { + t.Fatalf("Failed to unset GOOGLE_API_KEY: %v", err) + } + } + + if tt.configKey != "" { + viper.Set("google.api_key", tt.configKey) + } + + got := GetGoogleAPIKey() + if got != tt.expected { + t.Errorf("GetGoogleAPIKey() = %v, want %v", got, tt.expected) + } + }) + } +} + func TestBindFlagsToViper(t *testing.T) { // Save original viper state originalConfig := viper.New() diff --git a/internal/processor/processor.go b/internal/processor/processor.go index d52ca53..94b0448 100644 --- a/internal/processor/processor.go +++ b/internal/processor/processor.go @@ -32,12 +32,13 @@ type Processor struct { // NewProcessor creates a new word processor func NewProcessor(flags *cli.Flags) *Processor { - apiKey := cli.GetOpenAIKey() + openAIKey := cli.GetOpenAIKey() + translationProvider := translation.Provider(viper.GetString("translation.provider")) return &Processor{ flags: flags, - translator: translation.NewTranslator(apiKey), + translator: translation.NewTranslator(&translation.Config{Provider: translationProvider, OpenAIKey: openAIKey, GoogleAPIKey: cli.GetGoogleAPIKey()}), translationCache: translation.NewTranslationCache(), - phoneticFetcher: phonetic.NewFetcher(apiKey), + phoneticFetcher: phonetic.NewFetcher(openAIKey), } } diff --git a/internal/processor/processor_test.go b/internal/processor/processor_test.go index e835681..a9c9ab1 100644 --- a/internal/processor/processor_test.go +++ b/internal/processor/processor_test.go @@ -9,15 +9,8 @@ import ( ) func TestNewProcessor(t *testing.T) { - // Set up test environment - if err := os.Setenv("OPENAI_API_KEY", "test-key"); err != nil { - t.Fatalf("Failed to set OPENAI_API_KEY: %v", err) - } - defer func() { - if err := os.Unsetenv("OPENAI_API_KEY"); err != nil { - t.Errorf("Failed to unset OPENAI_API_KEY: %v", err) - } - }() + t.Setenv("OPENAI_API_KEY", "test-openai-key") + t.Setenv("GOOGLE_API_KEY", "test-google-key") flags := cli.NewFlags() p := NewProcessor(flags) @@ -63,8 +56,8 @@ func TestProcessSingleWord_InvalidWord(t *testing.T) { func TestProcessSingleWord_ValidWord(t *testing.T) { // Skip if no API key - if os.Getenv("OPENAI_API_KEY") == "" { - t.Skip("Skipping test: OPENAI_API_KEY not set") + if os.Getenv("OPENAI_API_KEY") == "" || os.Getenv("GOOGLE_API_KEY") == "" { + t.Skip("Skipping test: OPENAI_API_KEY and GOOGLE_API_KEY must be set") } flags := cli.NewFlags() @@ -114,9 +107,9 @@ func TestProcessBatch_ValidFile(t *testing.T) { flags.SkipImages = true p := NewProcessor(flags) - // Skip if no API key - if os.Getenv("OPENAI_API_KEY") == "" { - t.Skip("Skipping test: OPENAI_API_KEY not set") + // Skip if no API keys + if os.Getenv("OPENAI_API_KEY") == "" || os.Getenv("GOOGLE_API_KEY") == "" { + t.Skip("Skipping test: OPENAI_API_KEY and GOOGLE_API_KEY must be set") } err = p.ProcessBatch() diff --git a/internal/translation/doc.go b/internal/translation/doc.go index fac31ff..73dc73d 100644 --- a/internal/translation/doc.go +++ b/internal/translation/doc.go @@ -1,4 +1,4 @@ -// Package translation provides Bulgarian to English translation services -// using the OpenAI API. It includes translation caching for batch operations -// and file persistence for translated words. +// Package translation provides provider-aware Bulgarian and English translation +// services using OpenAI or Gemini. It includes translation caching for batch +// operations and file persistence for translated words. package translation diff --git a/internal/translation/translator.go b/internal/translation/translator.go index 5f63771..adcf6ae 100644 --- a/internal/translation/translator.go +++ b/internal/translation/translator.go @@ -6,100 +6,196 @@ import ( "os" "path/filepath" "strings" + "time" "github.com/sashabaranov/go-openai" + "google.golang.org/genai" ) -// Translator handles Bulgarian to English translation +const ( + // ProviderGemini routes translation requests to Gemini. + ProviderGemini Provider = "gemini" + // ProviderOpenAI routes translation requests to OpenAI. + ProviderOpenAI Provider = "openai" + + defaultGeminiModel = "gemini-2.5-flash" + translationTimeout = 30 * time.Second + translationMaxTokens = 50 + translationTemperature = 0.3 +) + +// Provider selects the translation backend. +type Provider string + +// Config holds translator settings and API credentials. +type Config struct { + Provider Provider + OpenAIKey string + GoogleAPIKey string + OpenAIModel string + GeminiModel string +} + +// DefaultConfig returns a translator configuration with Gemini as the default backend. +func DefaultConfig() *Config { + return &Config{ + Provider: ProviderGemini, + OpenAIModel: openai.GPT4oMini, + GeminiModel: defaultGeminiModel, + } +} + +// Translator handles Bulgarian and English translation using the configured backend. type Translator struct { - apiKey string - client *openai.Client + provider Provider + openAIKey string + googleAPIKey string + openAIClient *openai.Client + geminiClient *genai.Client + openAIModel string + geminiModel string } -// NewTranslator creates a new translator instance -func NewTranslator(apiKey string) *Translator { - return &Translator{ - apiKey: apiKey, - client: openai.NewClient(apiKey), +// NewTranslator creates a new translator instance from the provided config. +func NewTranslator(config *Config) *Translator { + if config == nil { + config = DefaultConfig() + } + + normalized := normalizeConfig(config) + translator := &Translator{ + provider: normalized.Provider, + openAIKey: normalized.OpenAIKey, + googleAPIKey: normalized.GoogleAPIKey, + openAIModel: normalized.OpenAIModel, + geminiModel: normalized.GeminiModel, } + + if normalized.OpenAIKey != "" { + translator.openAIClient = openai.NewClient(normalized.OpenAIKey) + } + + if normalized.GoogleAPIKey != "" { + client, err := genai.NewClient(context.Background(), &genai.ClientConfig{ + APIKey: normalized.GoogleAPIKey, + }) + if err == nil { + translator.geminiClient = client + } + } + + return translator } // TranslateWord translates a Bulgarian word to English. // // Example: // -// translator := translation.NewTranslator(os.Getenv("OPENAI_API_KEY")) +// translator := translation.NewTranslator(&translation.Config{ +// Provider: translation.ProviderGemini, +// GoogleAPIKey: os.Getenv("GOOGLE_API_KEY"), +// }) // english, err := translator.TranslateWord("ябълка") // if err != nil { // log.Fatal(err) // } // fmt.Println(english) func (t *Translator) TranslateWord(word string) (string, error) { - if t.apiKey == "" { + return t.translate(fmt.Sprintf( + "Translate the Bulgarian word '%s' to English. Respond with only the English translation, nothing else.", + word, + )) +} + +// TranslateEnglishToBulgarian translates an English word to Bulgarian. +func (t *Translator) TranslateEnglishToBulgarian(word string) (string, error) { + return t.translate(fmt.Sprintf( + "Translate the English word '%s' to Bulgarian. Respond with only the Bulgarian translation in Cyrillic script, nothing else.", + word, + )) +} + +func (t *Translator) translate(prompt string) (string, error) { + switch normalizeProvider(t.provider) { + case ProviderGemini: + return t.translateWithGemini(prompt) + case ProviderOpenAI: + return t.translateWithOpenAI(prompt) + default: + return "", fmt.Errorf("unknown translation provider: %s", t.provider) + } +} + +func (t *Translator) translateWithOpenAI(prompt string) (string, error) { + if t.openAIKey == "" { return "", fmt.Errorf("OpenAI API key not found") } + if t.openAIClient == nil { + return "", fmt.Errorf("OpenAI client not initialized") + } - ctx := context.Background() + ctx, cancel := context.WithTimeout(context.Background(), translationTimeout) + defer cancel() req := openai.ChatCompletionRequest{ - Model: openai.GPT4oMini, + Model: t.openAIModel, Messages: []openai.ChatCompletionMessage{ { Role: openai.ChatMessageRoleUser, - Content: fmt.Sprintf("Translate the Bulgarian word '%s' to English. Respond with only the English translation, nothing else.", word), + Content: prompt, }, }, - MaxTokens: 50, - Temperature: 0.3, + MaxTokens: translationMaxTokens, + Temperature: translationTemperature, } - resp, err := t.client.CreateChatCompletion(ctx, req) + resp, err := t.openAIClient.CreateChatCompletion(ctx, req) if err != nil { return "", fmt.Errorf("OpenAI API error: %w", err) } - if len(resp.Choices) == 0 { return "", fmt.Errorf("no translation returned") } translation := strings.TrimSpace(resp.Choices[0].Message.Content) + if translation == "" { + return "", fmt.Errorf("no translation returned") + } + return translation, nil } -// TranslateEnglishToBulgarian translates an English word to Bulgarian -func (t *Translator) TranslateEnglishToBulgarian(word string) (string, error) { - if t.apiKey == "" { - return "", fmt.Errorf("OpenAI API key not found") +func (t *Translator) translateWithGemini(prompt string) (string, error) { + if t.googleAPIKey == "" { + return "", fmt.Errorf("Google API key not found") } - - ctx := context.Background() - - req := openai.ChatCompletionRequest{ - Model: openai.GPT4oMini, - Messages: []openai.ChatCompletionMessage{ - { - Role: openai.ChatMessageRoleUser, - Content: fmt.Sprintf("Translate the English word '%s' to Bulgarian. Respond with only the Bulgarian translation in Cyrillic script, nothing else.", word), - }, - }, - MaxTokens: 50, - Temperature: 0.3, + if t.geminiClient == nil { + return "", fmt.Errorf("Gemini client not initialized") } - resp, err := t.client.CreateChatCompletion(ctx, req) + ctx, cancel := context.WithTimeout(context.Background(), translationTimeout) + defer cancel() + + temp := float32(translationTemperature) + resp, err := t.geminiClient.Models.GenerateContent(ctx, t.geminiModel, []*genai.Content{ + genai.NewContentFromText(prompt, genai.RoleUser), + }, &genai.GenerateContentConfig{ + Temperature: &temp, + MaxOutputTokens: translationMaxTokens, + }) if err != nil { - return "", fmt.Errorf("OpenAI API error: %w", err) + return "", fmt.Errorf("Gemini API error: %w", err) } - if len(resp.Choices) == 0 { + translation := strings.TrimSpace(resp.Text()) + if translation == "" { return "", fmt.Errorf("no translation returned") } - translation := strings.TrimSpace(resp.Choices[0].Message.Content) return translation, nil } -// SaveTranslation saves the translation to a file in the word directory +// SaveTranslation saves the translation to a file in the word directory. func SaveTranslation(wordDir, word, translation string) error { outputFile := filepath.Join(wordDir, "translation.txt") content := fmt.Sprintf("%s = %s\n", word, translation) @@ -111,35 +207,62 @@ func SaveTranslation(wordDir, word, translation string) error { return nil } -// TranslationCache stores translations in memory for batch operations +// TranslationCache stores translations in memory for batch operations. type TranslationCache struct { translations map[string]string } -// NewTranslationCache creates a new translation cache +// NewTranslationCache creates a new translation cache. func NewTranslationCache() *TranslationCache { return &TranslationCache{ translations: make(map[string]string), } } -// Add adds a translation to the cache +// Add adds a translation to the cache. func (tc *TranslationCache) Add(word, translation string) { tc.translations[word] = translation } -// Get retrieves a translation from the cache +// Get retrieves a translation from the cache. func (tc *TranslationCache) Get(word string) (string, bool) { translation, ok := tc.translations[word] return translation, ok } -// GetAll returns all cached translations +// GetAll returns all cached translations. func (tc *TranslationCache) GetAll() map[string]string { - // Return a copy to prevent external modification result := make(map[string]string) for k, v := range tc.translations { result[k] = v } return result } + +func normalizeConfig(config *Config) *Config { + normalized := *config + normalized.Provider = normalizeProvider(normalized.Provider) + + if normalized.Provider == "" { + normalized.Provider = ProviderGemini + } + if normalized.OpenAIModel == "" { + normalized.OpenAIModel = openai.GPT4oMini + } + if normalized.GeminiModel == "" { + normalized.GeminiModel = defaultGeminiModel + } + + return &normalized +} + +func normalizeProvider(provider Provider) Provider { + switch strings.ToLower(strings.TrimSpace(string(provider))) { + case "", string(ProviderGemini): + return ProviderGemini + case string(ProviderOpenAI): + return ProviderOpenAI + default: + return Provider(strings.ToLower(strings.TrimSpace(string(provider)))) + } +} diff --git a/internal/translation/translator_test.go b/internal/translation/translator_test.go index c87ad9c..1011bfa 100644 --- a/internal/translation/translator_test.go +++ b/internal/translation/translator_test.go @@ -7,94 +7,105 @@ import ( "testing" ) -func TestNewTranslator(t *testing.T) { - translator := NewTranslator("test-api-key") +func TestNewTranslator_DefaultsToGemini(t *testing.T) { + translator := NewTranslator(nil) if translator == nil { t.Fatal("NewTranslator returned nil") } - - if translator.apiKey != "test-api-key" { - t.Errorf("Expected API key 'test-api-key', got '%s'", translator.apiKey) + if translator.provider != ProviderGemini { + t.Fatalf("Expected default provider %q, got %q", ProviderGemini, translator.provider) + } + if translator.geminiModel != defaultGeminiModel { + t.Fatalf("Expected Gemini model %q, got %q", defaultGeminiModel, translator.geminiModel) + } + if translator.openAIModel != "gpt-4o-mini" { + t.Fatalf("Expected OpenAI model gpt-4o-mini, got %q", translator.openAIModel) } +} + +func TestNewTranslator_OpenAIProvider(t *testing.T) { + translator := NewTranslator(&Config{ + Provider: ProviderOpenAI, + OpenAIKey: "test-api-key", + }) - if translator.client == nil { - t.Error("OpenAI client not initialized") + if translator == nil { + t.Fatal("NewTranslator returned nil") + } + if translator.provider != ProviderOpenAI { + t.Fatalf("Expected provider %q, got %q", ProviderOpenAI, translator.provider) + } + if translator.openAIClient == nil { + t.Fatal("OpenAI client not initialized") } } -func TestTranslateWord_NoAPIKey(t *testing.T) { - translator := NewTranslator("") +func TestTranslateWord_NoGoogleAPIKey(t *testing.T) { + translator := NewTranslator(&Config{Provider: ProviderGemini}) _, err := translator.TranslateWord("ябълка") if err == nil { - t.Error("Expected error for missing API key") + t.Fatal("Expected error for missing Google API key") } - - if err.Error() != "OpenAI API key not found" { - t.Errorf("Expected 'OpenAI API key not found' error, got: %v", err) + if err.Error() != "Google API key not found" { + t.Fatalf("Expected 'Google API key not found' error, got: %v", err) } } -func TestTranslateWord_Integration(t *testing.T) { - // Skip if no API key - apiKey := os.Getenv("OPENAI_API_KEY") +func TestTranslateWord_IntegrationGemini(t *testing.T) { + apiKey := os.Getenv("GOOGLE_API_KEY") if apiKey == "" { - t.Skip("Skipping integration test: OPENAI_API_KEY not set") + t.Skip("Skipping integration test: GOOGLE_API_KEY not set") } - translator := NewTranslator(apiKey) + translator := NewTranslator(&Config{ + Provider: ProviderGemini, + GoogleAPIKey: apiKey, + }) - // Test with a simple word translation, err := translator.TranslateWord("ябълка") if err != nil { - t.Errorf("TranslateWord failed: %v", err) + t.Fatalf("TranslateWord failed: %v", err) } - - // Check that we got a reasonable translation - // The exact translation might vary, but it should contain "apple" if translation == "" { - t.Error("Got empty translation") + t.Fatal("Got empty translation") } t.Logf("Translation of 'ябълка': %s", translation) } -func TestTranslateEnglishToBulgarian_NoAPIKey(t *testing.T) { - translator := NewTranslator("") +func TestTranslateEnglishToBulgarian_NoOpenAIKey(t *testing.T) { + translator := NewTranslator(&Config{Provider: ProviderOpenAI}) _, err := translator.TranslateEnglishToBulgarian("apple") if err == nil { - t.Error("Expected error for missing API key") + t.Fatal("Expected error for missing OpenAI API key") } - if err.Error() != "OpenAI API key not found" { - t.Errorf("Expected 'OpenAI API key not found' error, got: %v", err) + t.Fatalf("Expected 'OpenAI API key not found' error, got: %v", err) } } -func TestTranslateEnglishToBulgarian_Integration(t *testing.T) { - // Skip if no API key +func TestTranslateEnglishToBulgarian_IntegrationOpenAI(t *testing.T) { apiKey := os.Getenv("OPENAI_API_KEY") if apiKey == "" { t.Skip("Skipping integration test: OPENAI_API_KEY not set") } - translator := NewTranslator(apiKey) + translator := NewTranslator(&Config{ + Provider: ProviderOpenAI, + OpenAIKey: apiKey, + }) - // Test with a simple word translation, err := translator.TranslateEnglishToBulgarian("apple") if err != nil { - t.Errorf("TranslateEnglishToBulgarian failed: %v", err) + t.Fatalf("TranslateEnglishToBulgarian failed: %v", err) } - - // Check that we got a reasonable translation - // The exact translation might vary, but it should be in Cyrillic if translation == "" { - t.Error("Got empty translation") + t.Fatal("Got empty translation") } - // Check that the result contains Cyrillic characters hasCyrillic := false for _, r := range translation { if r >= 'А' && r <= 'я' { @@ -103,7 +114,7 @@ func TestTranslateEnglishToBulgarian_Integration(t *testing.T) { } } if !hasCyrillic { - t.Errorf("Expected Cyrillic translation, got: %s", translation) + t.Fatalf("Expected Cyrillic translation, got: %s", translation) } t.Logf("Translation of 'apple': %s", translation) @@ -112,70 +123,61 @@ func TestTranslateEnglishToBulgarian_Integration(t *testing.T) { func TestSaveTranslation(t *testing.T) { tmpDir := t.TempDir() - err := SaveTranslation(tmpDir, "ябълка", "apple") - if err != nil { - t.Errorf("SaveTranslation failed: %v", err) + if err := SaveTranslation(tmpDir, "ябълка", "apple"); err != nil { + t.Fatalf("SaveTranslation failed: %v", err) } - // Check file was created translationFile := filepath.Join(tmpDir, "translation.txt") content, err := os.ReadFile(translationFile) if err != nil { - t.Errorf("Failed to read translation file: %v", err) + t.Fatalf("Failed to read translation file: %v", err) } expected := "ябълка = apple\n" if string(content) != expected { - t.Errorf("Expected content '%s', got '%s'", expected, string(content)) + t.Fatalf("Expected content %q, got %q", expected, string(content)) } } func TestSaveTranslation_InvalidPath(t *testing.T) { - err := SaveTranslation("/nonexistent/path", "ябълка", "apple") - if err == nil { - t.Error("Expected error for invalid path") + if err := SaveTranslation("/nonexistent/path", "ябълка", "apple"); err == nil { + t.Fatal("Expected error for invalid path") } } func TestTranslationCache(t *testing.T) { cache := NewTranslationCache() - // Test empty cache - _, found := cache.Get("ябълка") - if found { - t.Error("Expected not found in empty cache") + if _, found := cache.Get("ябълка"); found { + t.Fatal("Expected not found in empty cache") } - // Test adding and retrieving cache.Add("ябълка", "apple") cache.Add("котка", "cat") translation, found := cache.Get("ябълка") if !found { - t.Error("Expected to find 'ябълка' in cache") + t.Fatal("Expected to find 'ябълка' in cache") } if translation != "apple" { - t.Errorf("Expected 'apple', got '%s'", translation) + t.Fatalf("Expected 'apple', got %q", translation) } - // Test overwriting cache.Add("ябълка", "apple (fruit)") translation, found = cache.Get("ябълка") if !found || translation != "apple (fruit)" { - t.Errorf("Expected 'apple (fruit)', got '%s'", translation) + t.Fatalf("Expected 'apple (fruit)', got %q", translation) } } func TestTranslationCache_GetAll(t *testing.T) { cache := NewTranslationCache() - // Add some translations cache.Add("ябълка", "apple") cache.Add("котка", "cat") cache.Add("куче", "dog") all := cache.GetAll() - expected := map[string]string{ "ябълка": "apple", "котка": "cat", @@ -183,15 +185,14 @@ func TestTranslationCache_GetAll(t *testing.T) { } if !reflect.DeepEqual(all, expected) { - t.Errorf("GetAll() = %v, want %v", all, expected) + t.Fatalf("GetAll() = %v, want %v", all, expected) } - // Test that modifying returned map doesn't affect cache all["ябълка"] = "modified" translation, _ := cache.Get("ябълка") if translation != "apple" { - t.Error("Cache was modified through returned map") + t.Fatal("Cache was modified through returned map") } } @@ -200,6 +201,6 @@ func TestTranslationCache_EmptyCache(t *testing.T) { all := cache.GetAll() if len(all) != 0 { - t.Errorf("Expected empty map, got %v", all) + t.Fatalf("Expected empty map, got %v", all) } } -- cgit v1.2.3