package service import ( "context" "fmt" "log" "os" "path/filepath" "sdk-generator/internal/call" "sdk-generator/internal/config" "sdk-generator/internal/extractor" "sdk-generator/internal/models" "sdk-generator/internal/msg" "sdk-generator/internal/postprocess" "sdk-generator/internal/prompts" "sdk-generator/internal/push" "strings" "sync" "time" "github.com/google/uuid" ) type SDKGeneratorService struct { config *config.Config tasks map[string]*models.Task tasksMu sync.RWMutex } func NewSDKGeneratorService(cfg *config.Config) *SDKGeneratorService { // 创建 OpenAI 客户端配置 return &SDKGeneratorService{ config: cfg, tasks: make(map[string]*models.Task), } } // Generate 异步生成 SDK func (s *SDKGeneratorService) Generate(ctx context.Context, content []byte, sdkName string, llmSet *call.LlmCallSet) (string, error) { // 限制任务数量 s.tasksMu.RLock() if len(s.tasks) >= s.config.MaxTasks { s.tasksMu.RUnlock() return "", fmt.Errorf("任务队列已满,请稍后重试") } s.tasksMu.RUnlock() taskID := uuid.New().String() task := &models.Task{ ID: taskID, Status: models.StatusPending, SdkName: sdkName, OutputDir: filepath.Join(s.config.OutputDir, taskID), CreatedAt: time.Now(), UpdatedAt: time.Now(), CallLLM: call.NewCallLLM(llmSet), MarkdownContent: string(content), } s.tasksMu.Lock() s.tasks[taskID] = task s.tasksMu.Unlock() // 异步处理 go func() { defer func() { if r := recover(); r != nil { s.failTask(task, fmt.Sprintf("panic: %v", r)) } }() s.processTask(context.Background(), task, string(content)) }() return taskID, nil } func (s *SDKGeneratorService) processTask(ctx context.Context, task *models.Task, markdownContent string) { s.processTaskStatus(task, models.StatusProcessing) //0.创建目录 s.processTaskStatus(task, models.StatusMkdir) err := s.mkdir(ctx, task) if err != nil { s.failTask(task, fmt.Sprintf("[mkdir]: %v", err)) return } // 1. 分析文档 s.processTaskStatus(task, models.StatusAnaMd) err = s.anaMd(ctx, task) if err != nil { s.failTask(task, fmt.Sprintf("[anaMd]: %v", err)) return } // 2. 生成代码 s.processTaskStatus(task, models.StatusGenerateCode) err = s.generateCode(ctx, task) if err != nil { s.failTask(task, fmt.Sprintf("[generateCode]: %v", err)) return } // 4. 验证是否存在缺失 s.processTaskStatus(task, models.StatusValid) err = s.valid(ctx, task) if err != nil { // 后处理失败不中断流程 s.failTask(task, fmt.Sprintf("[valid]: %v", err)) return } // 3. 处理文件 s.processTaskStatus(task, models.StatusCreatFile) err = s.creatFile(ctx, task) if err != nil { // 后处理失败不中断流程 s.failTask(task, fmt.Sprintf("[creatFile]: %v", err)) return } // 5. 后处理 s.processTaskStatus(task, models.StatusProcess) err = postprocess.Process(ctx, task.SdkPackage) if err != nil { //如果运行报错则多轮重试 s.processTaskStatus(task, models.StatusFix) err = s.fix(ctx, task, err.Error(), 1) } // 6. 上传仓库 s.processTaskStatus(task, models.StatusSendToGit) err = s.sendToGit(ctx, task) if err != nil { s.failTask(task, fmt.Sprintf("[sendToGit]: %v", err)) } // 6. 发送钉钉消息 s.processTaskStatus(task, models.StatusMsgSend) msg.SendDingTalkAlert(ctx, task.SdkName, task.RepoURL) // 5. 完成 now := time.Now() task.Status = models.StatusCompleted task.UpdatedAt = now task.CompletedAt = &now s.updateTask(task) } func (s *SDKGeneratorService) sendToGit(ctx context.Context, task *models.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.Resp != "" { // 如果只有 Resp 内容,作为 main.go 推送 files = append(files, push.FileConfig{ Path: "main.go", Content: task.Resp, Message: "Initial commit: SDK code", Branch: "main", }) } //3. 如果有 MarkdownContent,添加为 README.md if task.MarkdownContent != "" { files = append(files, push.FileConfig{ Path: "README.md", Content: task.MarkdownContent, Message: "添加 README 文档", Branch: "main", }) } // 4. 如果没有文件,返回错误 if len(files) == 0 { return fmt.Errorf("没有可推送的文件内容") } // 5. 生成仓库名称(使用时间戳确保唯一性) repoName := fmt.Sprintf("%s-%s", task.SdkName, task.CreatedAt.Format("20060102-150405")) // 6. 配置标签(可选) var tagCfg *push.TagConfig if task.CompletedAt != nil { tagCfg = &push.TagConfig{ TagName: fmt.Sprintf("v%s", task.CompletedAt.Format("20060102.150405")), Message: fmt.Sprintf("SDK 生成完成: %s", task.SdkName), } } // 7. 创建仓库并推送代码 repoCfg := push.RepoConfig{ Name: repoName, Description: task.SdkName, Private: false, // 公共仓库 AutoInit: false, Gitignores: "", License: "", Readme: "", } err = giteaClient.CreateRepoAndPush(repoCfg, files, tagCfg, s.config.OrgName) if err != nil { s.failTask(task, fmt.Sprintf("推送代码到 Gitea 失败: %v", err)) return fmt.Errorf("推送代码到 Gitea 失败: %v", err) } task.RepoURL = fmt.Sprintf("%s/%s/%s", s.config.GiteaUrl, s.config.OrgName, repoName) return nil } func (s *SDKGeneratorService) mkdir(ctx context.Context, task *models.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.SdkPackage = filepath.Join(task.OutputDir, task.SdkName) return err } func (s *SDKGeneratorService) generateCode(ctx context.Context, task *models.Task) (err error) { promptSetSdtName := strings.Replace(prompts.GenSDKPrompt, "{{sdk_name}}", task.SdkName, -1) doc := task.Refine if len(doc) == 0 { doc = task.MarkdownContent } prompt := fmt.Sprintf(promptSetSdtName, doc) task.Resp, err = task.CallLLM.Do(ctx, prompts.SDKGeneratorPrompt(prompt)) if err != nil { s.failTask(task, fmt.Sprintf("调用大模型失败: %v", err)) return } return err } func (s *SDKGeneratorService) anaMd(ctx context.Context, task *models.Task) (err error) { task.Refine, err = task.CallLLM.Do(ctx, prompts.BuildRefinePrompt(task.MarkdownContent)) if err != nil { return fmt.Errorf("调用大模型失败: %v", err) } return err } func (s *SDKGeneratorService) valid(ctx context.Context, task *models.Task) (err error) { validRes, err := task.CallLLM.Do(ctx, prompts.GetValidatePrompt(task.Refine, task.SdkName, task.Resp)) if err != nil { return fmt.Errorf("调用大模型失败: %v", err) } // ✅ 清理返回结果 cleaned := s.cleanValidationResult(validRes) // ✅ 判断是否为 OK if cleaned == "OK" { task.Valid = "OK" return nil } // 验证不通过,保存修复后的代码 task.Valid = validRes return nil } // cleanValidationResult 清理验证结果 func (s *SDKGeneratorService) 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 *SDKGeneratorService) fix(ctx context.Context, task *models.Task, errMsg string, fixCount int) (err error) { log.Printf("代码存在问题,需要修复: %s,当前修复次数:%d", errMsg, fixCount) if fixCount >= s.config.MaxFixAttempts { return fmt.Errorf("修复失败,已达到最大尝试次数: %v", fixCount) } fix, err := task.CallLLM.Do(ctx, prompts.FixPrompt(task.Resp, errMsg, task.SdkName, task.Refine)) if err != nil { return fmt.Errorf("调用大模型失败: %v", err) } if len(task.Valid) > 0 && task.Valid != "OK" { task.Valid = fix } else { task.Resp = fix } if err = s.creatFile(ctx, task); err != nil { // 后处理失败不中断流程 return fmt.Errorf("创建文件失败: %v", err) } if err = postprocess.Process(ctx, task.SdkPackage); err != nil { err = s.fix(ctx, task, err.Error(), fixCount+1) } return err } func (s *SDKGeneratorService) creatFile(ctx context.Context, task *models.Task) error { // 3. 提取代码 extract := task.Resp 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.SdkName, "doc.md"), Content: task.MarkdownContent, }, extractor.File{ Path: filepath.Join(task.SdkName, "refine.md"), Content: task.Refine, }, extractor.File{ Path: filepath.Join(task.SdkName, "generate.md"), Content: task.Resp, }, extractor.File{ Path: filepath.Join(task.SdkName, "valid.md"), Content: task.Valid, }, ) // 3. 写入文件 if err = extractor.WriteFiles(files, task.OutputDir); err != nil { return fmt.Errorf("写入文件失败: %v", err) } task.Files = files return nil } func (s *SDKGeneratorService) failTask(task *models.Task, errMsg string) { task.Status = models.StatusFailed task.Error = errMsg task.UpdatedAt = time.Now() s.updateTask(task) } func (s *SDKGeneratorService) processTaskStatus(task *models.Task, process models.TaskStatus) { task.Status = process task.UpdatedAt = time.Now() s.tasksMu.Lock() defer s.tasksMu.Unlock() s.tasks[task.ID] = task } func (s *SDKGeneratorService) updateTask(task *models.Task) { s.tasksMu.Lock() defer s.tasksMu.Unlock() s.tasks[task.ID] = task } // GetTask 获取任务状态 func (s *SDKGeneratorService) GetTask(taskID string) (*models.Task, bool) { s.tasksMu.RLock() defer s.tasksMu.RUnlock() task, ok := s.tasks[taskID] return task, ok } // ListTasks 列出所有任务 func (s *SDKGeneratorService) ListTasks() []*models.Task { s.tasksMu.RLock() defer s.tasksMu.RUnlock() tasks := make([]*models.Task, 0, len(s.tasks)) for _, task := range s.tasks { tasks = append(tasks, task) } return tasks } // DeleteTask 删除任务 func (s *SDKGeneratorService) DeleteTask(taskID string) error { s.tasksMu.Lock() defer s.tasksMu.Unlock() task, ok := s.tasks[taskID] if !ok { return fmt.Errorf("任务不存在") } // 删除输出目录 if err := os.RemoveAll(task.OutputDir); err != nil { return fmt.Errorf("删除输出目录失败: %w", err) } delete(s.tasks, taskID) return nil } // GetTaskOutputDir 获取任务输出目录 func (s *SDKGeneratorService) GetTaskOutputDir(taskID string) string { return filepath.Join(s.config.OutputDir, taskID) }