fix(leonardo): 参考图走 uploadInitImage 永久桶;音频参考需搭配图/视频参考
This commit is contained in:
@@ -49,11 +49,112 @@ const mUploadImage = `mutation UploadImage($uploadImageInput: UploadImageInput!)
|
||||
}
|
||||
}`
|
||||
|
||||
// uploadInitImage uploads a reference (init) image for image-to-image: it asks
|
||||
// Leonardo for a presigned S3 POST, uploads the bytes, and returns the upload id
|
||||
// to reference in the Generate request's image_reference guidance.
|
||||
func (c *Client) uploadInitImage(ctx context.Context, cookie string, img []byte) (string, error) {
|
||||
return c.uploadAsset(ctx, cookie, "png", img)
|
||||
const mUploadInitImage = `mutation UploadInitImage($arg1: InitImageUploadInput!) {
|
||||
uploadInitImage(arg1: $arg1) {
|
||||
id
|
||||
url
|
||||
fields
|
||||
__typename
|
||||
}
|
||||
}`
|
||||
|
||||
// initImageExtension narrows a sniffed extension to what uploadInitImage accepts
|
||||
// (png / jpg / jpeg / webp); anything else is sent as png.
|
||||
func initImageExtension(extension string) string {
|
||||
switch strings.TrimPrefix(strings.ToLower(strings.TrimSpace(extension)), ".") {
|
||||
case "jpg":
|
||||
return "jpg"
|
||||
case "jpeg":
|
||||
return "jpeg"
|
||||
case "webp":
|
||||
return "webp"
|
||||
default:
|
||||
return "png"
|
||||
}
|
||||
}
|
||||
|
||||
// uploadInitImage uploads a reference image and returns the init image id to put
|
||||
// in a Generate request's image_reference guidance. It has to go through
|
||||
// uploadInitImage (permanent init-image bucket): the uploadImage mutation only
|
||||
// hands out temporary-bucket ids, which the generation service can't resolve.
|
||||
func (c *Client) uploadInitImage(ctx context.Context, cookie, extension string, img []byte) (string, error) {
|
||||
extension = initImageExtension(extension)
|
||||
payload, _ := json.Marshal(map[string]any{
|
||||
"operationName": "UploadInitImage",
|
||||
"query": mUploadInitImage,
|
||||
"variables": map[string]any{"arg1": map[string]any{"extension": extension}},
|
||||
})
|
||||
body, err := c.callGraphQL(ctx, cookie, payload, false, "upload-init-image")
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
var ur struct {
|
||||
Data struct {
|
||||
UploadInitImage struct {
|
||||
ID string `json:"id"`
|
||||
URL string `json:"url"`
|
||||
Fields string `json:"fields"`
|
||||
} `json:"uploadInitImage"`
|
||||
} `json:"data"`
|
||||
}
|
||||
if err := json.Unmarshal(body, &ur); err != nil {
|
||||
return "", fmt.Errorf("%w: upload-init-image non-json", ErrTemporaryUpstream)
|
||||
}
|
||||
up := ur.Data.UploadInitImage
|
||||
if up.ID == "" || up.URL == "" {
|
||||
return "", fmt.Errorf("%w: no upload url", ErrTemporaryUpstream)
|
||||
}
|
||||
if err := c.putPresigned(ctx, up.URL, up.Fields, "asset."+extension, img); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return up.ID, nil
|
||||
}
|
||||
|
||||
// putPresigned performs the presigned S3 POST: all policy fields first, the file
|
||||
// part LAST.
|
||||
func (c *Client) putPresigned(ctx context.Context, url, fieldsJSON, filename string, asset []byte) error {
|
||||
var fields map[string]string
|
||||
if err := json.Unmarshal([]byte(fieldsJSON), &fields); err != nil {
|
||||
return fmt.Errorf("%w: bad upload fields", ErrTemporaryUpstream)
|
||||
}
|
||||
var buf bytes.Buffer
|
||||
w := multipart.NewWriter(&buf)
|
||||
for k, v := range fields {
|
||||
_ = w.WriteField(k, v)
|
||||
}
|
||||
fw, err := w.CreateFormFile("file", filename)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if _, err := fw.Write(asset); err != nil {
|
||||
return err
|
||||
}
|
||||
_ = w.Close()
|
||||
|
||||
client, err := c.newDirectTLSClient()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
req, err := http.NewRequest(http.MethodPost, url, &buf)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
req = req.WithContext(ctx)
|
||||
req.Header = http.Header{
|
||||
"content-type": {w.FormDataContentType()},
|
||||
"user-agent": {userAgent},
|
||||
"origin": {appBase},
|
||||
"referer": {appBase + "/"},
|
||||
}
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return fmt.Errorf("%w: s3 upload: %s", ErrTemporaryUpstream, err.Error())
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
if resp.StatusCode != 204 && resp.StatusCode != 200 && resp.StatusCode != 201 {
|
||||
return fmt.Errorf("%w: s3 upload http %d", ErrTemporaryUpstream, resp.StatusCode)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// uploadAsset uploads one reference asset (extension png / mp3 / mp4 …) through
|
||||
@@ -66,7 +167,12 @@ func (c *Client) uploadAsset(ctx context.Context, cookie, extension string, asse
|
||||
payload, _ := json.Marshal(map[string]any{
|
||||
"operationName": "UploadImage",
|
||||
"query": mUploadImage,
|
||||
"variables": map[string]any{"uploadImageInput": map[string]any{"uploadType": "INIT", "extension": extension}},
|
||||
// originalFilename is mandatory for audio uploads and harmless otherwise.
|
||||
"variables": map[string]any{"uploadImageInput": map[string]any{
|
||||
"uploadType": "INIT",
|
||||
"extension": extension,
|
||||
"originalFilename": "asset." + extension,
|
||||
}},
|
||||
})
|
||||
body, err := c.callGraphQL(ctx, cookie, payload, false, "upload-init")
|
||||
if err != nil {
|
||||
@@ -88,49 +194,9 @@ func (c *Client) uploadAsset(ctx context.Context, cookie, extension string, asse
|
||||
if up.UploadID == "" || up.URL == "" {
|
||||
return "", fmt.Errorf("%w: no upload url", ErrTemporaryUpstream)
|
||||
}
|
||||
var fields map[string]string
|
||||
if err := json.Unmarshal([]byte(up.Fields), &fields); err != nil {
|
||||
return "", fmt.Errorf("%w: bad upload fields", ErrTemporaryUpstream)
|
||||
}
|
||||
|
||||
// Presigned S3 POST: all policy fields first, the file part LAST.
|
||||
var buf bytes.Buffer
|
||||
w := multipart.NewWriter(&buf)
|
||||
for k, v := range fields {
|
||||
_ = w.WriteField(k, v)
|
||||
}
|
||||
fw, err := w.CreateFormFile("file", "asset."+extension)
|
||||
if err != nil {
|
||||
if err := c.putPresigned(ctx, up.URL, up.Fields, "asset."+extension, asset); err != nil {
|
||||
return "", err
|
||||
}
|
||||
if _, err := fw.Write(asset); err != nil {
|
||||
return "", err
|
||||
}
|
||||
_ = w.Close()
|
||||
|
||||
client, err := c.newDirectTLSClient()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
req, err := http.NewRequest(http.MethodPost, up.URL, &buf)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
req = req.WithContext(ctx)
|
||||
req.Header = http.Header{
|
||||
"content-type": {w.FormDataContentType()},
|
||||
"user-agent": {userAgent},
|
||||
"origin": {appBase},
|
||||
"referer": {appBase + "/"},
|
||||
}
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("%w: s3 upload: %s", ErrTemporaryUpstream, err.Error())
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
if resp.StatusCode != 204 && resp.StatusCode != 200 && resp.StatusCode != 201 {
|
||||
return "", fmt.Errorf("%w: s3 upload http %d", ErrTemporaryUpstream, resp.StatusCode)
|
||||
}
|
||||
return up.UploadID, nil
|
||||
}
|
||||
|
||||
@@ -156,7 +222,7 @@ func (c *Client) GenerateImage(ctx context.Context, cookie, model, prompt string
|
||||
if len(img) == 0 {
|
||||
continue
|
||||
}
|
||||
uploadID, upErr := c.uploadInitImage(ctx, cookie, img)
|
||||
uploadID, upErr := c.uploadInitImage(ctx, cookie, assetExtension(img, "png"), img)
|
||||
if upErr != nil {
|
||||
return nil, nil, upErr
|
||||
}
|
||||
@@ -306,13 +372,23 @@ func (c *Client) pollImage(ctx context.Context, cookie, genID string) (string, e
|
||||
func graphqlError(body []byte) error {
|
||||
var env struct {
|
||||
Errors []struct {
|
||||
Message string `json:"message"`
|
||||
Message string `json:"message"`
|
||||
Extensions struct {
|
||||
Code string `json:"code"`
|
||||
Details struct {
|
||||
Message string `json:"message"`
|
||||
} `json:"details"`
|
||||
} `json:"extensions"`
|
||||
} `json:"errors"`
|
||||
}
|
||||
if err := json.Unmarshal(body, &env); err != nil || len(env.Errors) == 0 {
|
||||
return nil
|
||||
}
|
||||
msg := strings.TrimSpace(env.Errors[0].Message)
|
||||
// The generic "An error occurred." hides the real reason in extensions.
|
||||
if detail := strings.TrimSpace(env.Errors[0].Extensions.Details.Message); detail != "" && detail != msg {
|
||||
msg = msg + " (" + detail + ")"
|
||||
}
|
||||
low := strings.ToLower(msg)
|
||||
switch {
|
||||
case strings.Contains(low, "unauthor") || strings.Contains(low, "jwt") || strings.Contains(low, "token is") || strings.Contains(low, "forbidden"):
|
||||
|
||||
@@ -53,7 +53,7 @@ func (c *Client) GenerateVideo(ctx context.Context, cookie, model, prompt string
|
||||
if len(img) == 0 {
|
||||
continue
|
||||
}
|
||||
uploadID, upErr := c.uploadAsset(ctx, cookie, assetExtension(img, "png"), img)
|
||||
uploadID, upErr := c.uploadInitImage(ctx, cookie, assetExtension(img, "png"), img)
|
||||
if upErr != nil {
|
||||
return nil, nil, upErr
|
||||
}
|
||||
@@ -65,24 +65,6 @@ func (c *Client) GenerateVideo(ctx context.Context, cookie, model, prompt string
|
||||
if len(imageRefs) > 0 {
|
||||
guidances["image_reference"] = imageRefs
|
||||
}
|
||||
var audioRefs []map[string]any
|
||||
for _, aud := range refs.Audios {
|
||||
if len(aud) == 0 {
|
||||
continue
|
||||
}
|
||||
uploadID, upErr := c.uploadAsset(ctx, cookie, assetExtension(aud, "mp3"), aud)
|
||||
if upErr != nil {
|
||||
return nil, nil, upErr
|
||||
}
|
||||
audio := map[string]any{"id": uploadID, "type": "UPLOADED"}
|
||||
if secs := MediaDurationSeconds(aud); secs > 0 {
|
||||
audio["duration"] = secs
|
||||
}
|
||||
audioRefs = append(audioRefs, map[string]any{"audio": audio})
|
||||
}
|
||||
if len(audioRefs) > 0 {
|
||||
guidances["audio_reference"] = audioRefs
|
||||
}
|
||||
var videoRefs []map[string]any
|
||||
for _, vid := range refs.Videos {
|
||||
if len(vid) == 0 {
|
||||
@@ -101,6 +83,29 @@ func (c *Client) GenerateVideo(ctx context.Context, cookie, model, prompt string
|
||||
if len(videoRefs) > 0 {
|
||||
guidances["video_reference_base"] = videoRefs
|
||||
}
|
||||
// Leonardo rejects an audio reference that isn't paired with an image or
|
||||
// video reference (audio_reference_only_compatible_with_image_or_video_reference).
|
||||
var audioRefs []map[string]any
|
||||
for _, aud := range refs.Audios {
|
||||
if len(aud) == 0 {
|
||||
continue
|
||||
}
|
||||
if len(imageRefs) == 0 && len(videoRefs) == 0 {
|
||||
return nil, nil, errors.New("leonardo: audio reference requires an image or video reference")
|
||||
}
|
||||
uploadID, upErr := c.uploadAsset(ctx, cookie, assetExtension(aud, "mp3"), aud)
|
||||
if upErr != nil {
|
||||
return nil, nil, upErr
|
||||
}
|
||||
audio := map[string]any{"id": uploadID, "type": "UPLOADED"}
|
||||
if secs := MediaDurationSeconds(aud); secs > 0 {
|
||||
audio["duration"] = secs
|
||||
}
|
||||
audioRefs = append(audioRefs, map[string]any{"audio": audio})
|
||||
}
|
||||
if len(audioRefs) > 0 {
|
||||
guidances["audio_reference"] = audioRefs
|
||||
}
|
||||
|
||||
parameters := map[string]any{
|
||||
"height": height,
|
||||
|
||||
Reference in New Issue
Block a user