559 lines
16 KiB
Go
559 lines
16 KiB
Go
package biz
|
||
|
||
import (
|
||
"context"
|
||
"encoding/json"
|
||
"errors"
|
||
"fmt"
|
||
"log"
|
||
"os"
|
||
"path/filepath"
|
||
"sdk-generator/internal/config"
|
||
"sdk-generator/internal/data/file"
|
||
"sdk-generator/internal/data/impl"
|
||
"sdk-generator/internal/data/model"
|
||
"sdk-generator/internal/entitys"
|
||
"sdk-generator/internal/pkg"
|
||
"sdk-generator/internal/pkg/call"
|
||
"sdk-generator/internal/pkg/extractor"
|
||
"sdk-generator/internal/pkg/postprocess"
|
||
"sdk-generator/internal/pkg/push"
|
||
"sdk-generator/internal/prompts"
|
||
"sdk-generator/tmpl/dataTemp"
|
||
"sdk-generator/tmpl/errcode"
|
||
"strings"
|
||
|
||
"time"
|
||
|
||
"github.com/google/uuid"
|
||
"xorm.io/builder"
|
||
)
|
||
|
||
type SDKGeneratorBiz struct {
|
||
config *config.Config
|
||
docImpl *impl.AiGenerateDocImpl
|
||
taskImpl *impl.AiGenerateTaskImpl
|
||
}
|
||
|
||
func NewSDKGeneratorService(cfg *config.Config, docImpl *impl.AiGenerateDocImpl, taskImpl *impl.AiGenerateTaskImpl) *SDKGeneratorBiz {
|
||
// 创建 OpenAI 客户端配置
|
||
return &SDKGeneratorBiz{
|
||
config: cfg,
|
||
docImpl: docImpl,
|
||
taskImpl: taskImpl,
|
||
}
|
||
}
|
||
|
||
func (s *SDKGeneratorBiz) RefineDoc(ctx context.Context, content []byte, callSet *call.LlmCallSet, docDesc string) (*entitys.RefineResponse, error) {
|
||
|
||
if content == nil || len(content) == 0 {
|
||
return nil, errors.New("内容为空")
|
||
}
|
||
if callSet == nil {
|
||
return nil, errors.New("调用参数为空")
|
||
}
|
||
llmClient := call.NewCallLLM(callSet)
|
||
// 1. 分析文档
|
||
log.Println("开始分析文档")
|
||
res, err := llmClient.DoWithAll(ctx, prompts.BuildRefinePrompt(string(content)))
|
||
if err != nil {
|
||
return nil, fmt.Errorf("调用大模型失败: %v", err)
|
||
}
|
||
toolRes, exist := prompts.ExtractDocTypeFromToolCall(res)
|
||
if !exist {
|
||
return nil, errors.New("提取文档类型失败")
|
||
}
|
||
uuidStr := uuid.New().String()
|
||
s.docImpl.Add(ctx, &model.AiGenerateDoc{
|
||
Desc: docDesc,
|
||
InstanceID: uuidStr,
|
||
Interface: pkg.JsonStringIgonErr(toolRes.Interfaces),
|
||
RefinedDoc: toolRes.RefinedDoc,
|
||
CreatedAt: time.Now(),
|
||
})
|
||
return &entitys.RefineResponse{
|
||
InstanceId: uuidStr,
|
||
RefinedDoc: toolRes.RefinedDoc,
|
||
Interfaces: toolRes.Interfaces,
|
||
}, err
|
||
}
|
||
|
||
// Generate 异步生成 SDK
|
||
func (s *SDKGeneratorBiz) Generate(ctx context.Context, req *entitys.GenerateRequest, llmSet *call.LlmCallSet) (string, error) {
|
||
var doc model.AiGenerateDoc
|
||
|
||
err := s.docImpl.GetByKey(ctx, "instance_id", req.InstanceId, &doc)
|
||
if err != nil {
|
||
return "", errcode.NotFound("文档不存在")
|
||
}
|
||
if req.RefinedDoc == "" {
|
||
req.RefinedDoc = doc.RefinedDoc
|
||
}
|
||
|
||
// 限制任务数量
|
||
taskID := req.InstanceId + time.Now().Format("20060102150405")
|
||
task := entitys.Task{
|
||
TasKModel: &model.AiGenerateTask{
|
||
TaskID: req.InstanceId + "_" + time.Now().Format("20060102150405"),
|
||
CodeType: req.CodeType,
|
||
Interfaces: req.Interfaces,
|
||
GenerateType: req.GenerateType,
|
||
CreatedAt: time.Now(),
|
||
LlmModel: req.LlmModel,
|
||
TaskStatus: entitys.StatusPending.String(),
|
||
Desc: req.Desc,
|
||
},
|
||
RefinedDoc: req.RefinedDoc,
|
||
OutputDir: filepath.Join(s.config.OutputDir, taskID),
|
||
Name: req.Desc,
|
||
|
||
CallLLM: call.NewCallLLM(llmSet),
|
||
}
|
||
if err := s.taskImpl.Add(ctx, task.TasKModel); err != nil {
|
||
return "", errors.New("添加任务失败")
|
||
}
|
||
// 异步处理
|
||
//s.processTask(context.Background(), &task)
|
||
go func() {
|
||
newCtx := context.Background()
|
||
defer func() {
|
||
if r := recover(); r != nil {
|
||
s.failTask(newCtx, &task, fmt.Sprintf("panic: %v", r))
|
||
}
|
||
}()
|
||
s.processTask(newCtx, &task)
|
||
}()
|
||
|
||
return taskID, nil
|
||
}
|
||
|
||
func (s *SDKGeneratorBiz) processTask(ctx context.Context, task *entitys.Task) {
|
||
var usAges []*entitys.Usage
|
||
defer func() {
|
||
var total = &entitys.UsageTask{
|
||
Detail: make([]*entitys.Usage, len(usAges)),
|
||
}
|
||
|
||
for k, usage := range usAges {
|
||
total.PromptTokens += usage.PromptTokens
|
||
total.CompletionTokens += usage.CompletionTokens
|
||
total.TotalTokens += usage.TotalTokens
|
||
total.Detail[k] = usage
|
||
}
|
||
task.TasKModel.UseAge = pkg.JsonStringIgonErr(total)
|
||
s.updateTask(ctx, task.TasKModel)
|
||
}()
|
||
s.processTaskStatus(ctx, task, entitys.StatusProcessing)
|
||
|
||
//0.创建目录
|
||
s.processTaskStatus(ctx, task, entitys.StatusMkdir)
|
||
err := s.mkdir(ctx, task)
|
||
if err != nil {
|
||
s.failTask(ctx, task, fmt.Sprintf("[mkdir]: %v", err))
|
||
return
|
||
}
|
||
|
||
// 2. 生成代码
|
||
s.processTaskStatus(ctx, task, entitys.StatusGenerateCode)
|
||
usAge, err := s.generateCode(ctx, task)
|
||
if usAge != nil {
|
||
usAge.StateName = entitys.StatusGenerateCode.Desc()
|
||
usAges = append(usAges, usAge)
|
||
}
|
||
if err != nil {
|
||
s.failTask(ctx, task, fmt.Sprintf("[generateCode]: %v", err))
|
||
return
|
||
}
|
||
|
||
// 4. 验证是否存在缺失
|
||
s.processTaskStatus(ctx, task, entitys.StatusValid)
|
||
usAge, err = s.valid(ctx, task)
|
||
if usAge != nil {
|
||
usAge.StateName = entitys.StatusValid.Desc()
|
||
usAges = append(usAges, usAge)
|
||
}
|
||
|
||
if err != nil {
|
||
// 后处理失败不中断流程
|
||
s.failTask(ctx, task, fmt.Sprintf("[valid]: %v", err))
|
||
return
|
||
}
|
||
|
||
// 3. 处理文件
|
||
s.processTaskStatus(ctx, task, entitys.StatusCreatFile)
|
||
err = s.creatFile(ctx, task)
|
||
if err != nil {
|
||
// 后处理失败不中断流程
|
||
|
||
s.failTask(ctx, task, fmt.Sprintf("[creatFile]: %v", err))
|
||
return
|
||
}
|
||
|
||
// 5. 后处理
|
||
s.processTaskStatus(ctx, task, entitys.StatusProcess)
|
||
err = postprocess.Process(ctx, task.Package)
|
||
if err != nil {
|
||
//如果运行报错则多轮重试
|
||
var failUsage = &entitys.Usage{
|
||
StateName: entitys.StatusProcess.Desc(),
|
||
}
|
||
s.processTaskStatus(ctx, task, entitys.StatusFix)
|
||
err = s.fix(ctx, task, err.Error(), 1, failUsage)
|
||
if err != nil {
|
||
s.failTask(ctx, task, fmt.Sprintf("[fix]: %v", err))
|
||
}
|
||
usAges = append(usAges, failUsage)
|
||
}
|
||
|
||
// 6. 上传仓库
|
||
s.processTaskStatus(ctx, task, entitys.StatusSendToGit)
|
||
err = s.sendToGit(ctx, task)
|
||
if err != nil {
|
||
s.failTask(ctx, task, fmt.Sprintf("[sendToGit]: %v", err))
|
||
}
|
||
|
||
// 6. 发送钉钉消息
|
||
//s.processTaskStatus(ctx, task, entitys.StatusMsgSend)
|
||
//msg.SendDingTalkAlert(ctx, task.Name, task.TasKModel.RepoURL)
|
||
//// 5. 完成
|
||
now := time.Now()
|
||
task.TasKModel.TaskStatus = entitys.StatusCompleted.String()
|
||
task.TasKModel.CompletedAt = &now
|
||
s.updateTask(ctx, task.TasKModel)
|
||
}
|
||
|
||
func (s *SDKGeneratorBiz) sendToGit(ctx context.Context, task *entitys.Task) error {
|
||
// 1. 获取 Gitea 客户端
|
||
giteaClient, err := push.GetClient()
|
||
if err != nil || giteaClient == nil {
|
||
return fmt.Errorf("[GetClient]: %v", err)
|
||
}
|
||
|
||
// 2. 准备文件列表
|
||
var files []push.FileConfig
|
||
|
||
// 如果有提取的文件列表,推送所有文件
|
||
if len(task.Files) > 0 {
|
||
for _, file := range task.Files {
|
||
files = append(files, push.FileConfig{
|
||
Path: file.Path,
|
||
Content: file.Content,
|
||
Message: fmt.Sprintf("添加文件: %s", file.Path),
|
||
Branch: "main",
|
||
})
|
||
}
|
||
} else if task.Generate != "" {
|
||
// 如果只有 Resp 内容,作为 main.go 推送
|
||
files = append(files, push.FileConfig{
|
||
Path: "main.go",
|
||
Content: task.Generate,
|
||
Message: "Initial commit: SDK code",
|
||
Branch: "main",
|
||
})
|
||
}
|
||
|
||
//3. 如果有 MarkdownContent,添加为 README.md
|
||
files = append(files, push.FileConfig{
|
||
Path: "README.md",
|
||
Content: task.RefinedDoc,
|
||
Message: "添加 README 文档",
|
||
Branch: "main",
|
||
})
|
||
|
||
// 4. 如果没有文件,返回错误
|
||
if len(files) == 0 {
|
||
return fmt.Errorf("没有可推送的文件内容")
|
||
}
|
||
|
||
// 5. 生成仓库名称(使用时间戳确保唯一性)
|
||
repoName := fmt.Sprintf("%s-%s", task.Name, time.Now().Format("20060102150405"))
|
||
|
||
// 6. 配置标签(可选)
|
||
var tagCfg *push.TagConfig
|
||
if task.TasKModel.CompletedAt != nil {
|
||
tagCfg = &push.TagConfig{
|
||
TagName: fmt.Sprintf("v%s", time.Now().Format("20060102150405")),
|
||
Message: fmt.Sprintf("生成完成: %s", task.Name),
|
||
}
|
||
}
|
||
|
||
// 7. 创建仓库并推送代码
|
||
repoCfg := push.RepoConfig{
|
||
Name: repoName,
|
||
Description: task.Name,
|
||
Private: false, // 公共仓库
|
||
AutoInit: false,
|
||
Gitignores: "",
|
||
License: "",
|
||
Readme: "",
|
||
}
|
||
|
||
err = giteaClient.CreateRepoAndPush(repoCfg, files, tagCfg, s.config.OrgName)
|
||
if err != nil {
|
||
|
||
return fmt.Errorf("推送代码到 Gitea 失败: %v", err)
|
||
}
|
||
|
||
task.TasKModel.RepoURL = fmt.Sprintf("%s/%s/%s", s.config.GiteaUrl, s.config.OrgName, repoName)
|
||
return nil
|
||
}
|
||
|
||
func (s *SDKGeneratorBiz) mkdir(ctx context.Context, task *entitys.Task) (err error) {
|
||
if _, err = os.Stat(task.OutputDir); err == nil {
|
||
if err = os.RemoveAll(task.OutputDir); err != nil {
|
||
return fmt.Errorf("删除旧目录失败: %v", err)
|
||
}
|
||
}
|
||
if err = os.MkdirAll(task.OutputDir, 0755); err != nil {
|
||
return fmt.Errorf("创建输出目录失败: %v", err)
|
||
}
|
||
task.Package = filepath.Join(task.OutputDir, task.Name)
|
||
return err
|
||
}
|
||
|
||
func (s *SDKGeneratorBiz) generateCode(ctx context.Context, task *entitys.Task) (usage *entitys.Usage, err error) {
|
||
var NeedImplement string
|
||
if len(task.TasKModel.Interfaces) != 0 {
|
||
var Interfaces []entitys.Interface
|
||
_err := json.Unmarshal([]byte(task.TasKModel.Interfaces), &Interfaces)
|
||
if _err == nil {
|
||
NeedImplement = "\n ## ⚠️ 接口范围限制(最高优先级,必须严格遵守)\n\n**只实现以下指定的接口,其他接口不要生成:**:"
|
||
for _, v := range Interfaces {
|
||
NeedImplement += fmt.Sprintf("\n\t%s", v.Summary)
|
||
}
|
||
}
|
||
}
|
||
var (
|
||
finalPrompt strings.Builder
|
||
)
|
||
finalPrompt.WriteString(task.RefinedDoc)
|
||
//finalPrompt.WriteString(NeedImplement)
|
||
|
||
switch entitys.DocType(task.TasKModel.GenerateType) {
|
||
case entitys.DocTypeSdk:
|
||
sdkGen := prompts.NewSDKGenerator(
|
||
task.CallLLM.Client,
|
||
prompts.WithMaxIterations(10),
|
||
prompts.WithTimeout(10*time.Minute),
|
||
prompts.WithModel(task.TasKModel.LlmModel),
|
||
)
|
||
task.Generate, usage, err = sdkGen.GenerateSDK(ctx, finalPrompt.String(), task.Name, NeedImplement)
|
||
case entitys.DocTypeServerBoilerplate:
|
||
serverGen := prompts.NewServerGenerator(
|
||
task.CallLLM.Client,
|
||
prompts.WithServerMaxIterations(10),
|
||
prompts.WithServerTimeout(10*time.Minute),
|
||
prompts.WithServerModel(task.TasKModel.LlmModel),
|
||
)
|
||
task.Generate, usage, err = serverGen.GenerateServer(ctx, finalPrompt.String(), task.Name, NeedImplement)
|
||
}
|
||
|
||
return
|
||
}
|
||
|
||
func (s *SDKGeneratorBiz) valid(ctx context.Context, task *entitys.Task) (usage *entitys.Usage, err error) {
|
||
log.Printf("valid(): task.CallLLM=%p, task.CallLLM.Client=%p, task.CallLLM.ApiKey prefix=%s, task.Name=%s",
|
||
task.CallLLM, task.CallLLM.Client, prefix(task.CallLLM.ApiKey), task.Name)
|
||
|
||
validRes, useAge, err := task.CallLLM.Do(ctx, prompts.GetValidatePrompt(task.RefinedDoc, task.Name, task.Generate))
|
||
if err != nil {
|
||
log.Printf("valid(): error calling LLM: %v", err)
|
||
return nil, fmt.Errorf("调用大模型失败: %v", err)
|
||
}
|
||
|
||
usage = &entitys.Usage{
|
||
PromptTokens: useAge.PromptTokens,
|
||
CompletionTokens: useAge.CompletionTokens,
|
||
TotalTokens: useAge.TotalTokens,
|
||
}
|
||
|
||
// ✅ 清理返回结果
|
||
cleaned := s.cleanValidationResult(validRes)
|
||
|
||
// ✅ 判断是否为 OK
|
||
if cleaned == "OK" {
|
||
task.Valid = task.Generate
|
||
return usage, nil
|
||
}
|
||
|
||
// 验证不通过,保存修复后的代码
|
||
task.Valid = validRes
|
||
return usage, nil
|
||
}
|
||
|
||
// cleanValidationResult 清理验证结果
|
||
func (s *SDKGeneratorBiz) cleanValidationResult(result string) string {
|
||
// 1. 去除首尾空白
|
||
cleaned := strings.TrimSpace(result)
|
||
|
||
// 2. 去除 Markdown 代码块标记
|
||
cleaned = strings.TrimPrefix(cleaned, "```")
|
||
cleaned = strings.TrimSuffix(cleaned, "```")
|
||
cleaned = strings.TrimSpace(cleaned)
|
||
|
||
// 3. 去除可能的引号
|
||
cleaned = strings.Trim(cleaned, "\"")
|
||
cleaned = strings.Trim(cleaned, "'")
|
||
|
||
// 4. 只取前两行(防止 OK 后面跟了其他内容)
|
||
lines := strings.Split(cleaned, "\n")
|
||
if len(lines) > 0 {
|
||
cleaned = strings.TrimSpace(lines[0])
|
||
}
|
||
|
||
// 5. 去除可能的尾随标点
|
||
cleaned = strings.TrimSuffix(cleaned, ".")
|
||
cleaned = strings.TrimSuffix(cleaned, "!")
|
||
cleaned = strings.TrimSuffix(cleaned, "!")
|
||
|
||
return cleaned
|
||
}
|
||
func (s *SDKGeneratorBiz) fix(ctx context.Context, task *entitys.Task, errMsg string, fixCount int, usage *entitys.Usage) (err error) {
|
||
log.Printf("代码存在问题,需要修复: %s,当前修复次数:%d", errMsg, fixCount)
|
||
if fixCount >= s.config.MaxFixAttempts {
|
||
return fmt.Errorf("修复失败,已达到最大尝试次数: %v", fixCount)
|
||
}
|
||
fix, useAge, err := task.CallLLM.Do(ctx, prompts.FixPrompt(task.Valid, errMsg, task.Name, task.RefinedDoc))
|
||
if err != nil {
|
||
return fmt.Errorf("调用大模型失败: %v", err)
|
||
}
|
||
usage.PromptTokens += useAge.PromptTokens
|
||
usage.CompletionTokens += useAge.CompletionTokens
|
||
usage.TotalTokens += useAge.TotalTokens
|
||
|
||
task.Valid = fix
|
||
if err = s.creatFile(ctx, task); err != nil {
|
||
// 后处理失败不中断流程
|
||
return fmt.Errorf("创建文件失败: %v", err)
|
||
}
|
||
|
||
if err = postprocess.Process(ctx, task.Package); err != nil {
|
||
err = s.fix(ctx, task, err.Error(), fixCount+1, usage)
|
||
}
|
||
return err
|
||
}
|
||
|
||
func (s *SDKGeneratorBiz) creatFile(ctx context.Context, task *entitys.Task) error {
|
||
// 3. 提取代码
|
||
extract := task.Generate
|
||
if task.Valid != "" {
|
||
extract = task.Valid
|
||
}
|
||
files, err := extractor.Extract(extract)
|
||
if err != nil {
|
||
return fmt.Errorf("提取代码失败: %v", err)
|
||
}
|
||
|
||
if len(files) == 0 {
|
||
return fmt.Errorf("未提取到任何代码文件")
|
||
}
|
||
files = append(files,
|
||
extractor.File{
|
||
Path: filepath.Join(task.Name, "refine.md"),
|
||
Content: task.RefinedDoc,
|
||
},
|
||
extractor.File{
|
||
Path: filepath.Join(task.Name, "generate.md"),
|
||
Content: task.Generate,
|
||
},
|
||
extractor.File{
|
||
Path: filepath.Join(task.Name, "valid.md"),
|
||
Content: task.Valid,
|
||
},
|
||
)
|
||
switch entitys.DocType(task.TasKModel.GenerateType) {
|
||
case entitys.DocTypeSdk:
|
||
|
||
case entitys.DocTypeServerBoilerplate:
|
||
//加上dockerfile
|
||
files = append(files, extractor.File{
|
||
Path: filepath.Join(task.Name, "Dockerfile"),
|
||
Content: file.GetDockerfile(),
|
||
})
|
||
}
|
||
// 3. 写入文件
|
||
if err = extractor.WriteFiles(files, task.OutputDir); err != nil {
|
||
return fmt.Errorf("写入文件失败: %v", err)
|
||
}
|
||
task.Files = files
|
||
return nil
|
||
}
|
||
|
||
func (s *SDKGeneratorBiz) failTask(ctx context.Context, task *entitys.Task, errMsg string) {
|
||
task.TasKModel.TaskStatus = entitys.StatusFailed.String()
|
||
task.TasKModel.Error = errMsg
|
||
s.updateTask(ctx, task.TasKModel)
|
||
}
|
||
|
||
func (s *SDKGeneratorBiz) processTaskStatus(ctx context.Context, task *entitys.Task, process entitys.TaskStatus) {
|
||
s.taskImpl.UpdateByKey(ctx, s.taskImpl.PrimaryKey(), task.TasKModel.ID, &model.AiGenerateTask{TaskStatus: process.String()})
|
||
}
|
||
|
||
func (s *SDKGeneratorBiz) updateTask(ctx context.Context, task *model.AiGenerateTask) error {
|
||
return s.taskImpl.UpdateByKey(ctx, s.taskImpl.PrimaryKey(), task.ID, task)
|
||
}
|
||
|
||
// GetTask 获取任务状态
|
||
func (s *SDKGeneratorBiz) GetTask(ctx context.Context, taskID string) (*model.AiGenerateTask, bool) {
|
||
var task model.AiGenerateTask
|
||
err := s.taskImpl.GetByKey(ctx, "task_id", taskID, &task)
|
||
if err != nil || task.ID == 0 {
|
||
return nil, false
|
||
}
|
||
return &task, true
|
||
}
|
||
|
||
// ListTasks 列出所有任务
|
||
func (s *SDKGeneratorBiz) ListTasks(ctx context.Context, cond *builder.Cond, page, pageSize int) ([]model.AiGenerateTask, int64, error) {
|
||
var tasks []model.AiGenerateTask
|
||
total, err := s.taskImpl.GetListToStruct(ctx, cond, &dataTemp.ReqPageBo{
|
||
Page: page,
|
||
Limit: pageSize,
|
||
}, &tasks, "updated_at DESC")
|
||
if err != nil {
|
||
return nil, 0, err
|
||
}
|
||
return tasks, total.Total, err
|
||
}
|
||
|
||
// DeleteTask 删除任务
|
||
func (s *SDKGeneratorBiz) DeleteTask(ctx context.Context, taskID string) error {
|
||
err := s.taskImpl.DeleteByKey(ctx, "task_id", taskID)
|
||
|
||
return err
|
||
}
|
||
|
||
// GetTaskOutputDir 获取任务输出目录
|
||
func (s *SDKGeneratorBiz) GetTaskOutputDir(taskID string) string {
|
||
return filepath.Join(s.config.OutputDir, taskID)
|
||
}
|
||
|
||
// ListDoc 列出所有任务
|
||
func (s *SDKGeneratorBiz) ListDoc(ctx context.Context, cond *builder.Cond, page, pageSize int) ([]model.AiGenerateDoc, int64, error) {
|
||
var tasks []model.AiGenerateDoc
|
||
total, err := s.docImpl.GetListToStruct(ctx, cond, &dataTemp.ReqPageBo{
|
||
Page: page,
|
||
Limit: pageSize,
|
||
}, &tasks, "updated_at DESC")
|
||
if err != nil {
|
||
return nil, 0, err
|
||
}
|
||
return tasks, total.Total, err
|
||
}
|
||
|
||
// ListDoc 列出所有任务
|
||
func (s *SDKGeneratorBiz) UpdateDocByInstanceId(ctx context.Context, docModel *model.AiGenerateDoc) error {
|
||
return s.docImpl.UpdateByKey(ctx, "instance_id", docModel.InstanceID, docModel)
|
||
}
|
||
|
||
// prefix returns a short prefix of the key for safe logging
|
||
func prefix(s string) string {
|
||
if s == "" {
|
||
return ""
|
||
}
|
||
if len(s) <= 8 {
|
||
return s
|
||
}
|
||
return s[:8]
|
||
}
|