sdk_generate/internal/handler/sdk.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
}