This commit is contained in:
parent
092da9fac4
commit
ecbbbd98f1
|
|
@ -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
140
README.md
|
|
@ -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
|
||||
|
|
@ -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
1
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
|
||||
|
|
|
|||
8
go.sum
8
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=
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -7,12 +7,12 @@ import (
|
|||
)
|
||||
|
||||
var ProviderSetBiz = wire.NewSet(
|
||||
|
||||
third_party.NewOpenAi,
|
||||
NewAdviceFileBiz,
|
||||
third_party.NewHsyq,
|
||||
NewAdviceAdvicerBiz,
|
||||
NewAdviceSkillBiz,
|
||||
NewAdviceProjectBiz,
|
||||
NewAdviceModelSupBiz,
|
||||
NewAdviceClientBiz,
|
||||
NewAdviceChatBiz,
|
||||
NewAdvicerIndustryBiz,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -13,5 +13,6 @@ var ProviderService = wire.NewSet(
|
|||
NewFileService,
|
||||
NewIndustryService,
|
||||
NewProjectService,
|
||||
NewModelSupService,
|
||||
NewTalkSkillService,
|
||||
)
|
||||
|
|
|
|||
Loading…
Reference in New Issue