296 lines
6.7 KiB
Go
296 lines
6.7 KiB
Go
package handler
|
|
|
|
import (
|
|
"archive/zip"
|
|
"fmt"
|
|
"io"
|
|
"os"
|
|
"path/filepath"
|
|
"sdk-generator/internal/call"
|
|
"sdk-generator/internal/models"
|
|
"sdk-generator/internal/service"
|
|
"strconv"
|
|
"strings"
|
|
"time"
|
|
|
|
"github.com/gofiber/fiber/v2"
|
|
)
|
|
|
|
type SDKHandler struct {
|
|
service *service.SDKGeneratorService
|
|
}
|
|
|
|
func NewSDKHandler(service *service.SDKGeneratorService) *SDKHandler {
|
|
return &SDKHandler{service: service}
|
|
}
|
|
|
|
// GenerateSDK 上传文档并生成 SDK
|
|
// POST /api/v1/generate
|
|
func (h *SDKHandler) GenerateSDK(c *fiber.Ctx) error {
|
|
//获取上传文件
|
|
file, err := c.FormFile("file")
|
|
if err != nil {
|
|
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
|
"error": "请上传文件 (field: file)",
|
|
})
|
|
}
|
|
|
|
// 检查文件大小
|
|
if file.Size > 10*1024*1024 {
|
|
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
|
"error": "文件大小超过 10MB 限制",
|
|
})
|
|
}
|
|
var content []byte
|
|
// 检查文件类型
|
|
filename := file.Filename
|
|
ext := strings.ToLower(filepath.Ext(filename))
|
|
src, err := file.Open()
|
|
if err != nil {
|
|
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
|
"error": "打开文件失败",
|
|
})
|
|
}
|
|
llmSet := call.LlmCallSet{
|
|
ApiKey: c.FormValue("llm_api_key", ""),
|
|
BaseUrl: c.FormValue("llm_base_url", ""),
|
|
ModelName: c.FormValue("llm_model", ""),
|
|
}
|
|
if llmSet.ModelName == "" || llmSet.ApiKey == "" || llmSet.BaseUrl == "" {
|
|
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
|
"error": "缺少 llm_model 或 llm_api_key 或 llm_base_url",
|
|
})
|
|
}
|
|
defer src.Close()
|
|
switch ext {
|
|
case ".md", ".txt", ".markdown":
|
|
// 读取文件内容
|
|
content, err = io.ReadAll(src)
|
|
if err != nil {
|
|
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
|
"error": "读取文件失败",
|
|
})
|
|
}
|
|
break
|
|
case ".docx", ".doc", ".pdf":
|
|
res, err := ConvertFile(ConvertRequest{
|
|
File: src,
|
|
FileName: filename,
|
|
LlmModel: llmSet.ModelName,
|
|
LlmApiKey: llmSet.ApiKey,
|
|
LlmBaseUrl: llmSet.BaseUrl,
|
|
//LlmPrompt: c.FormValue("llm_prompt", ""),
|
|
})
|
|
if err != nil {
|
|
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
|
"error": "读取文件失败" + err.Error(),
|
|
})
|
|
}
|
|
content = []byte(res.Markdown)
|
|
break
|
|
default:
|
|
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{"error": "不支持的文件类型"})
|
|
}
|
|
|
|
sdkName := c.FormValue("sdk_name", time.Now().Format(time.RFC3339))
|
|
// 调用服务生成
|
|
taskID, err := h.service.Generate(c.Context(), content, sdkName, &llmSet)
|
|
if err != nil {
|
|
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
|
"error": err.Error(),
|
|
})
|
|
}
|
|
|
|
return c.Status(fiber.StatusAccepted).JSON(models.GenerateResponse{
|
|
TaskID: taskID,
|
|
Status: string(models.StatusPending),
|
|
Message: "SDK 生成任务已提交,请通过 /api/v1/tasks/:task_id 查询状态",
|
|
})
|
|
}
|
|
|
|
// GetTaskStatus 查询任务状态
|
|
// GET /api/v1/tasks/:task_id
|
|
func (h *SDKHandler) GetTaskStatus(c *fiber.Ctx) error {
|
|
taskID := c.Params("task_id")
|
|
if taskID == "" {
|
|
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
|
"error": "缺少 task_id",
|
|
})
|
|
}
|
|
|
|
task, ok := h.service.GetTask(taskID)
|
|
if !ok {
|
|
return c.Status(fiber.StatusNotFound).JSON(fiber.Map{
|
|
"error": "任务不存在",
|
|
})
|
|
}
|
|
task.StatusName = task.Status.Desc()
|
|
task.Percent = task.Status.Percent()
|
|
response := models.TaskResponse{
|
|
Task: *task,
|
|
}
|
|
|
|
// 如果任务完成,列出生成的文件
|
|
if task.Status == models.StatusCompleted {
|
|
files, _ := listFiles(task.OutputDir)
|
|
response.Files = files
|
|
}
|
|
|
|
return c.JSON(response)
|
|
}
|
|
|
|
// DownloadSDK 下载生成的 SDK
|
|
// GET /api/v1/tasks/:task_id/download
|
|
func (h *SDKHandler) DownloadSDK(c *fiber.Ctx) error {
|
|
taskID := c.Params("task_id")
|
|
if taskID == "" {
|
|
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
|
"error": "缺少 task_id",
|
|
})
|
|
}
|
|
|
|
task, ok := h.service.GetTask(taskID)
|
|
if !ok {
|
|
return c.Status(fiber.StatusNotFound).JSON(fiber.Map{
|
|
"error": "任务不存在",
|
|
})
|
|
}
|
|
|
|
if task.Status != models.StatusCompleted {
|
|
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
|
"error": fmt.Sprintf("任务未完成,当前状态: %s", task.Status),
|
|
})
|
|
}
|
|
|
|
// 检查输出目录是否存在
|
|
if _, err := os.Stat(task.OutputDir); os.IsNotExist(err) {
|
|
return c.Status(fiber.StatusNotFound).JSON(fiber.Map{
|
|
"error": "SDK 文件不存在",
|
|
})
|
|
}
|
|
|
|
// 创建 ZIP 文件并流式返回
|
|
zipFilename := fmt.Sprintf("sdk_%s.zip", taskID)
|
|
c.Set("Content-Type", "application/zip")
|
|
c.Set("Content-Disposition", fmt.Sprintf("attachment; filename=%s", zipFilename))
|
|
|
|
// 使用 Fiber 的 Response 直接写入
|
|
zipWriter := zip.NewWriter(c.Response().BodyWriter())
|
|
defer zipWriter.Close()
|
|
|
|
// 遍历目录添加文件到 ZIP
|
|
err := filepath.Walk(task.OutputDir, func(path string, info os.FileInfo, err error) error {
|
|
if err != nil {
|
|
return err
|
|
}
|
|
|
|
if info.IsDir() {
|
|
return nil
|
|
}
|
|
|
|
// 计算相对路径
|
|
relPath, err := filepath.Rel(task.OutputDir, path)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
|
|
// 创建 ZIP 条目
|
|
zipFile, err := zipWriter.Create(relPath)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
|
|
// 读取文件内容
|
|
content, err := os.ReadFile(path)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
|
|
_, err = zipFile.Write(content)
|
|
return err
|
|
})
|
|
|
|
if err != nil {
|
|
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
|
"error": fmt.Sprintf("打包 SDK 失败: %v", err),
|
|
})
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
// ListTasks 列出所有任务
|
|
// GET /api/v1/tasks
|
|
func (h *SDKHandler) ListTasks(c *fiber.Ctx) error {
|
|
tasks := h.service.ListTasks()
|
|
|
|
// 分页
|
|
page, _ := strconv.Atoi(c.Query("page", "1"))
|
|
size, _ := strconv.Atoi(c.Query("size", "20"))
|
|
|
|
if page < 1 {
|
|
page = 1
|
|
}
|
|
if size < 1 {
|
|
size = 20
|
|
}
|
|
if size > 100 {
|
|
size = 100
|
|
}
|
|
|
|
total := len(tasks)
|
|
start := (page - 1) * size
|
|
end := start + size
|
|
|
|
if start > total {
|
|
start = total
|
|
}
|
|
if end > total {
|
|
end = total
|
|
}
|
|
|
|
return c.JSON(models.ListTasksResponse{
|
|
Data: tasks[start:end],
|
|
Total: total,
|
|
Page: page,
|
|
Size: size,
|
|
})
|
|
}
|
|
|
|
// DeleteTask 删除任务
|
|
// DELETE /api/v1/tasks/:task_id
|
|
func (h *SDKHandler) DeleteTask(c *fiber.Ctx) error {
|
|
taskID := c.Params("task_id")
|
|
if taskID == "" {
|
|
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
|
"error": "缺少 task_id",
|
|
})
|
|
}
|
|
|
|
if err := h.service.DeleteTask(taskID); err != nil {
|
|
return c.Status(fiber.StatusNotFound).JSON(fiber.Map{
|
|
"error": err.Error(),
|
|
})
|
|
}
|
|
|
|
return c.JSON(fiber.Map{
|
|
"message": "任务已删除",
|
|
})
|
|
}
|
|
|
|
// 辅助函数:列出目录下的文件
|
|
func listFiles(dir string) ([]string, error) {
|
|
var files []string
|
|
err := filepath.Walk(dir, func(path string, info os.FileInfo, err error) error {
|
|
if err != nil {
|
|
return err
|
|
}
|
|
if !info.IsDir() {
|
|
relPath, _ := filepath.Rel(dir, path)
|
|
files = append(files, relPath)
|
|
}
|
|
return nil
|
|
})
|
|
return files, err
|
|
}
|