diff --git a/1.jpg b/1.jpg deleted file mode 100644 index c98e4c5..0000000 Binary files a/1.jpg and /dev/null differ diff --git a/Dockerfile b/Dockerfile index 8b02af6..d092dea 100644 --- a/Dockerfile +++ b/Dockerfile @@ -36,10 +36,14 @@ RUN echo 'http://mirrors.ustc.edu.cn/alpine/v3.5/main' > /etc/apk/repositories \ # 设置工作目录 WORKDIR /app -# 将编译好的二进制文件从构建阶段复制到运行阶段 +# 复制编译好的二进制文件从构建阶段复制到运行阶段 COPY --from=builder /app/server ./server # 复制配置文件夹 COPY --from=builder /app/config ./config +# 复制前端静态页面 +COPY --from=builder /app/web ./web +# 创建上传目录(聊天记录等文件) +RUN mkdir -p upload ENV TZ=Asia/Shanghai diff --git a/README.md b/README.md deleted file mode 100644 index 2cae5d7..0000000 --- a/README.md +++ /dev/null @@ -1,140 +0,0 @@ -# AI Scheduler - 智能路由调度系统 - -基于Go语言开发的智能AI助手,支持Function Calling和工具调用,可以智能路由用户请求到合适的工具进行处理。 - -## 功能特性 - -- 🤖 **智能对话**: 基于Ollama的AI对话能力 -- 🔧 **工具调用**: 支持天气查询、计算器等工具 -- 🎯 **智能路由**: 自动判断是否需要调用工具 -- 📚 **API文档**: 集成Swagger文档 -- ⚡ **高性能**: 基于Fiber框架的Websocket服务 -- 🏗️ **依赖注入**: 使用Wire进行依赖管理 - -## 项目结构 - -``` -ai_scheduler/ -├── cmd/ # 应用程序入口 -│ └── main.go -├── internal/ # 内部包 -│ ├── config/ # 配置管理 -│ ├── handlers/ # HTTP处理器 -│ ├── models/ # 数据模型 -│ ├── services/ # 业务服务 -│ ├── tools/ # 工具实现 -│ └── wire/ # 依赖注入 -├── pkg/ # 公共包 -│ ├── ollama/ # Ollama客户端 -│ └── types/ # 类型定义 -├── docs/ # API文档 -├── config.yaml # 配置文件 -└── go.mod # Go模块 -``` - -## 快速开始 - -### 1. 环境要求 - -- Go 1.23+ -- Ollama服务运行中 - -### 2. 安装依赖 - -```bash -go mod tidy -``` - -### 3. 配置文件 - -编辑 `config.yaml` 文件,确保Ollama服务地址正确: - -```yaml -server: - port: "8080" - host: "localhost" - -ollama: - base_url: "http://localhost:11434" - model: "llama2" - timeout: 30s - -tools: - weather: - enabled: true - mock_data: true - calculator: - enabled: true - mock_data: false - -logging: - level: "info" - format: "json" -``` - -### 4. 启动服务 - -```json -go run cmd/main.go -``` - -服务启动后,可以访问: - -- API服务: http://localhost:8080/api/v1/chat -- Swagger文档: http://localhost:8080/swagger/index.html -- 健康检查: http://localhost:8080/health - -## API使用示例 - -### 聊天接口 - -```bash -curl -X POST http://localhost:8080/api/v1/chat \ - -H "Content-Type: application/json" \ - -d '{ - "message": "北京今天天气怎么样?", - "model": "llama2" - }' -``` - -### 计算器示例 - -```bash -curl -X POST http://localhost:8080/api/v1/chat \ - -H "Content-Type: application/json" \ - -d '{ - "message": "计算 15 + 25 * 3", - "model": "llama2" - }' -``` - -## 支持的工具 - -### 1. 天气查询工具 -- 功能:查询指定城市的天气信息 -- 示例:"北京今天天气怎么样?" - -### 2. 计算器工具 -- 功能:执行数学计算 -- 支持:加减乘除、幂运算 -- 示例:"计算 2 + 3 * 4" - -## 开发说明 - -### 添加新工具 - -1. 在 `internal/tools/` 目录下创建新工具文件 -2. 实现 `types.Tool` 接口 -3. 在 `tools.Manager` 中注册新工具 - -### 配置管理 - -配置文件使用Viper加载,支持环境变量覆盖。 - -### 依赖注入 - -使用Google Wire进行依赖注入,修改依赖关系后需要重新生成代码。 - -## 许可证 - -MIT License \ No newline at end of file diff --git a/cmd/server/wire.go b/cmd/server/wire.go index 56f5868..7ee4a39 100644 --- a/cmd/server/wire.go +++ b/cmd/server/wire.go @@ -23,13 +23,11 @@ import ( func InitializeApp(context.Context, *config.Config, log.AllLogger) (*server.Servers, func(), error) { panic(wire.Build( server.ProviderSetServer, - pkg.ProviderSetClient, advice.ProviderService, biz.ProviderSetBiz, impl.ProviderImpl, utils.ProviderUtils, - mongo_model.ProviderSetMongo, )) diff --git a/go.mod b/go.mod index f42c49e..32d3a27 100644 --- a/go.mod +++ b/go.mod @@ -85,7 +85,6 @@ require ( github.com/hashicorp/hcl v1.0.0 // indirect github.com/jinzhu/inflection v1.0.0 // indirect github.com/jinzhu/now v1.1.5 // indirect - github.com/jmespath/go-jmespath v0.4.0 // indirect github.com/json-iterator/go v1.1.12 // indirect github.com/klauspost/compress v1.17.9 // indirect github.com/klauspost/cpuid/v2 v2.4.0 // indirect diff --git a/go.sum b/go.sum index 1d4eaa7..eb9516b 100644 --- a/go.sum +++ b/go.sum @@ -280,6 +280,8 @@ github.com/google/pprof v0.0.0-20201023163331-3e6fc7fc9c4c/go.mod h1:kpwsk12EmLe github.com/google/pprof v0.0.0-20201203190320-1bf35d6f28c2/go.mod h1:kpwsk12EmLew5upagYY7GY0pfYCcupk39gWOCRROcvE= github.com/google/pprof v0.0.0-20201218002935-b9804c9f04c2/go.mod h1:kpwsk12EmLew5upagYY7GY0pfYCcupk39gWOCRROcvE= github.com/google/renameio v0.1.0/go.mod h1:KWCgfxg9yswjAJkECMjeO8J8rahYeXnNhOm40UhjYkI= +github.com/google/subcommands v1.2.0 h1:vWQspBTo2nEqTUFita5/KeEWlUL8kQObDFbub/EN9oE= +github.com/google/subcommands v1.2.0/go.mod h1:ZjhPrFU+Olkh9WazFPsl27BQ4UPiG37m3yTrtFlrHVk= github.com/google/uuid v1.1.2/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= github.com/google/uuid v1.3.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0= @@ -309,9 +311,7 @@ github.com/jinzhu/inflection v1.0.0 h1:K317FqzuhWc8YvSVlFMCCUb36O/S9MCKRDI7QkRKD github.com/jinzhu/inflection v1.0.0/go.mod h1:h+uFLlag+Qp1Va5pdKtLDYj+kHp5pxUVkryuEj+Srlc= github.com/jinzhu/now v1.1.5 h1:/o9tlHleP7gOFmsnYNz3RGnqzefHA47wQpKrrdTIwXQ= github.com/jinzhu/now v1.1.5/go.mod h1:d3SSVoowX0Lcu0IBviAWJpolVfI5UJVZZ7cO71lE/z8= -github.com/jmespath/go-jmespath v0.4.0 h1:BEgLn5cpjn8UN1mAw4NjwDrS35OdebyEtFe+9YPoQUg= github.com/jmespath/go-jmespath v0.4.0/go.mod h1:T8mJZnbsbmF+m6zOOFylbeCJqk5+pHWvzYPziyZiYoo= -github.com/jmespath/go-jmespath/internal/testify v1.5.1 h1:shLQSRRSCCPj3f2gpwzGwWFoC7ycTf1rcQZHOlsJ6N8= github.com/jmespath/go-jmespath/internal/testify v1.5.1/go.mod h1:L3OGu8Wl2/fWfCI6z80xFu9LTZmf1ZRjMHUOPmWr69U= github.com/josharian/intern v1.0.0/go.mod h1:5DoeVV0s6jJacbCEi61lwdGj/aVlrQvzHFFd8Hwg//Y= github.com/json-iterator/go v1.1.10/go.mod h1:KdQUCv79m/52Kvf8AW2vK1V8akMuk1QjK/uOdHXbAo4= @@ -583,6 +583,8 @@ golang.org/x/mod v0.4.0/go.mod h1:s0Qsj1ACt9ePp/hMypM3fl4fZqREWJwdYDEqhRiZZUA= golang.org/x/mod v0.4.1/go.mod h1:s0Qsj1ACt9ePp/hMypM3fl4fZqREWJwdYDEqhRiZZUA= golang.org/x/mod v0.6.0-dev.0.20220419223038-86c51ed26bb4/go.mod h1:jJ57K6gSWd91VN4djpZkiMVwK6gcyfeH4XE8wZrZaV4= golang.org/x/mod v0.8.0/go.mod h1:iBbtSCu2XBx23ZKBPSOrRkjjQPZFPuis4dIYUhu/chs= +golang.org/x/mod v0.28.0 h1:gQBtGhjxykdjY9YhZpSlZIsbnaE2+PgjfLWUQTnoZ1U= +golang.org/x/mod v0.28.0/go.mod h1:yfB/L0NOf/kmEbXjzCPOx1iK1fRutOydrCMsqRhEBxI= golang.org/x/net v0.0.0-20180724234803-3673e40ba225/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4= golang.org/x/net v0.0.0-20180826012351-8a410e7b638d/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4= golang.org/x/net v0.0.0-20180906233101-161cd47e91fd/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4= @@ -790,6 +792,8 @@ golang.org/x/tools v0.0.0-20210108195828-e2f9c7f1fc8e/go.mod h1:emZCQorbCU4vsT4f golang.org/x/tools v0.1.0/go.mod h1:xkSsbof2nBLbhDlRMhhhyNLN/zl3eTqcnHD5viDpcZ0= golang.org/x/tools v0.1.12/go.mod h1:hNGJHUnrk76NpqgfD5Aqm5Crs+Hm0VOH/i9J2+nxYbc= golang.org/x/tools v0.6.0/go.mod h1:Xwgl3UAJ/d3gWutnCtw505GrjyAbvKui8lOU390QaIU= +golang.org/x/tools v0.37.0 h1:DVSRzp7FwePZW356yEAChSdNcQo6Nsp+fex1SUW09lE= +golang.org/x/tools v0.37.0/go.mod h1:MBN5QPQtLMHVdvsbtarmTNukZDdgwdwlO5qGacAzF0w= golang.org/x/xerrors v0.0.0-20190717185122-a985d3407aa7/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0= golang.org/x/xerrors v0.0.0-20191011141410-1b5146add898/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0= golang.org/x/xerrors v0.0.0-20191204190536-9bdfabe68543/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0= diff --git a/internal/biz/advice_advicer.go b/internal/biz/advice_advicer.go index a5ddcf8..05ef207 100644 --- a/internal/biz/advice_advicer.go +++ b/internal/biz/advice_advicer.go @@ -116,7 +116,7 @@ func (a *AdviceAdvicerBiz) VersionUpdate(ctx context.Context, param *entitys.Adv return res.Err() } -func (a *AdviceAdvicerBiz) VersionList(ctx context.Context, param *entitys.AdvicerVersionListReq) (list []mongo_model.AdvicerVersionMongo, err error) { +func (a *AdviceAdvicerBiz) VersionList(ctx context.Context, param *entitys.AdvicerVersionListReq) (list []mongo_model.AdvicerVersionItem, err error) { filter := bson.M{} // 1. advicer_id 条件 if param.AdvicerId != 0 { @@ -149,7 +149,7 @@ func (a *AdviceAdvicerBiz) VersionList(ctx context.Context, param *entitys.Advic } // 遍历结果 for cursor.Next(ctx) { - var advicerVersion mongo_model.AdvicerVersionMongo + var advicerVersion mongo_model.AdvicerVersionItem if err := cursor.Decode(&advicerVersion); err != nil { return nil, err } diff --git a/internal/biz/advice_chat.go b/internal/biz/advice_chat.go index 946b53f..b628a1e 100644 --- a/internal/biz/advice_chat.go +++ b/internal/biz/advice_chat.go @@ -16,9 +16,7 @@ import ( "fmt" "github.com/google/uuid" - "github.com/volcengine/volcengine-go-sdk/service/arkruntime/model" - "github.com/volcengine/volcengine-go-sdk/service/arkruntime/model/responses" - "github.com/volcengine/volcengine-go-sdk/volcengine" + "github.com/sashabaranov/go-openai" "go.mongodb.org/mongo-driver/bson" "go.mongodb.org/mongo-driver/mongo/options" "xorm.io/builder" @@ -28,7 +26,7 @@ import ( ) type AdviceChatBiz struct { - hsyq *third_party.Hsyq + openai *third_party.OpenAi rdb *utils.Rdb aiAdviceSessionImpl *impl.AiAdviceSessionImpl aiAdviceModelSupImpl *impl.AiAdviceModelSupImpl @@ -37,7 +35,7 @@ type AdviceChatBiz struct { } func NewAdviceChatBiz( - hsyq *third_party.Hsyq, + openai *third_party.OpenAi, rdb *utils.Rdb, aiAdviceSessionImpl *impl.AiAdviceSessionImpl, aiAdviceModelSupImpl *impl.AiAdviceModelSupImpl, @@ -45,7 +43,7 @@ func NewAdviceChatBiz( mongo *pkg.Mongo, ) *AdviceChatBiz { return &AdviceChatBiz{ - hsyq: hsyq, + openai: openai, rdb: rdb, aiAdviceSessionImpl: aiAdviceSessionImpl, aiAdviceModelSupImpl: aiAdviceModelSupImpl, @@ -56,26 +54,17 @@ func NewAdviceChatBiz( func (a *AdviceChatBiz) contextCache(ctx context.Context, chatData *entitys.ChatData, req *entitys.AdvicerChatRegistReq, projectInfo *entitys.AdvicerProjectInfoRes) (promptJson string, contextCache string, err error) { switch constants.Mode(projectInfo.ModelInfo.Mode) { - case constants.ModeResponse: + case constants.ModeResponse, constants.ModeContext: + //ModeContext依赖火山专有Context缓存API,统一降级为Responses存储+PreviousResponseID续写 prompt, err := a.buildBasePromptResponse(ctx, chatData, req) if err != nil { return "", "", err } - cache, err := a.hsyq.CreateResponse(ctx, projectInfo.ModelInfo.Key, projectInfo.ModelInfo.ChatModel, prompt, "", true) - if err != nil { - return "", "", err - } - contextCache = cache.Id - promptJson = pkg.JsonStringIgonErr(prompt) - case constants.ModeContext: - prompt, err := a.buildBasePromptContext(ctx, chatData, req) - if err != nil { - return "", "", err - } - contextCache, err = a.hsyq.CreateContextCache(ctx, projectInfo.ModelInfo.Key, projectInfo.ModelInfo.ChatModel, prompt) + cache, err := a.openai.CreateResponseMessages(ctx, projectInfo.ModelInfo.Key, projectInfo.ModelInfo.URL, projectInfo.ModelInfo.ChatModel, prompt, "") if err != nil { return "", "", err } + contextCache = cache.ID promptJson = pkg.JsonStringIgonErr(prompt) default: return "", "", fmt.Errorf("未知的mode类型:%d", projectInfo.ModelInfo.Mode) @@ -133,24 +122,28 @@ func (a *AdviceChatBiz) Chat(ctx context.Context, chat *entitys.AdvicerChatReq) if modelInfo.SupID == 0 { return assistant, errors.New("未找到模型信息") } - //basePromptJson, err := a.getChatDataFromStringSessionId(ctx, chat.SessionId) - //if err != nil { - // return nil, err - //} chatHis, err := a.getChatHis(ctx, session.SessionID, 6) + if err != nil { + return assistant, err + } prompt, err := a.buildChatPromptResponse(ctx, chat, &session, chatHis) if err != nil { return assistant, err } - resContent, err := a.callLlmResponse(ctx, prompt, modelInfo.Key, modelInfo.ChatModel, session.ContextCache) + resContent, err := a.callLlmResponse(ctx, prompt, modelInfo.Key, modelInfo.URL, modelInfo.ChatModel, session.ContextCache) if err != nil { return assistant, err } - result := resContent.Output[0].GetOutputMessage().Content[0].GetText().GetText() + result := resContent.GetOutputText() if err = json.Unmarshal([]byte(result), &assistant); err != nil { return assistant, err } + var inToken, outToken int64 + if resContent.Usage != nil { + inToken = int64(resContent.Usage.InputTokens) + outToken = int64(resContent.Usage.OutputTokens) + } chatCtx, cancel := context.WithCancel(context.Background()) go func(session dbmodel.AiAdviceSession) { defer cancel() @@ -158,8 +151,8 @@ func (a *AdviceChatBiz) Chat(ctx context.Context, chat *entitys.AdvicerChatReq) SessionId: chat.SessionId, User: chat.Content, Assistant: assistant, - InToken: resContent.Usage.InputTokens, - OutToken: resContent.Usage.OutputTokens, + InToken: inToken, + OutToken: outToken, CreatAt: time.Now(), }) if assistant.MissionStatus == "fail" || assistant.MissionStatus == "complete" { @@ -173,34 +166,21 @@ func (a *AdviceChatBiz) Chat(ctx context.Context, chat *entitys.AdvicerChatReq) return } -func (a *AdviceChatBiz) buildChatPromptResponse(ctx context.Context, chat *entitys.AdvicerChatReq, session *dbmodel.AiAdviceSession, chatList []mongo_model.AdvicerChatHisMongoEntity) ([]*responses.InputItem, error) { - - var message = make([]*responses.InputItem, 3) - message[0] = &responses.InputItem{ - Union: &responses.InputItem_EasyMessage{ - EasyMessage: &responses.ItemEasyMessage{ - Role: responses.MessageRole_system, - Content: &responses.MessageContent{Union: &responses.MessageContent_StringValue{StringValue: a.taskPrompt(session)}}, - }, +func (a *AdviceChatBiz) buildChatPromptResponse(ctx context.Context, chat *entitys.AdvicerChatReq, session *dbmodel.AiAdviceSession, chatList []mongo_model.AdvicerChatHisMongoEntity) ([]openai.ResponseInputMessage, error) { + message := []openai.ResponseInputMessage{ + { + Role: openai.ChatMessageRoleSystem, + Content: a.taskPrompt(session), + }, + { + Role: openai.ChatMessageRoleSystem, + Content: "历史聊天记录:\n" + pkg.JsonStringIgonErr(chatList), + }, + { + Role: openai.ChatMessageRoleUser, + Content: chat.Content, }, } - message[1] = &responses.InputItem{ - Union: &responses.InputItem_EasyMessage{ - EasyMessage: &responses.ItemEasyMessage{ - Role: responses.MessageRole_system, - Content: &responses.MessageContent{Union: &responses.MessageContent_StringValue{StringValue: "历史聊天记录:\n" + pkg.JsonStringIgonErr(chatList)}}, - }, - }, - } - message[2] = &responses.InputItem{ - Union: &responses.InputItem_EasyMessage{ - EasyMessage: &responses.ItemEasyMessage{ - Role: responses.MessageRole_user, - Content: &responses.MessageContent{Union: &responses.MessageContent_StringValue{StringValue: chat.Content}}, - }, - }, - } - return message, nil } func (a *AdviceChatBiz) getChatHis(ctx context.Context, sessionId string, limit int64) (chatList []mongo_model.AdvicerChatHisMongoEntity, err error) { @@ -225,100 +205,29 @@ func (a *AdviceChatBiz) getChatHis(ctx context.Context, sessionId string, limit return } -func (a *AdviceChatBiz) buildChatPrompt(ctx context.Context, chat *entitys.AdvicerChatReq, session *dbmodel.AiAdviceSession, modelInfo *dbmodel.AiAdviceModelSup) (model.ContextChatCompletionRequest, error) { - var message = make([]*model.ChatCompletionMessage, 2) - message[0] = &model.ChatCompletionMessage{ - Role: model.ChatMessageRoleUser, - Content: &model.ChatCompletionMessageContent{ - StringValue: volcengine.String(chat.Content), - }, - } - message[1] = &model.ChatCompletionMessage{ - Role: model.ChatMessageRoleAssistant, - Content: &model.ChatCompletionMessageContent{ - StringValue: volcengine.String(a.taskPrompt(session)), - }, - } - - req := model.ContextChatCompletionRequest{ - ContextID: session.ContextCache, - Model: modelInfo.ChatModel, - Messages: message, - Stream: false, - } - return req, nil -} - func (a *AdviceChatBiz) taskPrompt(session *dbmodel.AiAdviceSession) string { - //type mission struct { - // missionName string - // status string - // missionCompleteDesc string - //} - //var m = &mission{ - // missionName: session.Mission, - // status: pkg.Ter(session.MissionStatus == 1, "进行中", "已完成"), - // missionCompleteDesc: session.MissionCompleteDesc, - //} - //missionJon, _ := json.Marshal(m) return "[当前时间]" + time.Now().Format("2006-01-02 15:04:05") } -func (a *AdviceChatBiz) buildBasePromptContext(ctx context.Context, chatData *entitys.ChatData, req *entitys.AdvicerChatRegistReq) ([]*model.ChatCompletionMessage, error) { - var message = make([]*model.ChatCompletionMessage, 2) - message[0] = &model.ChatCompletionMessage{ - Role: model.ChatMessageRoleSystem, - Content: &model.ChatCompletionMessageContent{ - StringValue: volcengine.String(a.sysPrompt(chatData, req)), +func (a *AdviceChatBiz) buildBasePromptResponse(ctx context.Context, chatData *entitys.ChatData, req *entitys.AdvicerChatRegistReq) ([]openai.ResponseInputMessage, error) { + message := []openai.ResponseInputMessage{ + { + Role: openai.ChatMessageRoleSystem, + Content: a.sysPrompt(chatData, req), }, - } - message[1] = &model.ChatCompletionMessage{ - Role: model.ChatMessageRoleSystem, - Content: &model.ChatCompletionMessageContent{ - StringValue: volcengine.String(a.assistantPrompt(chatData)), + { + Role: openai.ChatMessageRoleSystem, + Content: a.assistantPrompt(chatData), }, } return message, nil } -func (a *AdviceChatBiz) buildBasePromptResponse(ctx context.Context, chatData *entitys.ChatData, req *entitys.AdvicerChatRegistReq) ([]*responses.InputItem, error) { - var message = make([]*responses.InputItem, 2) - message[0] = &responses.InputItem{ - Union: &responses.InputItem_EasyMessage{ - EasyMessage: &responses.ItemEasyMessage{ - Role: responses.MessageRole_system, - Content: &responses.MessageContent{Union: &responses.MessageContent_StringValue{StringValue: a.sysPrompt(chatData, req)}}, - }, - }, - } - message[1] = &responses.InputItem{ - Union: &responses.InputItem_EasyMessage{ - EasyMessage: &responses.ItemEasyMessage{ - Role: responses.MessageRole_system, - Content: &responses.MessageContent{Union: &responses.MessageContent_StringValue{StringValue: a.assistantPrompt(chatData)}}, - }, - }, - } - return message, nil -} - -func (a *AdviceChatBiz) setContent(ctx context.Context, basePromptJson string, content string, session *dbmodel.AiAdviceSession) ([]*model.ChatCompletionMessage, error) { - promptJson := strings.ReplaceAll(basePromptJson, "{{chat_content}}", content) - var basePrompt []*model.ChatCompletionMessage - err := json.Unmarshal([]byte(promptJson), &basePrompt) - if err != nil { - return nil, err - } - - return basePrompt, nil -} - func (a *AdviceChatBiz) sysPrompt(chatData *entitys.ChatData, req *entitys.AdvicerChatRegistReq) string { var prompt strings.Builder prompt.WriteString(constants.BasePrompt) prompt.WriteString(req.Mission) prompt.WriteString(constants.BasePrompt2) - return prompt.String() } @@ -327,26 +236,8 @@ func (a *AdviceChatBiz) assistantPrompt(chatData *entitys.ChatData) string { return pkg.JsonStringIgonErr(chatData) } -func (a *AdviceChatBiz) getChatDataFromStringSessionId(ctx context.Context, sessionId string) (basePromptJson string, err error) { - cache := a.rdb.Rdb.Get(ctx, sessionId) - if cache.Err() != nil { - err = cache.Err() - return - } - - return cache.Val(), cache.Err() -} - -func (a *AdviceChatBiz) callLlm(ctx context.Context, request model.ContextChatCompletionRequest, key string) (string, error) { - res, err := a.hsyq.ChatWithRequest(ctx, key, request) - if err != nil { - return "", err - } - return *res.Choices[0].Message.Content.StringValue, nil -} - -func (a *AdviceChatBiz) callLlmResponse(ctx context.Context, request []*responses.InputItem, key string, modelName string, id string) (*responses.ResponseObject, error) { - res, err := a.hsyq.CreateResponse(ctx, key, modelName, request, id, false) +func (a *AdviceChatBiz) callLlmResponse(ctx context.Context, request []openai.ResponseInputMessage, key string, url string, modelName string, id string) (*openai.CreateResponseResponse, error) { + res, err := a.openai.CreateResponseMessages(ctx, key, url, modelName, request, id) if err != nil { return nil, err } diff --git a/internal/biz/advice_client.go b/internal/biz/advice_client.go index 391f585..38cb4e2 100644 --- a/internal/biz/advice_client.go +++ b/internal/biz/advice_client.go @@ -73,11 +73,11 @@ func (a *AdviceClientBiz) Update(ctx context.Context, param *entitys.AdvicerrCli return res.Err() } -func (a *AdviceClientBiz) List(ctx context.Context, param *entitys.AdvicerClientListReq) (list []mongo_model.AdvicerClientMongo, err error) { +func (a *AdviceClientBiz) List(ctx context.Context, param *entitys.AdvicerClientListReq) (list []mongo_model.AdvicerClientItem, err error) { filter := bson.M{} // 1. advicer_id 条件 if param.AdvicerId != 0 { - filter["AdvicerId"] = param.AdvicerId + filter["advicerId"] = param.AdvicerId } if param.ProjectId != 0 { @@ -99,7 +99,7 @@ func (a *AdviceClientBiz) List(ctx context.Context, param *entitys.AdvicerClient } // 遍历结果 for cursor.Next(ctx) { - var advicerVersion mongo_model.AdvicerClientMongo + var advicerVersion mongo_model.AdvicerClientItem if err := cursor.Decode(&advicerVersion); err != nil { return nil, err } diff --git a/internal/biz/advice_file.go b/internal/biz/advice_file.go index 6ccba98..669d824 100644 --- a/internal/biz/advice_file.go +++ b/internal/biz/advice_file.go @@ -12,19 +12,15 @@ import ( "os" "strings" "time" - - "github.com/volcengine/volcengine-go-sdk/service/arkruntime/model" - "github.com/volcengine/volcengine-go-sdk/service/arkruntime/model/responses" - "github.com/volcengine/volcengine-go-sdk/volcengine" ) type AdviceFileBiz struct { - hsyq *third_party.Hsyq + openai *third_party.OpenAi } -func NewAdviceFileBiz(hsyq *third_party.Hsyq) *AdviceFileBiz { +func NewAdviceFileBiz(openai *third_party.OpenAi) *AdviceFileBiz { return &AdviceFileBiz{ - hsyq: hsyq, + openai: openai, } } @@ -53,20 +49,22 @@ func (a *AdviceFileBiz) WordAna(ctx context.Context, wordContent string, project } timeSte := time.Now().Format("200601021504") dir := "./cache/" + timeSte - os.Mkdir(dir, 0755) + //缓存文件仅用于排查问题,写入失败不影响主流程 + _ = os.MkdirAll(dir, 0755) + //获取示例 examples := a.getAllExamples() //构建提示词 prompt := a.buildSimplePrompt(wordContent, examples) - os.WriteFile(dir+"/requset.json", []byte(prompt), 0644) + _ = os.WriteFile(dir+"/requset.json", []byte(prompt), 0644) //llm提取信息 anaContent, err := a.callLlm2(ctx, prompt, &projectInfo.ModelInfo) if err != nil { return nil, err } - os.WriteFile(dir+"/res.json", []byte(anaContent), 0644) + _ = os.WriteFile(dir+"/res.json", []byte(anaContent), 0644) //格式整理 data, err := a.parseResponse(ctx, []byte(anaContent)) @@ -76,8 +74,8 @@ func (a *AdviceFileBiz) WordAna(ctx context.Context, wordContent string, project //组装数据 resData := a.cateData(data) - os.WriteFile("./cache/"+timeSte+"/extracted.json", pkg.JsonByteIgonErr(resData), 0644) - return resData, err + _ = os.WriteFile(dir+"/extracted.json", pkg.JsonByteIgonErr(resData), 0644) + return resData, nil } func (a *AdviceFileBiz) cateData(data map[string]mongo_model.AdviceData) map[mongo_model.AdviceRole]map[string]mongo_model.AdviceData { @@ -113,8 +111,9 @@ func (a *AdviceFileBiz) parseResponse(ctx context.Context, responseByte []byte) return } for k, v := range result { + //跳过预定义之外的字段,避免模型输出的额外说明导致整体解析中断 if _, ok := DataMap[k]; !ok { - return + continue } var vbyte []byte if vbyte, err = json.Marshal(v); err != nil { @@ -123,56 +122,30 @@ func (a *AdviceFileBiz) parseResponse(ctx context.Context, responseByte []byte) newData := DataMap[k].Copy() if err = json.Unmarshal(vbyte, newData); err != nil { + err = fmt.Errorf("字段%s解析失败: %w", k, err) return } resultOutPut[k] = newData } + if len(resultOutPut) == 0 { + err = fmt.Errorf("响应中未包含有效数据") + return + } return } -//func (a *AdviceFileBiz) fixJson(ctx context.Context, json []byte) ([]byte, error) { -// prompt := "你是一个专业的JSON修复专家。请帮我修复以下错误的JSON格式。\n\n要求:\n1. 保持原有数据的结构和内容不变\n2. 修复JSON语法错误\n3. 输出格式化的正确JSON\n4. 简要说明修复了哪些问题\n\n错误的JSON:\n" + string(json) + "\n\n请直接输出修复后的JSON。" -// call, err := a.callLlm(ctx, prompt, jsonModel) -// if err != nil { -// return nil, err -// } -// -// return []byte(call), nil -//} - -func (a *AdviceFileBiz) callLlm(ctx context.Context, prompt string, modelInfo *dbmodel.AiAdviceModelSup) (string, error) { - var message = make([]*model.ChatCompletionMessage, 1) - message[0] = &model.ChatCompletionMessage{ - Role: model.ChatMessageRoleUser, - Content: &model.ChatCompletionMessageContent{ - StringValue: volcengine.String(prompt), - }, - } - res, err := a.hsyq.Chat(ctx, modelInfo.Key, modelInfo.FileModel, message) - if err != nil { - return "", err - } - - return *res.Choices[0].Message.Content.StringValue, nil -} - +// callLlm2 调用Responses接口提取文件信息 +// 走json_object格式输出,保证返回内容可直接反序列化 func (a *AdviceFileBiz) callLlm2(ctx context.Context, prompt string, modelInfo *dbmodel.AiAdviceModelSup) (string, error) { - var message = make([]*responses.InputItem, 3) - message[0] = &responses.InputItem{ - Union: &responses.InputItem_EasyMessage{ - EasyMessage: &responses.ItemEasyMessage{ - Role: responses.MessageRole_system, - Content: &responses.MessageContent{Union: &responses.MessageContent_StringValue{StringValue: prompt}}, - }, - }, - } - res, err := a.hsyq.CreateResponse(ctx, modelInfo.Key, modelInfo.FileModel, message, "", false) + content, err := a.openai.RequestResponsesJson(ctx, modelInfo.Key, modelInfo.URL, modelInfo.FileModel, prompt) if err != nil { - return "", err + return "", fmt.Errorf("文件分析模型调用失败: %w", err) } - - return res.Output[0].GetOutputMessage().Content[0].GetText().GetText(), nil + if strings.TrimSpace(content) == "" { + return "", fmt.Errorf("文件分析模型返回内容为空") + } + return content, nil } func (a *AdviceFileBiz) getAllExamples() map[string]mongo_model.AdviceData { diff --git a/internal/biz/advice_project.go b/internal/biz/advice_project.go index 9b1a433..743f93a 100644 --- a/internal/biz/advice_project.go +++ b/internal/biz/advice_project.go @@ -1,11 +1,14 @@ package biz import ( + errorcode "ai_scheduler/internal/data/error" "ai_scheduler/internal/data/impl" "ai_scheduler/internal/data/model" "ai_scheduler/internal/data/mongo_model" "ai_scheduler/internal/entitys" "ai_scheduler/internal/pkg" + "ai_scheduler/tmpl/dataTemp" + "errors" "fmt" "time" @@ -13,6 +16,7 @@ import ( "go.mongodb.org/mongo-driver/bson" "go.mongodb.org/mongo-driver/bson/primitive" + "go.mongodb.org/mongo-driver/mongo" "xorm.io/builder" ) @@ -20,6 +24,7 @@ type AdviceProjectBiz struct { AdvicerProjectMongo *mongo_model.AdvicerProjectMongo adviceProjectImpl *impl.AdviceProjectImpl aiAdviceModelSupImpl *impl.AiAdviceModelSupImpl + industryImpl *impl.AdviceIndustryImpl mongo *pkg.Mongo } @@ -27,6 +32,7 @@ func NewAdviceProjectBiz( advicerProjectMongo *mongo_model.AdvicerProjectMongo, adviceProjectImpl *impl.AdviceProjectImpl, aiAdviceModelSupImpl *impl.AiAdviceModelSupImpl, + industryImpl *impl.AdviceIndustryImpl, mongo *pkg.Mongo, ) *AdviceProjectBiz { return &AdviceProjectBiz{ @@ -34,14 +40,21 @@ func NewAdviceProjectBiz( mongo: mongo, adviceProjectImpl: adviceProjectImpl, aiAdviceModelSupImpl: aiAdviceModelSupImpl, + industryImpl: industryImpl, } } +// BaseAdd 新建项目基础信息;IndustryId>0 时自动复制行业模板的维度内容 func (a *AdviceProjectBiz) BaseAdd(ctx context.Context, param *entitys.AdvicerProjectBaseAddReq) (res *entitys.AdvicerProjectBaseAddRes, err error) { add := &model.AiAdviceProject{ Name: param.Name, ModelSupID: param.ModelSupId, } + if param.IndustryId > 0 { + if err = a.fillTemplateFromIndustry(ctx, add, param.IndustryId); err != nil { + return nil, err + } + } err = a.adviceProjectImpl.AddWithData(add) if err != nil { return nil, err @@ -51,19 +64,127 @@ func (a *AdviceProjectBiz) BaseAdd(ctx context.Context, param *entitys.AdvicerPr }, err } +// BaseUpdate 更新项目基础信息与项目级模板字段(仅更新传入非空字段) func (a *AdviceProjectBiz) BaseUpdate(ctx context.Context, param *entitys.AdvicerProjectBaseUpdateReq) (err error) { if param.ProjectId == 0 { - return + return errorcode.ParamErr("projectId is empty") + } + updates := make(map[string]interface{}) + if param.Name != "" { + updates["name"] = param.Name + } + if param.ModelSupId != 0 { + updates["model_sup_id"] = param.ModelSupId + } + if param.Desc != "" { + updates["template_desc"] = param.Desc + } + if param.AdvicerDesc != "" { + updates["template_advicer_desc"] = param.AdvicerDesc + } + if param.ClientDimension != "" { + updates["client_dimension"] = param.ClientDimension + } + if param.ProjectDimension != "" { + updates["project_dimension"] = param.ProjectDimension + } + if param.AdvicerDimension != "" { + updates["advicer_dimension"] = param.AdvicerDimension + } + if param.TalkSkillDimension != "" { + updates["talk_skill_dimension"] = param.TalkSkillDimension + } + if param.RuleDimension != "" { + updates["rule_dimension"] = param.RuleDimension + } + // IndustryId>0 时重新从行业模板复制(覆盖模板字段) + if param.IndustryId > 0 { + var industry model.AiAdviceIndustryTemp + if err = a.industryImpl.GetByKey(ctx, "industry_id", param.IndustryId, &industry); err != nil { + return err + } + if industry.IndustryId == 0 { + return errorcode.ParamErr("行业模板不存在") + } + a.mergeIndustryTemplate(updates, &industry) + } + if len(updates) == 0 { + return nil } cond := builder.NewCond() cond = cond.And(builder.Eq{"project_id": param.ProjectId}) - err = a.adviceProjectImpl.UpdateByCond(&cond, &model.AiAdviceProject{ - Name: param.Name, - ModelSupID: param.ModelSupId, - }) + err = a.adviceProjectImpl.UpdateByCond(&cond, updates) return err } +// TemplateCopy 重新应用行业模板(全量覆盖项目模板字段) +func (a *AdviceProjectBiz) TemplateCopy(ctx context.Context, param *entitys.AdvicerProjectTemplateCopyReq) (err error) { + if param.ProjectId == 0 { + return errorcode.ParamErr("projectId is empty") + } + if param.IndustryId == 0 { + return errorcode.ParamErr("industryId is empty") + } + var industry model.AiAdviceIndustryTemp + if err = a.industryImpl.GetByKey(ctx, "industry_id", param.IndustryId, &industry); err != nil { + return err + } + if industry.IndustryId == 0 { + return errorcode.ParamErr("行业模板不存在") + } + updates := make(map[string]interface{}) + a.mergeIndustryTemplate(updates, &industry) + cond := builder.NewCond() + cond = cond.And(builder.Eq{"project_id": param.ProjectId}) + return a.adviceProjectImpl.UpdateByCond(&cond, updates) +} + +// fillTemplateFromIndustry 从行业模板复制维度内容到项目实体 +func (a *AdviceProjectBiz) fillTemplateFromIndustry(ctx context.Context, project *model.AiAdviceProject, industryId int32) error { + var industry model.AiAdviceIndustryTemp + if err := a.industryImpl.GetByKey(ctx, "industry_id", industryId, &industry); err != nil { + return err + } + if industry.IndustryId == 0 { + return errorcode.ParamErr("行业模板不存在") + } + project.IndustryID = industry.IndustryId + project.TemplateDesc = industry.Desc + project.TemplateAdvicerDesc = industry.AdvicerDesc + project.ClientDimension = industry.ClientDimension + project.ProjectDimension = industry.ProjectDimension + project.AdvicerDimension = industry.AdvicerDimension + project.TalkSkillDimension = industry.TalkSkillDimension + project.RuleDimension = industry.RuleDimension + return nil +} + +// mergeIndustryTemplate 将行业模板内容合并到更新 map(全量覆盖语义) +func (a *AdviceProjectBiz) mergeIndustryTemplate(updates map[string]interface{}, industry *model.AiAdviceIndustryTemp) { + updates["industry_id"] = industry.IndustryId + updates["template_desc"] = industry.Desc + updates["template_advicer_desc"] = industry.AdvicerDesc + updates["client_dimension"] = industry.ClientDimension + updates["project_dimension"] = industry.ProjectDimension + updates["advicer_dimension"] = industry.AdvicerDimension + updates["talk_skill_dimension"] = industry.TalkSkillDimension + updates["rule_dimension"] = industry.RuleDimension +} + +// List 分页查询项目列表 +func (a *AdviceProjectBiz) List(ctx context.Context, param *entitys.AdvicerProjectListReq) (list []model.AiAdviceProject, total int64, err error) { + list = make([]model.AiAdviceProject, 0) + cond := builder.NewCond() + if param.Name != "" { + cond = cond.And(builder.Like{"name", "%" + param.Name + "%"}) + } + page, err := a.adviceProjectImpl.GetListToStruct(ctx, &cond, &dataTemp.ReqPageBo{Page: param.Page, Limit: param.PageSize}, &list, "") + if err != nil { + return nil, 0, err + } + return list, page.Total, nil +} + func (a *AdviceProjectBiz) Add(ctx context.Context, param *entitys.AdvicerProjectAddReq) (id interface{}, err error) { res, err := a.mongo.Co(a.AdvicerProjectMongo).InsertOne(ctx, &mongo_model.AdvicerProjectMongo{ @@ -115,7 +236,12 @@ func (a *AdviceProjectBiz) Info(ctx context.Context, param *entitys.AdvicerProje if err != nil { return nil, err } - baseInfo, err := a.BaseInfo(configInfo.ProjectId) + baseProjectId := configInfo.ProjectId + // mongo 中暂无配置记录时,回退用请求参数中的 projectId 查询基础信息 + if baseProjectId == 0 { + baseProjectId = param.ProjectId + } + baseInfo, err := a.BaseInfo(baseProjectId) if err != nil { return nil, err } @@ -173,12 +299,12 @@ func (a *AdviceProjectBiz) ConfigInfo(ctx context.Context, param *entitys.Advice } res := a.mongo.Co(a.AdvicerProjectMongo).FindOne(ctx, filter) - if res.Err() != nil { + if res.Err() != nil && !errors.Is(res.Err(), mongo.ErrNoDocuments) { return info, res.Err() } // 遍历结果 - if err := res.Decode(&info); err != nil { + if err := res.Decode(&info); err != nil && !errors.Is(err, mongo.ErrNoDocuments) { return info, err } diff --git a/internal/biz/advice_skill.go b/internal/biz/advice_skill.go index eeb682f..410b84f 100644 --- a/internal/biz/advice_skill.go +++ b/internal/biz/advice_skill.go @@ -75,11 +75,11 @@ func (a *AdviceSkillBiz) VersionUpdate(ctx context.Context, param *entitys.Advic return res.Err() } -func (a *AdviceSkillBiz) VersionList(ctx context.Context, param *entitys.AdvicerTalkSkillListReq) (list []mongo_model.AdvicerTalkSkillMongo, err error) { +func (a *AdviceSkillBiz) VersionList(ctx context.Context, param *entitys.AdvicerTalkSkillListReq) (list []mongo_model.AdvicerTalkSkillItem, err error) { filter := bson.M{} // 1. advicer_id 条件 if param.AdvicerId != 0 { - filter["AdvicerId"] = param.AdvicerId + filter["advicerId"] = param.AdvicerId } if param.ProjectId != 0 { @@ -112,7 +112,7 @@ func (a *AdviceSkillBiz) VersionList(ctx context.Context, param *entitys.Advicer } // 遍历结果 for cursor.Next(ctx) { - var advicerVersion mongo_model.AdvicerTalkSkillMongo + var advicerVersion mongo_model.AdvicerTalkSkillItem if err := cursor.Decode(&advicerVersion); err != nil { return nil, err } diff --git a/internal/biz/llm_service/third_party/hsyq.go b/internal/biz/llm_service/third_party/hsyq.go deleted file mode 100644 index 8b73b9f..0000000 --- a/internal/biz/llm_service/third_party/hsyq.go +++ /dev/null @@ -1,142 +0,0 @@ -package third_party - -import ( - "context" - - "time" - - "github.com/gofiber/fiber/v2/log" - "github.com/volcengine/volcengine-go-sdk/service/arkruntime" - "github.com/volcengine/volcengine-go-sdk/service/arkruntime/model" - "github.com/volcengine/volcengine-go-sdk/service/arkruntime/model/responses" - "github.com/volcengine/volcengine-go-sdk/volcengine" -) - -type Hsyq struct { - mapClient map[string]*arkruntime.Client -} - -func NewHsyq() *Hsyq { - return &Hsyq{ - mapClient: make(map[string]*arkruntime.Client), - } -} - -func (h *Hsyq) getClient(key string) *arkruntime.Client { - var client *arkruntime.Client - if _, ok := h.mapClient[key]; ok { - client = h.mapClient[key] - } else { - client = arkruntime.NewClientWithApiKey( - key, - arkruntime.WithRegion("cn-beijing"), - arkruntime.WithTimeout(2*time.Minute), - arkruntime.WithRetryTimes(2), - ) - h.mapClient[key] = client - } - return client -} - -// 火山引擎 -func (h *Hsyq) Chat(ctx context.Context, key string, modelName string, prompt []*model.ChatCompletionMessage) (model.ChatCompletionResponse, error) { - req := model.CreateChatCompletionRequest{ - Model: modelName, - Messages: prompt, - Stream: new(bool), - Thinking: &model.Thinking{Type: model.ThinkingTypeDisabled}, - } - - resp, err := h.getClient(key).CreateChatCompletion(ctx, req) - if err != nil { - return model.ChatCompletionResponse{ID: ""}, err - } - log.Info("token用量:", resp.Usage.TotalTokens, "输入:", resp.Usage.PromptTokens, "输出:", resp.Usage.CompletionTokens) - return resp, err -} - -// 火山引擎 -func (h *Hsyq) ChatWithRequest(ctx context.Context, key string, request model.ContextChatCompletionRequest) (model.ChatCompletionResponse, error) { - - resp, err := h.getClient(key).CreateContextChatCompletion(ctx, request) - if err != nil { - return model.ChatCompletionResponse{ID: ""}, err - } - log.Info("token用量:", resp.Usage.TotalTokens, "输入:", resp.Usage.PromptTokens, "输出:", resp.Usage.CompletionTokens) - return resp, err -} - -func (h *Hsyq) CreateContextCache(ctx context.Context, key string, modelName string, prompt []*model.ChatCompletionMessage) (string, error) { - req := model.CreateContextRequest{ - Model: modelName, - Messages: prompt, - TTL: volcengine.Int(3600), - Mode: model.ContextModeSession, - TruncationStrategy: &model.TruncationStrategy{Type: model.TruncationStrategyTypeRollingTokens}, - } - - resp, err := h.getClient(key).CreateContext(ctx, req) - if err != nil { - return "", err - } - log.Info("token用量:", resp.Usage.TotalTokens, "输入:", resp.Usage.PromptTokens, "输出:", resp.Usage.CompletionTokens) - return resp.ID, err -} - -func (h *Hsyq) CreateResponse(ctx context.Context, key string, modelName string, prompt []*responses.InputItem, id string, isRegis bool) (*responses.ResponseObject, error) { - - req := &responses.ResponsesRequest{ - Model: modelName, - Input: &responses.ResponsesInput{ - Union: &responses.ResponsesInput_ListValue{ - ListValue: &responses.InputItemList{ListValue: prompt}, - }, - }, - Stream: new(bool), - Reasoning: &responses.ResponsesReasoning{Effort: responses.ReasoningEffort_minimal}, - Thinking: &responses.ResponsesThinking{Type: responses.ThinkingType_disabled.Enum()}, - Text: &responses.ResponsesText{Format: &responses.TextFormat{Type: responses.TextType_json_object}}, - } - if isRegis { - prefix := true - req.Caching = &responses.ResponsesCaching{Type: responses.CacheType_enabled.Enum(), Prefix: &prefix} - req.ExpireAt = volcengine.Int64(time.Now().Unix() + 3600) - } - if len(id) != 0 { - req.PreviousResponseId = &id - //req.Text = &responses.ResponsesText{ - // Format: &responses.TextFormat{ - // Type: responses.TextType_json_object, - // Schema: - // } - //} - } - resp, err := h.getClient(key).CreateResponses(ctx, req) - if err != nil { - return nil, err - } - log.Info("token用量:", resp.Usage.TotalTokens, "输入:", resp.Usage.InputTokens, "输出:", resp.Usage.OutputTokens) - return resp, err -} - -func (h *Hsyq) RequestHsyqJson(ctx context.Context, key string, modelName string, prompt []*responses.InputItem) (*responses.ResponseObject, error) { - req := responses.ResponsesRequest{ - Model: modelName, - Input: &responses.ResponsesInput{ - Union: &responses.ResponsesInput_ListValue{ - ListValue: &responses.InputItemList{ListValue: prompt}, - }, - }, - Stream: new(bool), - - Thinking: &responses.ResponsesThinking{Type: responses.ThinkingType_disabled.Enum()}, - Text: &responses.ResponsesText{Format: &responses.TextFormat{Type: responses.TextType_json_object}}, - } - - resp, err := h.getClient(key).CreateResponses(ctx, &req) - if err != nil { - return resp, err - } - log.Info("token用量:", resp.Usage.TotalTokens) - return resp, err -} diff --git a/internal/biz/llm_service/third_party/openai.go b/internal/biz/llm_service/third_party/openai.go index 6782009..26dabbb 100644 --- a/internal/biz/llm_service/third_party/openai.go +++ b/internal/biz/llm_service/third_party/openai.go @@ -84,6 +84,43 @@ func (o *OpenAi) CreateResponse( return &resp, nil } +// CreateResponseMessages 对应 Hsyq.CreateResponse 的多消息版本 +// json_object格式输出;响应需存储后才能通过 PreviousResponseID 续写 +// id 不为空时走 PreviousResponseID 续写对话 +func (o *OpenAi) CreateResponseMessages( + ctx context.Context, + key string, + url string, + modelName string, + input []openai.ResponseInputMessage, + id string, +) (*openai.CreateResponseResponse, error) { + store := true + req := openai.CreateResponseRequest{ + Model: modelName, + Input: input, + Store: &store, + Text: &openai.ResponseTextConfig{ + Format: &openai.ResponseTextFormat{ + Type: "json_object", + }, + }, + } + + if len(id) != 0 { + req.PreviousResponseID = id + } + + resp, err := o.getClient(key, url).CreateResponse(ctx, req) + if err != nil { + return nil, err + } + if resp.Usage != nil { + log.Info("token用量:", resp.Usage.TotalTokens, "输入:", resp.Usage.InputTokens, "输出:", resp.Usage.OutputTokens) + } + return &resp, nil +} + // RequestResponsesJson 对应 Hsyq.RequestHsyqJson // 通过 Text.Format 指定 JSON 输出,直接返回文字内容 func (o *OpenAi) RequestResponsesJson( diff --git a/internal/biz/provider_set.go b/internal/biz/provider_set.go index 5cf6729..5ab75d7 100644 --- a/internal/biz/provider_set.go +++ b/internal/biz/provider_set.go @@ -7,12 +7,12 @@ import ( ) var ProviderSetBiz = wire.NewSet( - + third_party.NewOpenAi, NewAdviceFileBiz, - third_party.NewHsyq, NewAdviceAdvicerBiz, NewAdviceSkillBiz, NewAdviceProjectBiz, + NewAdviceModelSupBiz, NewAdviceClientBiz, NewAdviceChatBiz, NewAdvicerIndustryBiz, diff --git a/internal/data/impl/advice_admin_impl.go b/internal/data/impl/advice_admin_impl.go index c18885d..304652e 100644 --- a/internal/data/impl/advice_admin_impl.go +++ b/internal/data/impl/advice_admin_impl.go @@ -30,7 +30,7 @@ func (a *AdviceAdminImpl) GetByAccount(ctx context.Context, account string) (*mo var admin model.AiAdviceAdmin err := a.Db.WithContext(ctx). Model(&model.AiAdviceAdmin{}). - Where("account = ?", account). + Where("account = ? and type=1", account). First(&admin).Error if err != nil { return nil, err diff --git a/internal/data/model/ai_advice_project.gen.go b/internal/data/model/ai_advice_project.gen.go index 15573c6..a57fa4b 100644 --- a/internal/data/model/ai_advice_project.gen.go +++ b/internal/data/model/ai_advice_project.gen.go @@ -12,10 +12,18 @@ const TableNameAiAdviceProject = "ai_advice_project" // AiAdviceProject mapped from table type AiAdviceProject struct { - ProjectID int32 `gorm:"column:project_id;primaryKey;autoIncrement:true" json:"project_id"` - Name string `gorm:"column:name;not null;comment:姓名" json:"name"` // 姓名 - ModelSupID int32 `gorm:"column:model_sup_id;not null;comment:模型提供方配置,关联advicer_model_sup" json:"model_sup_id"` // 模型提供方配置,关联advicer_model_sup - CreateAt time.Time `gorm:"column:create_at;default:CURRENT_TIMESTAMP" json:"create_at"` + ProjectID int32 `gorm:"column:project_id;primaryKey;autoIncrement:true" json:"project_id"` + Name string `gorm:"column:name;not null;comment:姓名" json:"name"` // 姓名 + ModelSupID int32 `gorm:"column:model_sup_id;not null;comment:模型提供方配置,关联advicer_model_sup" json:"model_sup_id"` // 模型提供方配置,关联advicer_model_sup + IndustryID int32 `gorm:"column:industry_id;not null;default:0;comment:来源行业模板ID" json:"industry_id"` // 来源行业模板ID(ai_advice_industry_temp.industry_id),0=未关联 + TemplateDesc string `gorm:"column:template_desc;type:text;comment:行业/项目介绍(模板副本)" json:"desc"` // 行业/项目介绍(模板副本) + TemplateAdvicerDesc string `gorm:"column:template_advicer_desc;type:text;comment:销售作用说明(模板副本)" json:"advicer_desc"` // 销售作用说明(模板副本) + ClientDimension string `gorm:"column:client_dimension;type:text;comment:客户维度(模板副本)" json:"client_dimension"` // 客户维度(模板副本) + ProjectDimension string `gorm:"column:project_dimension;type:text;comment:项目维度(模板副本)" json:"project_dimension"` // 项目维度(模板副本) + AdvicerDimension string `gorm:"column:advicer_dimension;type:text;comment:顾问维度(模板副本)" json:"advicer_dimension"` // 顾问维度(模板副本) + TalkSkillDimension string `gorm:"column:talk_skill_dimension;type:text;comment:聊天技巧维度(模板副本)" json:"talk_skill_dimension"` // 聊天技巧维度(模板副本) + RuleDimension string `gorm:"column:rule_dimension;type:text;comment:风控维度(模板副本)" json:"rule_dimension"` // 风控维度(模板副本) + CreateAt time.Time `gorm:"column:create_at;default:CURRENT_TIMESTAMP" json:"create_at"` } // TableName AiAdviceProject's table name diff --git a/internal/data/mongo_model/common.go b/internal/data/mongo_model/common.go index 3b20438..21d8f8a 100644 --- a/internal/data/mongo_model/common.go +++ b/internal/data/mongo_model/common.go @@ -1,5 +1,27 @@ package mongo_model +import "go.mongodb.org/mongo-driver/bson/primitive" + +// 以下 Item 包装类型用于列表查询返回记录 _id(原始模型不包含 _id 字段) + +// AdvicerVersionItem 销售版本(含 _id) +type AdvicerVersionItem struct { + Id primitive.ObjectID `bson:"_id" json:"id"` + AdvicerVersionMongo `bson:",inline"` +} + +// AdvicerTalkSkillItem 聊天技巧(含 _id) +type AdvicerTalkSkillItem struct { + Id primitive.ObjectID `bson:"_id" json:"id"` + AdvicerTalkSkillMongo `bson:",inline"` +} + +// AdvicerClientItem 客户(含 _id) +type AdvicerClientItem struct { + Id primitive.ObjectID `bson:"_id" json:"id"` + AdvicerClientMongo `bson:",inline"` +} + type AdviceRole string const ( diff --git a/internal/entitys/advicer_data.go b/internal/entitys/advicer_data.go index 78383d6..d791c87 100644 --- a/internal/entitys/advicer_data.go +++ b/internal/entitys/advicer_data.go @@ -103,6 +103,8 @@ type AdvicerTalkSkillInfoReq struct { type AdvicerProjectBaseAddReq struct { Name string `json:"name"` ModelSupId int32 `json:"modelSupId"` + // IndustryId 来源行业模板ID(>0 时自动复制该行业模板的维度内容) + IndustryId int32 `json:"industryId"` } type AdvicerProjectBaseAddRes struct { @@ -113,6 +115,29 @@ type AdvicerProjectBaseUpdateReq struct { ProjectId int32 `json:"projectId"` Name string `json:"name"` ModelSupId int32 `json:"modelSupId"` + // IndustryId >0 时重新从该行业模板复制维度内容(覆盖现有模板字段) + IndustryId int32 `json:"industryId"` + // 以下为项目级模板字段(可选,传入非空值时更新) + Desc string `json:"desc"` + AdvicerDesc string `json:"advicer_desc"` + ClientDimension string `json:"client_dimension"` + ProjectDimension string `json:"project_dimension"` + AdvicerDimension string `json:"advicer_dimension"` + TalkSkillDimension string `json:"talk_skill_dimension"` + RuleDimension string `json:"rule_dimension"` +} + +// AdvicerProjectListReq 查询项目列表 +type AdvicerProjectListReq struct { + Name string `json:"name"` + Page int `json:"page"` + PageSize int `json:"page_size"` +} + +// AdvicerProjectTemplateCopyReq 重新应用行业模板(覆盖项目模板字段) +type AdvicerProjectTemplateCopyReq struct { + ProjectId int32 `json:"projectId"` + IndustryId int32 `json:"industryId"` } type AdvicerProjectAddReq struct { diff --git a/internal/server/http.go b/internal/server/http.go index bedfea6..35c9461 100644 --- a/internal/server/http.go +++ b/internal/server/http.go @@ -23,17 +23,23 @@ func NewHTTPServer( adviceClient *advice.ClientService, industry *advice.IndustryService, admin *advice.AdminService, + modelSup *advice.ModelSupService, ) *fiber.App { //构建 server app := initRoute() // 提供静态前端:/ -> web/index.html, /admin.html, /project.html + // 上传文件静态托管:/upload/...(聊天记录等) + app.Static("/upload", "./upload") app.Static("/", "./web") - router.SetupRoutes(cfg, app, adviceFile, adviceData, adviceChat, adviceProject, adviceTalkSkill, adviceClient, industry, admin) + router.SetupRoutes(cfg, app, adviceFile, adviceData, adviceChat, adviceProject, adviceTalkSkill, adviceClient, industry, admin, modelSup) return app } func initRoute() *fiber.App { - app := fiber.New() + app := fiber.New(fiber.Config{ + // 提升请求体上限,支持上传聊天记录文件(默认 4MB) + BodyLimit: 32 * 1024 * 1024, + }) app.Use( recover.New(), logger.New(), diff --git a/internal/server/router/advicer.go b/internal/server/router/advicer.go index 39f9045..9ecb30a 100644 --- a/internal/server/router/advicer.go +++ b/internal/server/router/advicer.go @@ -19,6 +19,7 @@ func AdvicerRouterRegist( adviceClient *advice.ClientService, industry *advice.IndustryService, admin *advice.AdminService, + modelSup *advice.ModelSupService, ) { advicer := r.Group("advice/admin") // 登录放行 @@ -29,6 +30,7 @@ func AdvicerRouterRegist( advicer.Post("del", Vali(admin.Del, &entitys.AdvicerAdminDelReq{})) advicer.Post("file/word/ana", adviceFile.WordAna) + advicer.Post("file/upload", adviceFile.Upload) //销售 advicer.Post("advicer/add", adviceData.AdvicerUpdate) advicer.Post("advicer/update", adviceData.AdvicerUpdate) @@ -50,6 +52,8 @@ func AdvicerRouterRegist( advicer.Post("project/info/add", adviceProject.Add) advicer.Post("project/info/update", adviceProject.Update) advicer.Post("project/info", adviceProject.Info) + advicer.Post("project/list", adviceProject.List) + advicer.Post("project/template/copy", adviceProject.TemplateCopy) //客户 advicer.Post("client/add", adviceClient.Add) @@ -67,4 +71,10 @@ func AdvicerRouterRegist( advicer.Post("industry/generate", Vali(industry.Generate, &entitys.AdvicerIndustryGenerateReq{})) advicer.Post("industry/list", Vali(industry.List, &entitys.AdvicerIndustryListReq{})) advicer.Post("industry/del", Vali(industry.Del, &entitys.AdvicerIndustryDelReq{})) + + //模型配置 + advicer.Post("modelsup/add", modelSup.Add) + advicer.Post("modelsup/update", modelSup.Update) + advicer.Post("modelsup/list", modelSup.List) + advicer.Post("modelsup/del", modelSup.Del) } diff --git a/internal/server/router/router.go b/internal/server/router/router.go index de618d2..3d38960 100644 --- a/internal/server/router/router.go +++ b/internal/server/router/router.go @@ -4,6 +4,7 @@ import ( "ai_scheduler/internal/config" errorcode "ai_scheduler/internal/data/error" errors "ai_scheduler/internal/data/error" + "ai_scheduler/internal/middleware" "ai_scheduler/internal/pkg" "ai_scheduler/internal/services/advice" "encoding/json" @@ -17,13 +18,13 @@ import ( // SetupRoutes 设置路由 func SetupRoutes(cfg *config.Config, app *fiber.App, adviceFile *advice.FileService, adviceData *advice.AdvicerService, adviceChat *advice.ChatService, adviceProject *advice.ProjectService, adviceTalkSkill *advice.TalkSkillService, adviceClient *advice.ClientService, - industry *advice.IndustryService, admin *advice.AdminService, + industry *advice.IndustryService, admin *advice.AdminService, modelSup *advice.ModelSupService, ) { app.Use(func(c *fiber.Ctx) error { // 设置 CORS 头 c.Set("Access-Control-Allow-Origin", "*") c.Set("Access-Control-Allow-Methods", "GET, POST, PUT, DELETE, OPTIONS") - c.Set("Access-Control-Allow-Headers", "Content-Type, Authorization") + c.Set("Access-Control-Allow-Headers", "Content-Type, Authorization, X-finder-TOKEN") // AI能力调用路由,设置不同的 CORS 头 if strings.HasPrefix(c.Path(), "/api/v1/capability") { @@ -46,8 +47,14 @@ func SetupRoutes(cfg *config.Config, app *fiber.App, adviceFile *advice.FileServ return nil }) + // 微信项目路由:/api/v1/project/wx/*(通过 X-finder-TOKEN 转发,不走 JWT) + projectR := r.Group("project") + WxRegist(projectR) + // 管理后台路由(需要登录): /api/v1/admin/... adminR := r.Group("admin") + // 登录接口放行,其余接口校验 JWT 登录态 + adminR.Use(middleware.AuthMiddleware(cfg.JwtSecret, "/api/v1/admin/advice/admin/login")) AdvicerRouterRegist( cfg, adminR, @@ -59,32 +66,11 @@ func SetupRoutes(cfg *config.Config, app *fiber.App, adviceFile *advice.FileServ adviceClient, industry, admin, - 1, // requiredType = 1 for management backend - "/api/v1/admin/advice/admin/login", + modelSup, ) - // 项目后台路由: /api/v1/project/... - projectR := r.Group("project") - AdvicerRouterRegist( - cfg, - projectR, - adviceFile, - adviceData, - adviceChat, - adviceProject, - adviceTalkSkill, - adviceClient, - industry, - admin, - 2, // requiredType = 2 for project backend - "/api/v1/project/advice/admin/login", - ) - - // 把微信能力注册到项目后台下 - WxRegist(projectR) } - func registerResponse(router fiber.Router) { // 自定义返回 router.Use(func(c *fiber.Ctx) error { @@ -121,7 +107,9 @@ func registerCommon(c *fiber.Ctx, err error) error { body := c.Response().Body() if c.Locals("skip_response_wrap") == true { - return c.JSON(string(body)) + // handler 已自行写入原始 JSON(如微信代理透传上游 data),原样输出不做包装 + c.Set(fiber.HeaderContentType, fiber.MIMEApplicationJSONCharsetUTF8) + return c.Send(body) } var rawData json.RawMessage if len(body) > 0 { diff --git a/internal/services/advice/advicer_test.go b/internal/services/advice/advicer_test.go index cd89b7a..03b7a79 100644 --- a/internal/services/advice/advicer_test.go +++ b/internal/services/advice/advicer_test.go @@ -123,15 +123,15 @@ func Run(ctx context.Context, reqBody []byte) { advicerTalkSkillMongo := mongo_model.NewAdvicerTalkSkillMongo() advicerClientMongo := mongo_model.NewAdvicerClientMongo() advicerProjectMongo := mongo_model.NewAdvicerProjectMongo() - hsyq := third_party.NewHsyq() + openai := third_party.NewOpenAi() adviceProjectImpl := impl.NewAdviceProjectImpl(db) mongo, _ := pkg.NewMongoDb(ctx, configConfig) adviceAdvicerBiz := biz.NewAdviceAdvicerBiz(advicerImpl, advicerVersionMongo, mongo) - adviceChatBiz := biz.NewAdviceChatBiz(hsyq, rdb, aiAdviceSessionImpl, aiAdviceModelSupImpl, advicerChatHisMongo, mongo) + adviceChatBiz := biz.NewAdviceChatBiz(openai, rdb, aiAdviceSessionImpl, aiAdviceModelSupImpl, advicerChatHisMongo, mongo) skillBiz := biz.NewAdviceSkillBiz(advicerTalkSkillMongo, mongo) clientBiz := biz.NewAdviceClientBiz(advicerClientMongo, mongo) - adviceFileBiz := biz.NewAdviceFileBiz(hsyq) + adviceFileBiz := biz.NewAdviceFileBiz(openai) adviceProjectBiz := biz.NewAdviceProjectBiz(advicerProjectMongo, adviceProjectImpl, aiAdviceModelSupImpl, mongo) adviceClientBiz := biz.NewAdviceClientBiz(advicerClientMongo, mongo) adviceSkillBiz := biz.NewAdviceSkillBiz(advicerTalkSkillMongo, mongo) diff --git a/internal/services/advice/file.go b/internal/services/advice/file.go index f2a7c7b..c61c8f7 100644 --- a/internal/services/advice/file.go +++ b/internal/services/advice/file.go @@ -9,10 +9,15 @@ import ( "ai_scheduler/internal/pkg/file_download" "context" "errors" + "os" + "path/filepath" + "strings" + "time" "net/url" "github.com/gofiber/fiber/v2" + "github.com/google/uuid" ) // FileService 文件处理 @@ -88,3 +93,41 @@ func (a *FileService) WordAnat(path string) ([]byte, error) { return pkg.JsonByteIgonErr(ana), err } + +// Upload 本地上传文件(聊天记录等),返回可直接访问的 URL +func (a *FileService) Upload(c *fiber.Ctx) error { + file, err := c.FormFile("file") + if err != nil { + return errorcode.ParamErr("file 不能为空") + } + if file.Size > 32*1024*1024 { + return errorcode.ParamErr("文件过大,最大支持 32MB") + } + ext := strings.ToLower(filepath.Ext(file.Filename)) + allowed := map[string]bool{ + ".doc": true, ".docx": true, ".pdf": true, ".txt": true, ".md": true, + ".xls": true, ".xlsx": true, ".csv": true, ".json": true, + } + if !allowed[ext] { + return errorcode.ParamErr("不支持的文件类型: " + ext) + } + // 按日期分目录存储 + day := time.Now().Format("20060102") + dir := filepath.Join("upload", day) + if err = os.MkdirAll(dir, 0755); err != nil { + return err + } + saveName := uuid.NewString() + ext + dst := filepath.Join(dir, saveName) + if err = c.SaveFile(file, dst); err != nil { + return err + } + relPath := "/" + filepath.ToSlash(dst) + res := &entitys.FileUploadRes{ + Path: relPath, + Url: c.Protocol() + "://" + c.Hostname() + relPath, + Name: file.Filename, + Size: file.Size, + } + return pkg.HandleResponse(c, res, nil) +} diff --git a/internal/services/advice/project.go b/internal/services/advice/project.go index 9c28c49..07f72f1 100644 --- a/internal/services/advice/project.go +++ b/internal/services/advice/project.go @@ -72,3 +72,23 @@ func (d *ProjectService) Info(c *fiber.Ctx) error { list, err := d.adviceProjectBiz.Info(c.UserContext(), req) return pkg.HandleResponse(c, list, err) } + +// List 项目分页列表 +func (d *ProjectService) List(c *fiber.Ctx) error { + req := &entitys.AdvicerProjectListReq{} + if err := c.BodyParser(req); err != nil { + return err + } + list, total, err := d.adviceProjectBiz.List(c.UserContext(), req) + return pkg.SuccessWithPageMsg(c, list, total, req.Page, req.PageSize, err) +} + +// TemplateCopy 重新复制行业模板到项目(覆盖项目模板字段) +func (d *ProjectService) TemplateCopy(c *fiber.Ctx) error { + req := &entitys.AdvicerProjectTemplateCopyReq{} + if err := c.BodyParser(req); err != nil { + return err + } + err := d.adviceProjectBiz.TemplateCopy(c.UserContext(), req) + return pkg.HandleResponse(c, nil, err) +} diff --git a/internal/services/advice/provider_set.go b/internal/services/advice/provider_set.go index bff0381..c800cfd 100644 --- a/internal/services/advice/provider_set.go +++ b/internal/services/advice/provider_set.go @@ -13,5 +13,6 @@ var ProviderService = wire.NewSet( NewFileService, NewIndustryService, NewProjectService, + NewModelSupService, NewTalkSkillService, )