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