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 } llmSet.BaseUrl = "https://ark.cn-beijing.volces.com/api/v3" // 限制任务数量 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] }