296 lines
7.9 KiB
Go
296 lines
7.9 KiB
Go
package handler
|
|
|
|
import (
|
|
"encoding/json"
|
|
"net/http"
|
|
"time"
|
|
|
|
"jimeng/client"
|
|
"jimeng/config"
|
|
"jimeng/model"
|
|
"jimeng/store"
|
|
)
|
|
|
|
// Handler HTTP处理器
|
|
type Handler struct {
|
|
client *client.Client
|
|
store *store.Store
|
|
cfg *config.Config
|
|
}
|
|
|
|
// NewHandler 创建处理器
|
|
func NewHandler(c *client.Client, s *store.Store, cfg *config.Config) *Handler {
|
|
return &Handler{client: c, store: s, cfg: cfg}
|
|
}
|
|
|
|
// RegisterRoutes 注册路由
|
|
func (h *Handler) RegisterRoutes(mux *http.ServeMux) {
|
|
mux.HandleFunc("/api/novel/submit", h.submitNovel)
|
|
mux.HandleFunc("/api/novel/query", h.queryNovel)
|
|
mux.HandleFunc("/api/video/submit", h.submitVideo)
|
|
mux.HandleFunc("/api/video/query", h.queryVideo)
|
|
mux.HandleFunc("/api/tasks", h.listTasks)
|
|
mux.HandleFunc("/api/task", h.getTask)
|
|
mux.HandleFunc("/api/task/delete", h.deleteTask)
|
|
mux.HandleFunc("/api/testdata", h.getTestData)
|
|
}
|
|
|
|
func (h *Handler) submitNovel(w http.ResponseWriter, r *http.Request) {
|
|
if r.Method != http.MethodPost {
|
|
http.Error(w, "Method not allowed", http.StatusMethodNotAllowed)
|
|
return
|
|
}
|
|
var input map[string]interface{}
|
|
if err := json.NewDecoder(r.Body).Decode(&input); err != nil {
|
|
http.Error(w, err.Error(), http.StatusBadRequest)
|
|
return
|
|
}
|
|
result, err := h.client.SubmitNovelTask(input)
|
|
if err != nil {
|
|
http.Error(w, err.Error(), http.StatusInternalServerError)
|
|
return
|
|
}
|
|
// 创建本地任务记录
|
|
task := &model.Task{
|
|
ID: store.GenerateID(),
|
|
Type: "novel",
|
|
Capability: "pippit_novel_agent",
|
|
Action: getString(input, "action_key"),
|
|
Status: "CREATED",
|
|
ThreadID: getString(input, "thread_id"),
|
|
AssetID: getString(input, "asset_id"),
|
|
RunID: getString(result, "run_id"),
|
|
Input: input,
|
|
CreatedAt: time.Now(),
|
|
UpdatedAt: time.Now(),
|
|
}
|
|
if data, ok := result["data"].(map[string]interface{}); ok {
|
|
task.Status = getString(data, "status")
|
|
if task.RunID == "" {
|
|
task.RunID = getString(data, "run_id")
|
|
}
|
|
if task.ThreadID == "" {
|
|
task.ThreadID = getString(data, "thread_id")
|
|
}
|
|
}
|
|
if err := h.store.CreateTask(task); err != nil {
|
|
http.Error(w, err.Error(), http.StatusInternalServerError)
|
|
return
|
|
}
|
|
// 将本地 task_id 和 run_id 注入响应
|
|
if data, ok := result["data"].(map[string]interface{}); ok {
|
|
if getString(data, "task_id") == "" {
|
|
data["task_id"] = task.ID
|
|
}
|
|
if getString(data, "run_id") == "" {
|
|
data["run_id"] = task.RunID
|
|
}
|
|
}
|
|
json.NewEncoder(w).Encode(result)
|
|
}
|
|
|
|
func (h *Handler) queryNovel(w http.ResponseWriter, r *http.Request) {
|
|
if r.Method != http.MethodPost {
|
|
http.Error(w, "Method not allowed", http.StatusMethodNotAllowed)
|
|
return
|
|
}
|
|
var req struct {
|
|
TaskID string `json:"task_id"`
|
|
RunID string `json:"run_id"`
|
|
}
|
|
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
|
http.Error(w, err.Error(), http.StatusBadRequest)
|
|
return
|
|
}
|
|
result, err := h.client.QueryNovelTask(req.TaskID, req.RunID)
|
|
if err != nil {
|
|
http.Error(w, err.Error(), http.StatusInternalServerError)
|
|
return
|
|
}
|
|
// 更新本地任务状态
|
|
if data, ok := result["data"].(map[string]interface{}); ok {
|
|
taskID := req.TaskID
|
|
if taskID == "" {
|
|
// 通过run_id查找任务
|
|
tasks := h.store.ListTasks()
|
|
for _, t := range tasks {
|
|
if t.RunID == req.RunID {
|
|
taskID = t.ID
|
|
break
|
|
}
|
|
}
|
|
}
|
|
if task, ok := h.store.GetTask(taskID); ok {
|
|
task.Status = getString(data, "status")
|
|
task.ThreadID = getString(data, "thread_id")
|
|
task.UpdatedAt = time.Now()
|
|
if usage, ok := data["usage"].(map[string]interface{}); ok {
|
|
task.Usage = usage
|
|
}
|
|
if novelData, ok := data["novel_data"].(map[string]interface{}); ok {
|
|
task.Output = novelData
|
|
}
|
|
h.store.UpdateTask(task)
|
|
}
|
|
}
|
|
json.NewEncoder(w).Encode(result)
|
|
}
|
|
|
|
func (h *Handler) submitVideo(w http.ResponseWriter, r *http.Request) {
|
|
if r.Method != http.MethodPost {
|
|
http.Error(w, "Method not allowed", http.StatusMethodNotAllowed)
|
|
return
|
|
}
|
|
var req struct {
|
|
CapabilityKey string `json:"capability_key"`
|
|
Input map[string]interface{} `json:"input"`
|
|
}
|
|
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
|
http.Error(w, err.Error(), http.StatusBadRequest)
|
|
return
|
|
}
|
|
result, err := h.client.SubmitVideoTask(req.CapabilityKey, req.Input)
|
|
if err != nil {
|
|
http.Error(w, err.Error(), http.StatusInternalServerError)
|
|
return
|
|
}
|
|
// 创建本地任务记录
|
|
task := &model.Task{
|
|
ID: store.GenerateID(),
|
|
Type: "video",
|
|
Capability: req.CapabilityKey,
|
|
Action: "generate_video",
|
|
Status: "CREATED",
|
|
RunID: getString(result, "run_id"),
|
|
Input: req.Input,
|
|
CreatedAt: time.Now(),
|
|
UpdatedAt: time.Now(),
|
|
}
|
|
if data, ok := result["data"].(map[string]interface{}); ok {
|
|
task.Status = getString(data, "status")
|
|
if task.RunID == "" {
|
|
task.RunID = getString(data, "run_id")
|
|
}
|
|
}
|
|
if err := h.store.CreateTask(task); err != nil {
|
|
http.Error(w, err.Error(), http.StatusInternalServerError)
|
|
return
|
|
}
|
|
// 将本地 task_id 和 run_id 注入响应,确保前端能拿到
|
|
if data, ok := result["data"].(map[string]interface{}); ok {
|
|
if getString(data, "task_id") == "" {
|
|
data["task_id"] = task.ID
|
|
}
|
|
if getString(data, "run_id") == "" {
|
|
data["run_id"] = task.RunID
|
|
}
|
|
}
|
|
json.NewEncoder(w).Encode(result)
|
|
}
|
|
|
|
func (h *Handler) queryVideo(w http.ResponseWriter, r *http.Request) {
|
|
if r.Method != http.MethodPost {
|
|
http.Error(w, "Method not allowed", http.StatusMethodNotAllowed)
|
|
return
|
|
}
|
|
var req struct {
|
|
TaskID string `json:"task_id"`
|
|
RunID string `json:"run_id"`
|
|
}
|
|
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
|
http.Error(w, err.Error(), http.StatusBadRequest)
|
|
return
|
|
}
|
|
result, err := h.client.QueryVideoTask(req.TaskID, req.RunID)
|
|
if err != nil {
|
|
http.Error(w, err.Error(), http.StatusInternalServerError)
|
|
return
|
|
}
|
|
// 更新本地任务状态
|
|
if data, ok := result["data"].(map[string]interface{}); ok {
|
|
taskID := req.TaskID
|
|
if taskID == "" {
|
|
tasks := h.store.ListTasks()
|
|
for _, t := range tasks {
|
|
if t.RunID == req.RunID {
|
|
taskID = t.ID
|
|
break
|
|
}
|
|
}
|
|
}
|
|
if task, ok := h.store.GetTask(taskID); ok {
|
|
task.Status = getString(data, "status")
|
|
task.UpdatedAt = time.Now()
|
|
if usage, ok := data["usage"].(map[string]interface{}); ok {
|
|
task.Usage = usage
|
|
}
|
|
if artifact, ok := data["video_artifact"].(map[string]interface{}); ok {
|
|
task.Output = artifact
|
|
}
|
|
if errCode, ok := data["err_code"].(string); ok {
|
|
task.ErrorCode = errCode
|
|
}
|
|
if errMsg, ok := data["err_msg"].(string); ok {
|
|
task.ErrorMsg = errMsg
|
|
}
|
|
h.store.UpdateTask(task)
|
|
}
|
|
}
|
|
json.NewEncoder(w).Encode(result)
|
|
}
|
|
|
|
func (h *Handler) listTasks(w http.ResponseWriter, r *http.Request) {
|
|
tasks := h.store.ListTasks()
|
|
json.NewEncoder(w).Encode(tasks)
|
|
}
|
|
|
|
func (h *Handler) getTask(w http.ResponseWriter, r *http.Request) {
|
|
id := r.URL.Query().Get("id")
|
|
if id == "" {
|
|
http.Error(w, "id required", http.StatusBadRequest)
|
|
return
|
|
}
|
|
task, ok := h.store.GetTask(id)
|
|
if !ok {
|
|
http.Error(w, "task not found", http.StatusNotFound)
|
|
return
|
|
}
|
|
json.NewEncoder(w).Encode(task)
|
|
}
|
|
|
|
func (h *Handler) deleteTask(w http.ResponseWriter, r *http.Request) {
|
|
if r.Method != http.MethodPost {
|
|
http.Error(w, "Method not allowed", http.StatusMethodNotAllowed)
|
|
return
|
|
}
|
|
var req struct {
|
|
ID string `json:"id"`
|
|
}
|
|
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
|
http.Error(w, err.Error(), http.StatusBadRequest)
|
|
return
|
|
}
|
|
if err := h.store.DeleteTask(req.ID); err != nil {
|
|
http.Error(w, err.Error(), http.StatusInternalServerError)
|
|
return
|
|
}
|
|
w.Write([]byte(`{"success":true}`))
|
|
}
|
|
|
|
func getString(m map[string]interface{}, key string) string {
|
|
if v, ok := m[key]; ok {
|
|
if s, ok := v.(string); ok {
|
|
return s
|
|
}
|
|
}
|
|
return ""
|
|
}
|
|
|
|
func (h *Handler) getTestData(w http.ResponseWriter, r *http.Request) {
|
|
resp := map[string]interface{}{
|
|
"test_product": h.cfg.TestProduct,
|
|
"test_model": h.cfg.TestModel,
|
|
}
|
|
json.NewEncoder(w).Encode(resp)
|
|
}
|