summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorPaul Buetow <paul@buetow.org>2026-04-02 09:10:57 +0300
committerPaul Buetow <paul@buetow.org>2026-04-02 09:10:57 +0300
commitfaee049f51a3e6125a710fae3e713638bd68dd48 (patch)
treeb76365d8108ce98d8a7bfa0fee3ff1571215e5d9
parent7b26fefeae30c5e02d5848ae24405f9b2436decb (diff)
Add Gemini model listing support
-rw-r--r--cmd/totalrecall/main.go2
-rw-r--r--internal/cli/command.go2
-rw-r--r--internal/models/doc.go4
-rw-r--r--internal/models/lister.go178
-rw-r--r--internal/models/lister_test.go205
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)
}
}