Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
48 changes: 41 additions & 7 deletions internal/cmd/generate.go
Original file line number Diff line number Diff line change
Expand Up @@ -15,8 +15,8 @@ import (
)

var generateCmd = &cobra.Command{
Use: "generate",
Short: "Generate AI content (image, video, audio, 3D, chat)",
Use: "generate",
Short: "Generate AI content (image, video, audio, 3D, chat)",
Aliases: []string{"gen"},
}

Expand Down Expand Up @@ -54,6 +54,16 @@ func pollAndDownload(cmd *cobra.Command, genType, fetchEndpoint string, result m
return handleCompleted(result, download, outputDir, genType)
}

// The API reports a failed generation as HTTP 200 with status "error" and no
// job id, so the client sees no error. Polling then asked fetch/ for an
// empty id until the timeout, and --no-wait printed "Job queued" for it.
if status == "error" || status == "failed" {
return generationFailed(result)
}
if requestID == "" {
return fmt.Errorf("the API returned no job id to poll: %s", responseSummary(result))
}

if noWait {
outputResult(result, func() {
fmt.Printf("Job queued: %s\n", requestID)
Expand Down Expand Up @@ -96,11 +106,7 @@ func pollAndDownload(cmd *cobra.Command, genType, fetchEndpoint string, result m
case "success":
return handleCompleted(fetchResult, download, outputDir, genType)
case "error", "failed":
msg := "Generation failed"
if m, ok := fetchResult["message"].(string); ok {
msg = m
}
return fmt.Errorf("%s", msg)
return generationFailed(fetchResult)
case "processing":
elapsed := time.Since(startTime).Round(time.Second)
eta := ""
Expand All @@ -114,6 +120,34 @@ func pollAndDownload(cmd *cobra.Command, genType, fetchEndpoint string, result m
}
}

// generationFailed turns an API "error"/"failed" body into an error that carries
// its message, which may be a string or, for validation errors, an object.
func generationFailed(result map[string]interface{}) error {
switch m := result["message"].(type) {
case string:
if m != "" {
return fmt.Errorf("%s", m)
}
case nil:
default:
if encoded, err := json.Marshal(m); err == nil {
return fmt.Errorf("%s", encoded)
}
}

return fmt.Errorf("Generation failed")
}

// responseSummary is the raw body, for errors about a response the CLI could not use.
func responseSummary(result map[string]interface{}) string {
encoded, err := json.Marshal(result)
if err != nil {
return fmt.Sprintf("%v", result)
}

return string(encoded)
}

func hasOutputURLs(result map[string]interface{}) bool {
if output, ok := result["output"].([]interface{}); ok && len(output) > 0 {
return true
Expand Down
62 changes: 62 additions & 0 deletions internal/cmd/generate_poll_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,62 @@
package cmd

import (
"testing"
"time"

"github.com/spf13/cobra"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)

// A failed generation comes back as HTTP 200 with status "error" and no job id.
// pollAndDownload used to poll fetch/ with that empty id until the timeout.
func generationCommand(t *testing.T, noWait bool) *cobra.Command {
t.Helper()
cmd := &cobra.Command{}
addGenerationFlags(cmd)
require.NoError(t, cmd.Flags().Set("timeout", (2*time.Second).String()))
if noWait {
require.NoError(t, cmd.Flags().Set("no-wait", "true"))
}

return cmd
}

func TestPollAndDownload_ReturnsTheApiErrorWithoutPolling(t *testing.T) {
for _, noWait := range []bool{false, true} {
start := time.Now()

err := pollAndDownload(generationCommand(t, noWait), "video", "/v7/video-fusion/fetch", map[string]interface{}{
"status": "error",
"code": "provider_error",
"message": "The provider refused the request.",
})

require.Error(t, err)
assert.Equal(t, "The provider refused the request.", err.Error())
assert.Less(t, time.Since(start), time.Second, "it must not poll")
}
}

func TestPollAndDownload_ReportsAStructuredValidationMessage(t *testing.T) {
err := pollAndDownload(generationCommand(t, false), "image", "/v7/images/fetch", map[string]interface{}{
"status": "error",
"message": map[string]interface{}{"init_image": []interface{}{"The init image field is required."}},
})

require.Error(t, err)
assert.Contains(t, err.Error(), "The init image field is required.")
}

func TestPollAndDownload_RefusesToPollWithoutAJobID(t *testing.T) {
start := time.Now()

err := pollAndDownload(generationCommand(t, false), "video", "/v7/video-fusion/fetch", map[string]interface{}{
"status": "processing",
})

require.Error(t, err)
assert.Contains(t, err.Error(), "no job id")
assert.Less(t, time.Since(start), time.Second, "it must not poll")
}
Loading