jytest/client/client.go

134 lines
3.3 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

package client
import (
"bytes"
"encoding/json"
"fmt"
"io"
"log"
"net/http"
"time"
"jimeng/signer"
)
// Client 即梦API客户端
type Client struct {
signer *signer.Config
hc *http.Client
}
// NewClient 创建客户端
func NewClient(ak, sk, host string) *Client {
return &Client{
signer: &signer.Config{
AccessKey: ak,
SecretKey: sk,
Host: host,
},
hc: &http.Client{Timeout: 30 * time.Second},
}
}
// post 统一发送带签名的 POST 请求
func (c *Client) post(path string, body any, out any) error {
b, err := json.Marshal(body)
if err != nil {
return err
}
log.Printf("[HTTP] POST %s Body=%s", path, string(b))
hdrs, err := c.signer.Headers(&signer.Request{
Method: "POST",
Path: path,
Body: b,
ContentType: "application/json",
})
if err != nil {
return err
}
req, err := http.NewRequest("POST", c.signer.Host+path, bytes.NewReader(b))
if err != nil {
return err
}
req.Header.Set("Content-Type", "application/json")
for k, v := range hdrs {
req.Header.Set(k, v)
}
resp, err := c.hc.Do(req)
if err != nil {
return err
}
defer resp.Body.Close()
rb, _ := io.ReadAll(resp.Body)
log.Printf("[HTTP] Response status=%d body=%s", resp.StatusCode, string(rb))
if err := json.Unmarshal(rb, out); err != nil {
return fmt.Errorf("decode resp: %w, body=%s", err, string(rb))
}
return nil
}
// SubmitNovelTask 提交短剧任务
func (c *Client) SubmitNovelTask(input map[string]interface{}) (map[string]interface{}, error) {
req := map[string]interface{}{
"capability_key": "pippit_novel_agent",
"action_key": input["action_key"],
"run_id": fmt.Sprintf("biz-%d", time.Now().UnixNano()),
"input": input,
}
var out map[string]interface{}
if err := c.post("/agent_openapi/v1/novel/submit", req, &out); err != nil {
return nil, err
}
return out, nil
}
// QueryNovelTask 查询短剧任务
func (c *Client) QueryNovelTask(taskID, runID string) (map[string]interface{}, error) {
body := map[string]string{}
if taskID != "" {
body["task_id"] = taskID
} else {
body["run_id"] = runID
}
var out map[string]interface{}
if err := c.post("/agent_openapi/v1/novel/query", body, &out); err != nil {
return nil, err
}
return out, nil
}
// SubmitVideoTask 提交视频任务(带货营销/剧情营销/短片创作)
func (c *Client) SubmitVideoTask(capabilityKey string, input map[string]interface{}) (map[string]interface{}, error) {
req := map[string]interface{}{
"capability_key": capabilityKey,
"action_key": "generate_video",
"run_id": fmt.Sprintf("biz-%d", time.Now().UnixNano()),
}
// 带货营销和剧情营销使用 avatar_marketing_input,短片创作使用 input
if capabilityKey == "pippit_video_part_agent" {
req["input"] = input
} else {
req["avatar_marketing_input"] = input
}
var out map[string]interface{}
if err := c.post("/agent_openapi/v1/video/submit", req, &out); err != nil {
return nil, err
}
return out, nil
}
// QueryVideoTask 查询视频任务
func (c *Client) QueryVideoTask(taskID, runID string) (map[string]interface{}, error) {
body := map[string]string{}
if taskID != "" {
body["task_id"] = taskID
} else {
body["run_id"] = runID
}
var out map[string]interface{}
if err := c.post("/agent_openapi/v1/video/query", body, &out); err != nil {
return nil, err
}
return out, nil
}