参考图适配
This commit is contained in:
@@ -147,6 +147,10 @@ func (h *V1Handler) CreateVideo(c *gin.Context) {
|
|||||||
Prompt string `json:"prompt"`
|
Prompt string `json:"prompt"`
|
||||||
Seconds json.RawMessage `json:"seconds"`
|
Seconds json.RawMessage `json:"seconds"`
|
||||||
Size string `json:"size"`
|
Size string `json:"size"`
|
||||||
|
// Reference frames (image-to-video / first-last frames) as base64 or
|
||||||
|
// data-URI strings — the JSON equivalent of multipart input_reference.
|
||||||
|
InputReference []string `json:"input_reference"`
|
||||||
|
ReferenceImages []string `json:"reference_images"`
|
||||||
}
|
}
|
||||||
if err := c.ShouldBindJSON(&body); err != nil {
|
if err := c.ShouldBindJSON(&body); err != nil {
|
||||||
c.JSON(http.StatusBadRequest, gin.H{"detail": "invalid request body"})
|
c.JSON(http.StatusBadRequest, gin.H{"detail": "invalid request body"})
|
||||||
@@ -154,6 +158,7 @@ func (h *V1Handler) CreateVideo(c *gin.Context) {
|
|||||||
}
|
}
|
||||||
modelID, prompt, size = body.Model, body.Prompt, body.Size
|
modelID, prompt, size = body.Model, body.Prompt, body.Size
|
||||||
seconds = rawToString(body.Seconds)
|
seconds = rawToString(body.Seconds)
|
||||||
|
refs = append(body.InputReference, body.ReferenceImages...)
|
||||||
}
|
}
|
||||||
duration := strings.TrimSpace(seconds)
|
duration := strings.TrimSpace(seconds)
|
||||||
if duration != "" && !strings.HasSuffix(duration, "s") {
|
if duration != "" && !strings.HasSuffix(duration, "s") {
|
||||||
|
|||||||
@@ -122,21 +122,47 @@ func (c *Client) GenerateImage(ctx context.Context, baseURL, apiKey, model, prom
|
|||||||
|
|
||||||
// GenerateVideo drives the upstream Sora-style async video API:
|
// GenerateVideo drives the upstream Sora-style async video API:
|
||||||
// POST /v1/videos → poll GET /v1/videos/{id} → GET /v1/videos/{id}/content.
|
// POST /v1/videos → poll GET /v1/videos/{id} → GET /v1/videos/{id}/content.
|
||||||
// When downloadResult is false it returns the upstream content URL instead.
|
// Reference frames (image-to-video / first-last frames) are sent as multipart
|
||||||
func (c *Client) GenerateVideo(ctx context.Context, baseURL, apiKey, model, prompt, size string, seconds int, downloadResult bool) ([]byte, string, error) {
|
// input_reference[] files, matching the OpenAI videos API. When downloadResult
|
||||||
|
// is false it returns the upstream content URL instead.
|
||||||
|
func (c *Client) GenerateVideo(ctx context.Context, baseURL, apiKey, model, prompt, size string, seconds int, frames [][]byte, downloadResult bool) ([]byte, string, error) {
|
||||||
baseURL = strings.TrimRight(strings.TrimSpace(baseURL), "/")
|
baseURL = strings.TrimRight(strings.TrimSpace(baseURL), "/")
|
||||||
if baseURL == "" || apiKey == "" {
|
if baseURL == "" || apiKey == "" {
|
||||||
return nil, "", ErrAuth
|
return nil, "", ErrAuth
|
||||||
}
|
}
|
||||||
payload := map[string]any{"model": model, "prompt": prompt}
|
var created map[string]any
|
||||||
if size != "" {
|
var err error
|
||||||
payload["size"] = size
|
if len(frames) > 0 {
|
||||||
|
body := &bytes.Buffer{}
|
||||||
|
w := multipart.NewWriter(body)
|
||||||
|
_ = w.WriteField("model", model)
|
||||||
|
_ = w.WriteField("prompt", prompt)
|
||||||
|
if size != "" {
|
||||||
|
_ = w.WriteField("size", size)
|
||||||
|
}
|
||||||
|
if seconds > 0 {
|
||||||
|
_ = w.WriteField("seconds", fmt.Sprintf("%d", seconds))
|
||||||
|
}
|
||||||
|
for i, f := range frames {
|
||||||
|
fw, e := w.CreateFormFile("input_reference[]", fmt.Sprintf("frame_%d.png", i+1))
|
||||||
|
if e != nil {
|
||||||
|
return nil, "", e
|
||||||
|
}
|
||||||
|
_, _ = fw.Write(f)
|
||||||
|
}
|
||||||
|
_ = w.Close()
|
||||||
|
created, err = c.doMultipart(ctx, baseURL+"/v1/videos", apiKey, body, w.FormDataContentType())
|
||||||
|
} else {
|
||||||
|
payload := map[string]any{"model": model, "prompt": prompt}
|
||||||
|
if size != "" {
|
||||||
|
payload["size"] = size
|
||||||
|
}
|
||||||
|
if seconds > 0 {
|
||||||
|
payload["seconds"] = fmt.Sprintf("%d", seconds)
|
||||||
|
}
|
||||||
|
raw, _ := json.Marshal(payload)
|
||||||
|
created, err = c.doJSON(ctx, http.MethodPost, baseURL+"/v1/videos", apiKey, raw)
|
||||||
}
|
}
|
||||||
if seconds > 0 {
|
|
||||||
payload["seconds"] = fmt.Sprintf("%d", seconds)
|
|
||||||
}
|
|
||||||
raw, _ := json.Marshal(payload)
|
|
||||||
created, err := c.doJSON(ctx, http.MethodPost, baseURL+"/v1/videos", apiKey, raw)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, "", err
|
return nil, "", err
|
||||||
}
|
}
|
||||||
@@ -216,6 +242,33 @@ func (c *Client) doJSON(ctx context.Context, method, url, apiKey string, body []
|
|||||||
return out, nil
|
return out, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (c *Client) doMultipart(ctx context.Context, url, apiKey string, body io.Reader, contentType string) (map[string]any, error) {
|
||||||
|
req, err := http.NewRequest(http.MethodPost, url, body)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
req = req.WithContext(ctx)
|
||||||
|
req.Header.Set("Authorization", "Bearer "+apiKey)
|
||||||
|
req.Header.Set("Content-Type", contentType)
|
||||||
|
resp, err := httpClient().Do(req)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("%w: %s", ErrTemporaryUpstream, sanitizeErr(err))
|
||||||
|
}
|
||||||
|
defer resp.Body.Close()
|
||||||
|
raw, _ := io.ReadAll(resp.Body)
|
||||||
|
if e := mapStatus(resp.StatusCode, raw); e != nil {
|
||||||
|
return nil, e
|
||||||
|
}
|
||||||
|
var out map[string]any
|
||||||
|
if len(raw) == 0 {
|
||||||
|
return map[string]any{}, nil
|
||||||
|
}
|
||||||
|
if err := json.Unmarshal(raw, &out); err != nil {
|
||||||
|
return nil, fmt.Errorf("%w: non-json: %s", ErrTemporaryUpstream, clip(raw, 120))
|
||||||
|
}
|
||||||
|
return out, nil
|
||||||
|
}
|
||||||
|
|
||||||
func (c *Client) download(ctx context.Context, url, apiKey string) ([]byte, error) {
|
func (c *Client) download(ctx context.Context, url, apiKey string) ([]byte, error) {
|
||||||
req, _ := http.NewRequest(http.MethodGet, url, nil)
|
req, _ := http.NewRequest(http.MethodGet, url, nil)
|
||||||
req = req.WithContext(ctx)
|
req = req.WithContext(ctx)
|
||||||
|
|||||||
@@ -1989,6 +1989,12 @@ func (s *V1Service) generateCustomVideo(ctx context.Context, eventID string, mod
|
|||||||
return nil, "", ErrNoProviderAccount
|
return nil, "", ErrNoProviderAccount
|
||||||
}
|
}
|
||||||
size := upstreamVideoSize(aspectRatio, resolution)
|
size := upstreamVideoSize(aspectRatio, resolution)
|
||||||
|
// Optional reference frames (image-to-video / first-last frames) — forwarded
|
||||||
|
// to the upstream as multipart input_reference[] files.
|
||||||
|
frames, err := decodeReferenceImages(in.ReferenceImages, max(1, modelItem.MaxReferenceImages))
|
||||||
|
if err != nil {
|
||||||
|
return nil, "", err
|
||||||
|
}
|
||||||
var lastErr error
|
var lastErr error
|
||||||
var videoURL string
|
var videoURL string
|
||||||
busy := 0
|
busy := 0
|
||||||
@@ -2003,7 +2009,7 @@ func (s *V1Service) generateCustomVideo(ctx context.Context, eventID string, mod
|
|||||||
_ = s.events.SetAccount(ctx, eventID, token.ID, token.AccountEmail)
|
_ = s.events.SetAccount(ctx, eventID, token.ID, token.AccountEmail)
|
||||||
_ = s.tokens.TouchLastUsed(ctx, token.ID)
|
_ = s.tokens.TouchLastUsed(ctx, token.ID)
|
||||||
baseURL := stringValue(token.Meta["base_url"])
|
baseURL := stringValue(token.Meta["base_url"])
|
||||||
d, url, genErr := s.custom.GenerateVideo(ctx, baseURL, token.Value, modelItem.ID, in.Prompt, size, durationSeconds, downloadResult)
|
d, url, genErr := s.custom.GenerateVideo(ctx, baseURL, token.Value, modelItem.ID, in.Prompt, size, durationSeconds, frames, downloadResult)
|
||||||
if genErr == nil {
|
if genErr == nil {
|
||||||
_, _ = s.tokens.Update(ctx, "custom", token.ID, map[string]any{
|
_, _ = s.tokens.Update(ctx, "custom", token.ID, map[string]any{
|
||||||
"last_used_at": time.Now(), "success_total": gorm.Expr("success_total + 1"), "fails": 0,
|
"last_used_at": time.Now(), "success_total": gorm.Expr("success_total + 1"), "fails": 0,
|
||||||
|
|||||||
Reference in New Issue
Block a user