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
|
const mUploadInitImage = `mutation UploadInitImage($arg1: InitImageUploadInput!) {
|
||||||
// Leonardo for a presigned S3 POST, uploads the bytes, and returns the upload id
|
uploadInitImage(arg1: $arg1) {
|
||||||
// to reference in the Generate request's image_reference guidance.
|
id
|
||||||
func (c *Client) uploadInitImage(ctx context.Context, cookie string, img []byte) (string, error) {
|
url
|
||||||
return c.uploadAsset(ctx, cookie, "png", img)
|
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
|
// 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{
|
payload, _ := json.Marshal(map[string]any{
|
||||||
"operationName": "UploadImage",
|
"operationName": "UploadImage",
|
||||||
"query": mUploadImage,
|
"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")
|
body, err := c.callGraphQL(ctx, cookie, payload, false, "upload-init")
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -88,49 +194,9 @@ func (c *Client) uploadAsset(ctx context.Context, cookie, extension string, asse
|
|||||||
if up.UploadID == "" || up.URL == "" {
|
if up.UploadID == "" || up.URL == "" {
|
||||||
return "", fmt.Errorf("%w: no upload url", ErrTemporaryUpstream)
|
return "", fmt.Errorf("%w: no upload url", ErrTemporaryUpstream)
|
||||||
}
|
}
|
||||||
var fields map[string]string
|
if err := c.putPresigned(ctx, up.URL, up.Fields, "asset."+extension, asset); err != nil {
|
||||||
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 {
|
|
||||||
return "", err
|
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
|
return up.UploadID, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -156,7 +222,7 @@ func (c *Client) GenerateImage(ctx context.Context, cookie, model, prompt string
|
|||||||
if len(img) == 0 {
|
if len(img) == 0 {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
uploadID, upErr := c.uploadInitImage(ctx, cookie, img)
|
uploadID, upErr := c.uploadInitImage(ctx, cookie, assetExtension(img, "png"), img)
|
||||||
if upErr != nil {
|
if upErr != nil {
|
||||||
return nil, nil, upErr
|
return nil, nil, upErr
|
||||||
}
|
}
|
||||||
@@ -307,12 +373,22 @@ func graphqlError(body []byte) error {
|
|||||||
var env struct {
|
var env struct {
|
||||||
Errors []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"`
|
} `json:"errors"`
|
||||||
}
|
}
|
||||||
if err := json.Unmarshal(body, &env); err != nil || len(env.Errors) == 0 {
|
if err := json.Unmarshal(body, &env); err != nil || len(env.Errors) == 0 {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
msg := strings.TrimSpace(env.Errors[0].Message)
|
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)
|
low := strings.ToLower(msg)
|
||||||
switch {
|
switch {
|
||||||
case strings.Contains(low, "unauthor") || strings.Contains(low, "jwt") || strings.Contains(low, "token is") || strings.Contains(low, "forbidden"):
|
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 {
|
if len(img) == 0 {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
uploadID, upErr := c.uploadAsset(ctx, cookie, assetExtension(img, "png"), img)
|
uploadID, upErr := c.uploadInitImage(ctx, cookie, assetExtension(img, "png"), img)
|
||||||
if upErr != nil {
|
if upErr != nil {
|
||||||
return nil, nil, upErr
|
return nil, nil, upErr
|
||||||
}
|
}
|
||||||
@@ -65,24 +65,6 @@ func (c *Client) GenerateVideo(ctx context.Context, cookie, model, prompt string
|
|||||||
if len(imageRefs) > 0 {
|
if len(imageRefs) > 0 {
|
||||||
guidances["image_reference"] = imageRefs
|
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
|
var videoRefs []map[string]any
|
||||||
for _, vid := range refs.Videos {
|
for _, vid := range refs.Videos {
|
||||||
if len(vid) == 0 {
|
if len(vid) == 0 {
|
||||||
@@ -101,6 +83,29 @@ func (c *Client) GenerateVideo(ctx context.Context, cookie, model, prompt string
|
|||||||
if len(videoRefs) > 0 {
|
if len(videoRefs) > 0 {
|
||||||
guidances["video_reference_base"] = videoRefs
|
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{
|
parameters := map[string]any{
|
||||||
"height": height,
|
"height": height,
|
||||||
|
|||||||
Reference in New Issue
Block a user