diff options
| author | Paul Buetow <paul@buetow.org> | 2026-04-02 09:10:57 +0300 |
|---|---|---|
| committer | Paul Buetow <paul@buetow.org> | 2026-04-02 09:10:57 +0300 |
| commit | faee049f51a3e6125a710fae3e713638bd68dd48 (patch) | |
| tree | b76365d8108ce98d8a7bfa0fee3ff1571215e5d9 | |
| parent | 7b26fefeae30c5e02d5848ae24405f9b2436decb (diff) | |
Add Gemini model listing support
| -rw-r--r-- | cmd/totalrecall/main.go | 2 | ||||
| -rw-r--r-- | internal/cli/command.go | 2 | ||||
| -rw-r--r-- | internal/models/doc.go | 4 | ||||
| -rw-r--r-- | internal/models/lister.go | 178 | ||||
| -rw-r--r-- | internal/models/lister_test.go | 205 |
5 files changed, 336 insertions, 55 deletions
diff --git a/cmd/totalrecall/main.go b/cmd/totalrecall/main.go index d719be5..d81bf4a 100644 --- a/cmd/totalrecall/main.go +++ b/cmd/totalrecall/main.go @@ -50,7 +50,7 @@ func runCommand(cmd *cobra.Command, args []string, flags *cli.Flags) error { // Handle --list-models flag if flags.ListModels { - lister := models.NewLister(cli.GetOpenAIKey()) + lister := models.NewLister(cli.GetOpenAIKey(), cli.GetGoogleAPIKey(), os.Stdout) return lister.ListAvailableModels() } diff --git a/internal/cli/command.go b/internal/cli/command.go index 0e1e93b..621d43a 100644 --- a/internal/cli/command.go +++ b/internal/cli/command.go @@ -66,7 +66,7 @@ func setupFlags(cmd *cobra.Command, flags *Flags) { cmd.Flags().BoolVar(&flags.GenerateAnki, "anki", false, "Generate Anki import file (APKG format by default, use --anki-csv for legacy CSV)") cmd.Flags().BoolVar(&flags.AnkiCSV, "anki-csv", false, "Generate legacy CSV format instead of APKG when using --anki") cmd.Flags().StringVar(&flags.DeckName, "deck-name", flags.DeckName, "Deck name for APKG export") - cmd.Flags().BoolVar(&flags.ListModels, "list-models", false, "List available OpenAI models for the current API key") + cmd.Flags().BoolVar(&flags.ListModels, "list-models", false, "List available OpenAI and Gemini models for the configured API keys") cmd.Flags().BoolVar(&flags.AllVoices, "all-voices", false, "Generate audio in all available voices (creates multiple files)") cmd.Flags().BoolVar(&flags.NoAutoPlay, "no-auto-play", false, "Disable automatic audio playback in GUI mode (auto-play is enabled by default)") cmd.Flags().BoolVar(&flags.Archive, "archive", false, "Archive existing cards directory with timestamp") diff --git a/internal/models/doc.go b/internal/models/doc.go index 116a04a..0935b58 100644 --- a/internal/models/doc.go +++ b/internal/models/doc.go @@ -1,4 +1,4 @@ // Package models provides functionality for listing and categorizing -// available OpenAI models. It helps users discover which TTS, image -// generation, and chat models are available with their API key. +// available OpenAI and Gemini models. It helps users discover which TTS, +// image generation, chat, and Gemini models are available with their API keys. package models diff --git a/internal/models/lister.go b/internal/models/lister.go index bb383bc..c798dc6 100644 --- a/internal/models/lister.go +++ b/internal/models/lister.go @@ -3,40 +3,101 @@ package models import ( "context" "fmt" + "io" + "os" "sort" "strings" "github.com/sashabaranov/go-openai" + "google.golang.org/genai" ) -// Lister handles listing available OpenAI models +type openAIModelLister interface { + ListModels(context.Context) (openai.ModelsList, error) +} + +type geminiModelLister interface { + List(context.Context, *genai.ListModelsConfig) (genai.Page[genai.Model], error) +} + +// Lister handles listing available OpenAI and Gemini models. type Lister struct { - apiKey string - client *openai.Client + openAIKey string + geminiKey string + openAIClient openAIModelLister + geminiClient geminiModelLister + geminiInitErr error + out io.Writer } -// NewLister creates a new model lister -func NewLister(apiKey string) *Lister { - return &Lister{ - apiKey: apiKey, - client: openai.NewClient(apiKey), +// NewLister creates a new model lister. +func NewLister(openAIKey, geminiKey string, out io.Writer) *Lister { + lister := &Lister{ + openAIKey: strings.TrimSpace(openAIKey), + geminiKey: strings.TrimSpace(geminiKey), + out: out, + } + + if lister.out == nil { + lister.out = os.Stdout + } + + if lister.openAIKey != "" { + lister.openAIClient = openai.NewClient(lister.openAIKey) + } + + if lister.geminiKey != "" { + client, err := genai.NewClient(context.Background(), &genai.ClientConfig{ + APIKey: lister.geminiKey, + }) + if err != nil { + lister.geminiInitErr = err + } else { + lister.geminiClient = client.Models + } } + + return lister } -// ListAvailableModels lists all available OpenAI models categorized by type +// ListAvailableModels lists all available OpenAI and Gemini models categorized by provider. func (l *Lister) ListAvailableModels() error { - if l.apiKey == "" { - return fmt.Errorf("OpenAI API key not found. Set OPENAI_API_KEY environment variable or configure in .totalrecall.yaml") + if l.openAIKey == "" && l.geminiKey == "" { + return fmt.Errorf("no API keys found. Set OPENAI_API_KEY and/or GOOGLE_API_KEY environment variable(s) or configure them in .totalrecall.yaml") } - // List models - ctx := context.Background() - models, err := l.client.ListModels(ctx) + fmt.Fprintln(l.out, "Available Models:") + + printedSection := false + if l.openAIKey != "" { + if err := l.printOpenAIModels(); err != nil { + return err + } + printedSection = true + } + + if l.geminiKey != "" { + if printedSection { + fmt.Fprintln(l.out) + } + if err := l.printGeminiModels(); err != nil { + return err + } + } + + return nil +} + +func (l *Lister) printOpenAIModels() error { + if l.openAIClient == nil { + return fmt.Errorf("OpenAI client not initialized") + } + + models, err := l.openAIClient.ListModels(context.Background()) if err != nil { - return fmt.Errorf("failed to list models: %w", err) + return fmt.Errorf("failed to list OpenAI models: %w", err) } - // Categorize models ttsModels := []string{} imageModels := []string{} chatModels := []string{} @@ -57,27 +118,26 @@ func (l *Lister) ListAvailableModels() error { sort.Strings(imageModels) sort.Strings(chatModels) - // Print models - fmt.Println("Available OpenAI Models:") - fmt.Println("\nText-to-Speech (TTS) Models:") + fmt.Fprintln(l.out, "OpenAI Models:") + fmt.Fprintln(l.out, " Text-to-Speech (TTS) Models:") if len(ttsModels) == 0 { - fmt.Println(" No TTS models found") + fmt.Fprintln(l.out, " No TTS models found") } else { for _, model := range ttsModels { - fmt.Printf(" %s\n", model) + fmt.Fprintf(l.out, " %s\n", model) } } - fmt.Println("\nImage Generation Models:") + fmt.Fprintln(l.out, " Image Generation Models:") if len(imageModels) == 0 { - fmt.Println(" No image models found") + fmt.Fprintln(l.out, " No image models found") } else { for _, model := range imageModels { - fmt.Printf(" %s\n", model) + fmt.Fprintf(l.out, " %s\n", model) } } - fmt.Println("\nChat/Translation Models (for Bulgarian translation):") + fmt.Fprintln(l.out, " Chat/Translation Models (for Bulgarian translation):") if len(chatModels) > 10 { // Show only relevant models relevantModels := []string{} @@ -87,14 +147,76 @@ func (l *Lister) ListAvailableModels() error { } } for _, model := range relevantModels { - fmt.Printf(" %s\n", model) + fmt.Fprintf(l.out, " %s\n", model) } - fmt.Printf(" ... and %d more models\n", len(chatModels)-len(relevantModels)) + fmt.Fprintf(l.out, " ... and %d more models\n", len(chatModels)-len(relevantModels)) } else { for _, model := range chatModels { - fmt.Printf(" %s\n", model) + fmt.Fprintf(l.out, " %s\n", model) + } + } + + return nil +} + +func (l *Lister) printGeminiModels() error { + if l.geminiInitErr != nil { + return fmt.Errorf("failed to initialize Gemini client: %w", l.geminiInitErr) + } + if l.geminiClient == nil { + return fmt.Errorf("Gemini client not initialized") + } + + ctx := context.Background() + config := &genai.ListModelsConfig{ + QueryBase: genai.Ptr(true), + } + + geminiModels := []string{} + for { + models, err := l.geminiClient.List(ctx, config) + if err != nil { + return fmt.Errorf("failed to list Gemini models: %w", err) } + + geminiModels = append(geminiModels, collectGeminiModelIDs(models)...) + if models.NextPageToken == "" { + break + } + + config.PageToken = models.NextPageToken + } + + sort.Strings(geminiModels) + + fmt.Fprintln(l.out, "Gemini Models:") + if len(geminiModels) == 0 { + fmt.Fprintln(l.out, " No Gemini models found") + return nil + } + + for _, model := range geminiModels { + fmt.Fprintf(l.out, " %s\n", model) } return nil } + +func collectGeminiModelIDs(page genai.Page[genai.Model]) []string { + modelIDs := make([]string, 0, len(page.Items)) + for _, model := range page.Items { + if model == nil { + continue + } + + modelID := strings.TrimPrefix(strings.TrimSpace(model.Name), "models/") + if modelID == "" { + modelID = strings.TrimSpace(model.DisplayName) + } + if modelID != "" { + modelIDs = append(modelIDs, modelID) + } + } + + return modelIDs +} diff --git a/internal/models/lister_test.go b/internal/models/lister_test.go index b125981..d867ec4 100644 --- a/internal/models/lister_test.go +++ b/internal/models/lister_test.go @@ -1,53 +1,212 @@ package models import ( - "os" + "bytes" + "context" + "strings" "testing" + + "github.com/sashabaranov/go-openai" + "google.golang.org/genai" ) +type fakeOpenAIClient struct { + models openai.ModelsList + err error +} + +func (f *fakeOpenAIClient) ListModels(context.Context) (openai.ModelsList, error) { + return f.models, f.err +} + +type fakeGeminiClient struct { + pages map[string]genai.Page[genai.Model] + err error + calls []string +} + +func (f *fakeGeminiClient) List(_ context.Context, config *genai.ListModelsConfig) (genai.Page[genai.Model], error) { + token := "" + if config != nil { + token = config.PageToken + } + + f.calls = append(f.calls, token) + + if f.err != nil { + return genai.Page[genai.Model]{}, f.err + } + + if page, ok := f.pages[token]; ok { + return page, nil + } + + return genai.Page[genai.Model]{}, nil +} + func TestNewLister(t *testing.T) { - lister := NewLister("test-api-key") + lister := NewLister(" test-openai-key ", " test-gemini-key ", nil) if lister == nil { t.Fatal("NewLister returned nil") } - if lister.apiKey != "test-api-key" { - t.Errorf("Expected API key 'test-api-key', got '%s'", lister.apiKey) + if lister.openAIKey != "test-openai-key" { + t.Fatalf("openAIKey = %q, want %q", lister.openAIKey, "test-openai-key") } - if lister.client == nil { - t.Error("OpenAI client not initialized") + if lister.geminiKey != "test-gemini-key" { + t.Fatalf("geminiKey = %q, want %q", lister.geminiKey, "test-gemini-key") + } + + if lister.openAIKey != "" && lister.openAIClient == nil { + t.Fatal("OpenAI client not initialized") + } + + if lister.geminiKey != "" && lister.geminiClient == nil && lister.geminiInitErr == nil { + t.Fatal("Gemini client not initialized") } } -func TestListAvailableModels_NoAPIKey(t *testing.T) { - lister := NewLister("") +func TestListAvailableModels_NoAPIKeys(t *testing.T) { + var output bytes.Buffer + lister := &Lister{out: &output} err := lister.ListAvailableModels() if err == nil { - t.Error("Expected error for missing API key") + t.Fatal("Expected error for missing API keys") } - expectedError := "OpenAI API key not found. Set OPENAI_API_KEY environment variable or configure in .totalrecall.yaml" - if err.Error() != expectedError { - t.Errorf("Expected error '%s', got: %v", expectedError, err) + expectedError := "no API keys found" + if !strings.Contains(err.Error(), expectedError) { + t.Fatalf("Expected error containing %q, got %v", expectedError, err) } } -func TestListAvailableModels_Integration(t *testing.T) { - // Skip if no API key - apiKey := os.Getenv("OPENAI_API_KEY") - if apiKey == "" { - t.Skip("Skipping integration test: OPENAI_API_KEY not set") +func TestListAvailableModels_OpenAIOnly(t *testing.T) { + var output bytes.Buffer + lister := &Lister{ + openAIKey: "test-openai-key", + openAIClient: &fakeOpenAIClient{ + models: openai.ModelsList{ + Models: []openai.Model{ + {ID: "tts-1"}, + {ID: "dall-e-3"}, + {ID: "gpt-4o"}, + {ID: "gpt-4o-mini-tts"}, + {ID: "gpt-4.1"}, + }, + }, + }, + out: &output, } - lister := NewLister(apiKey) + if err := lister.ListAvailableModels(); err != nil { + t.Fatalf("ListAvailableModels failed: %v", err) + } - // This test just verifies the method runs without error - // The actual output goes to stdout which we don't capture in tests - err := lister.ListAvailableModels() - if err != nil { - t.Errorf("ListAvailableModels failed: %v", err) + got := output.String() + for _, want := range []string{ + "Available Models:", + "OpenAI Models:", + "Text-to-Speech (TTS) Models:", + "Image Generation Models:", + "Chat/Translation Models (for Bulgarian translation):", + "tts-1", + "dall-e-3", + "gpt-4o", + "gpt-4o-mini-tts", + } { + if !strings.Contains(got, want) { + t.Fatalf("output missing %q:\n%s", want, got) + } + } + + if strings.Contains(got, "Gemini Models:") { + t.Fatalf("output unexpectedly contained Gemini section:\n%s", got) + } +} + +func TestListAvailableModels_GeminiOnly(t *testing.T) { + var output bytes.Buffer + geminiClient := &fakeGeminiClient{ + pages: map[string]genai.Page[genai.Model]{ + "": { + Items: []*genai.Model{ + {Name: "models/gemini-2.5-pro"}, + {Name: "models/gemini-2.5-flash"}, + }, + NextPageToken: "page-2", + }, + "page-2": { + Items: []*genai.Model{ + {Name: "models/gemini-2.5-flash-preview-tts"}, + {DisplayName: "Gemini Experimental"}, + }, + }, + }, + } + + lister := &Lister{ + geminiKey: "test-gemini-key", + geminiClient: geminiClient, + out: &output, + } + + if err := lister.ListAvailableModels(); err != nil { + t.Fatalf("ListAvailableModels failed: %v", err) + } + + got := output.String() + for _, want := range []string{ + "Available Models:", + "Gemini Models:", + "gemini-2.5-flash", + "gemini-2.5-pro", + "gemini-2.5-flash-preview-tts", + "Gemini Experimental", + } { + if !strings.Contains(got, want) { + t.Fatalf("output missing %q:\n%s", want, got) + } + } + + if len(geminiClient.calls) != 2 { + t.Fatalf("expected 2 Gemini page requests, got %d", len(geminiClient.calls)) + } +} + +func TestListAvailableModels_BothProviders(t *testing.T) { + var output bytes.Buffer + lister := &Lister{ + openAIKey: "test-openai-key", + openAIClient: &fakeOpenAIClient{ + models: openai.ModelsList{ + Models: []openai.Model{{ID: "tts-1"}}, + }, + }, + geminiKey: "test-gemini-key", + geminiClient: &fakeGeminiClient{ + pages: map[string]genai.Page[genai.Model]{ + "": { + Items: []*genai.Model{{Name: "models/gemini-2.5-flash"}}, + }, + }, + }, + out: &output, + } + + if err := lister.ListAvailableModels(); err != nil { + t.Fatalf("ListAvailableModels failed: %v", err) + } + + got := output.String() + openAIIndex := strings.Index(got, "OpenAI Models:") + geminiIndex := strings.Index(got, "Gemini Models:") + if openAIIndex == -1 || geminiIndex == -1 { + t.Fatalf("expected both provider sections, got:\n%s", got) + } + if openAIIndex > geminiIndex { + t.Fatalf("expected OpenAI section before Gemini section, got:\n%s", got) } } |
