jytest/handler/handler.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)
}