252 lines
6.4 KiB
Go
252 lines
6.4 KiB
Go
package service
|
|
|
|
import (
|
|
"io"
|
|
"log"
|
|
"path/filepath"
|
|
"sdk-generator/internal/biz"
|
|
"sdk-generator/internal/data/model"
|
|
"sdk-generator/internal/entitys"
|
|
|
|
"sdk-generator/internal/pkg/call"
|
|
"sdk-generator/internal/pkg/response"
|
|
"sdk-generator/tmpl/errcode"
|
|
"strings"
|
|
|
|
"github.com/gofiber/fiber/v2"
|
|
"xorm.io/builder"
|
|
)
|
|
|
|
type SDKService struct {
|
|
generate *biz.SDKGeneratorBiz
|
|
}
|
|
|
|
func NewSDKService(service *biz.SDKGeneratorBiz) *SDKService {
|
|
return &SDKService{generate: service}
|
|
}
|
|
|
|
func (h *SDKService) Refine(c *fiber.Ctx, req *entitys.RefineRequest) 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": "打开文件失败",
|
|
})
|
|
}
|
|
// Debug: log received llm api key prefix and incoming Authorization header
|
|
if req != nil {
|
|
if req.LlmApiKey != "" {
|
|
log.Printf("Refine handler: received llm_api_key prefix=%s", req.LlmApiKey[:minLen(len(req.LlmApiKey), 8)])
|
|
}
|
|
}
|
|
incomingAuth := c.Get("Authorization")
|
|
if incomingAuth != "" {
|
|
log.Printf("Refine handler: incoming Authorization header prefix=%s", incomingAuth[:minLen(len(incomingAuth), 8)])
|
|
}
|
|
|
|
llmSet := call.LlmCallSet{
|
|
ApiKey: req.LlmApiKey,
|
|
BaseUrl: req.LlmBaseUrl,
|
|
ModelName: req.LlmModel,
|
|
}
|
|
defer func() { _ = src.Close() }()
|
|
switch ext {
|
|
case ".md", ".txt", ".markdown":
|
|
// 读取文件内容
|
|
content, err = io.ReadAll(src)
|
|
if err != nil {
|
|
return errcode.ParamErr("读取文件失败")
|
|
}
|
|
break
|
|
case ".docx", ".doc", ".pdf":
|
|
res, err := biz.ConvertFile(biz.ConvertRequest{
|
|
File: src,
|
|
FileName: filename,
|
|
LlmModel: llmSet.ModelName,
|
|
LlmApiKey: llmSet.ApiKey,
|
|
LlmBaseUrl: llmSet.BaseUrl,
|
|
LlmPrompt: req.LlmPrompt,
|
|
})
|
|
if err != nil {
|
|
return errcode.ParamErr("读取文件失败" + err.Error())
|
|
}
|
|
content = []byte(res.Markdown)
|
|
break
|
|
default:
|
|
return errcode.ParamErr("不支持的文件类型")
|
|
}
|
|
desc := req.Desc
|
|
if desc == "" {
|
|
desc = filename
|
|
}
|
|
res, err := h.generate.RefineDoc(c.UserContext(), content, &llmSet, desc)
|
|
if err != nil {
|
|
return errcode.BadReq("文档分析失败:" + err.Error())
|
|
}
|
|
return response.HandleResponse(c, res)
|
|
}
|
|
|
|
// ListRefine
|
|
func (h *SDKService) ListRefine(c *fiber.Ctx, req *entitys.ListDocRequest) error {
|
|
cond := builder.NewCond()
|
|
cond = cond.And(builder.IsNull{"deleted_at"})
|
|
if req.InstanceId != "" {
|
|
cond = cond.And(builder.Eq{"instance_id": req.InstanceId})
|
|
}
|
|
if req.Desc != "" {
|
|
cond = cond.And(builder.Like{"`desc`", req.Desc})
|
|
}
|
|
if req.Page < 1 {
|
|
req.Page = 1
|
|
}
|
|
if req.PageSize < 1 {
|
|
req.PageSize = 20
|
|
}
|
|
|
|
docs, total, err := h.generate.ListDoc(c.UserContext(), &cond, req.Page, req.PageSize)
|
|
if err != nil {
|
|
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
|
"error": err.Error(),
|
|
})
|
|
}
|
|
|
|
return response.HandleResponse(c, entitys.ListDocResponse{
|
|
Data: docs,
|
|
Total: total,
|
|
Page: req.Page,
|
|
Size: req.PageSize,
|
|
})
|
|
|
|
}
|
|
|
|
// ListRefine
|
|
func (h *SDKService) UpdateRefine(c *fiber.Ctx, req *entitys.UpdateDocRequest) error {
|
|
|
|
err := h.generate.UpdateDocByInstanceId(c.UserContext(), &model.AiGenerateDoc{
|
|
InstanceID: req.InstanceId,
|
|
RefinedDoc: req.RefinedDoc,
|
|
Desc: req.Desc,
|
|
})
|
|
if err != nil {
|
|
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
|
"error": err.Error(),
|
|
})
|
|
}
|
|
|
|
return response.HandleResponse(c, "ok")
|
|
}
|
|
|
|
// Generate 上传文档并生成 SDK
|
|
func (h *SDKService) Generate(c *fiber.Ctx, req *entitys.GenerateRequest) error {
|
|
// Debug: log received llm api key prefix and incoming Authorization header
|
|
if req != nil {
|
|
if req.LlmApiKey != "" {
|
|
log.Printf("Generate handler: received llm_api_key prefix=%s", req.LlmApiKey[:minLen(len(req.LlmApiKey), 8)])
|
|
}
|
|
}
|
|
incomingAuth := c.Get("Authorization")
|
|
if incomingAuth != "" {
|
|
log.Printf("Generate handler: incoming Authorization header prefix=%s", incomingAuth[:minLen(len(incomingAuth), 8)])
|
|
}
|
|
|
|
llmSet := call.LlmCallSet{
|
|
ApiKey: req.LlmApiKey,
|
|
BaseUrl: req.LlmBaseUrl,
|
|
ModelName: req.LlmModel,
|
|
}
|
|
// 调用服务生成
|
|
taskID, err := h.generate.Generate(c.Context(), req, &llmSet)
|
|
if err != nil {
|
|
return errcode.BadReq("任务提交失败:" + err.Error())
|
|
}
|
|
return response.HandleResponse(c, &entitys.GenerateResponse{
|
|
TaskID: taskID,
|
|
Status: string(entitys.StatusPending),
|
|
Message: "SDK 生成任务已提交,请通过 /api/v1/tasks/" + taskID + " 查询状态",
|
|
})
|
|
}
|
|
|
|
// GetTaskStatus 查询任务状态
|
|
// GET /api/v1/tasks/:task_id
|
|
func (h *SDKService) GetTaskStatus(c *fiber.Ctx) error {
|
|
taskID := c.Params("task_id")
|
|
if taskID == "" {
|
|
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
|
"error": "缺少 task_id",
|
|
})
|
|
}
|
|
var resp entitys.TaskResponse
|
|
task, ok := h.generate.GetTask(c.UserContext(), taskID)
|
|
if !ok {
|
|
return c.Status(fiber.StatusNotFound).JSON(fiber.Map{
|
|
"error": "任务不存在",
|
|
})
|
|
}
|
|
resp.Task = task
|
|
resp.StatusDesc = entitys.TaskStatus(task.TaskStatus).Desc()
|
|
resp.Percent = entitys.TaskStatus(task.TaskStatus).Percent()
|
|
|
|
return response.HandleResponse(c, resp)
|
|
}
|
|
|
|
// ListTasks 列出所有任务
|
|
// GET /api/v1/tasks
|
|
func (h *SDKService) ListTasks(c *fiber.Ctx, req *entitys.ListTaskRequest) error {
|
|
|
|
cond := builder.NewCond()
|
|
cond = cond.And(builder.IsNull{"deleted_at"})
|
|
if req.TaskId != "" {
|
|
cond = cond.And(builder.Like{"task_id", req.TaskId})
|
|
}
|
|
if req.Desc != "" {
|
|
cond = cond.And(builder.Like{"`desc`", req.Desc})
|
|
}
|
|
if req.Page < 1 {
|
|
req.Page = 1
|
|
}
|
|
if req.PageSize < 1 {
|
|
req.PageSize = 20
|
|
}
|
|
tasks, total, err := h.generate.ListTasks(c.UserContext(), &cond, req.Page, req.PageSize)
|
|
if err != nil {
|
|
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
|
"error": err.Error(),
|
|
})
|
|
}
|
|
for i := range tasks {
|
|
tasks[i].TaskStatus = entitys.TaskStatus(tasks[i].TaskStatus).Desc()
|
|
}
|
|
return c.JSON(entitys.ListTasksResponse{
|
|
Data: tasks,
|
|
Total: total,
|
|
Page: req.Page,
|
|
Size: req.PageSize,
|
|
})
|
|
}
|
|
|
|
// minLen returns the smaller of a and b
|
|
func minLen(a, b int) int {
|
|
if a < b {
|
|
return a
|
|
}
|
|
return b
|
|
}
|