This commit is contained in:
renzhiyuan 2026-09-22 01:39:43 +08:00
parent 092da9fac4
commit ecbbbd98f1
26 changed files with 418 additions and 545 deletions

BIN
1.jpg

Binary file not shown.

Before

Width:  |  Height:  |  Size: 309 KiB

View File

@ -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

140
README.md
View File

@ -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

View File

@ -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,
))

1
go.mod
View File

@ -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

8
go.sum
View File

@ -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=

View File

@ -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
}

View File

@ -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
}

View File

@ -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
}

View File

@ -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 {

View File

@ -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
}

View File

@ -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
}

View File

@ -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
}

View File

@ -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(

View File

@ -7,12 +7,12 @@ import (
)
var ProviderSetBiz = wire.NewSet(
third_party.NewOpenAi,
NewAdviceFileBiz,
third_party.NewHsyq,
NewAdviceAdvicerBiz,
NewAdviceSkillBiz,
NewAdviceProjectBiz,
NewAdviceModelSupBiz,
NewAdviceClientBiz,
NewAdviceChatBiz,
NewAdvicerIndustryBiz,

View File

@ -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

View File

@ -12,10 +12,18 @@ const TableNameAiAdviceProject = "ai_advice_project"
// AiAdviceProject mapped from table <ai_advice_project>
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

View File

@ -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 (

View File

@ -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 {

View File

@ -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(),

View File

@ -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)
}

View File

@ -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 {

View File

@ -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)

View File

@ -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)
}

View File

@ -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)
}

View File

@ -13,5 +13,6 @@ var ProviderService = wire.NewSet(
NewFileService,
NewIndustryService,
NewProjectService,
NewModelSupService,
NewTalkSkillService,
)