enable to set the translate model

This commit is contained in:
RockYang 2024-11-08 18:06:39 +08:00
parent 5be4e83876
commit 135755d21d
25 changed files with 328 additions and 207 deletions

View File

@ -142,7 +142,6 @@ type SystemConfig struct {
OrderPayTimeout int `json:"order_pay_timeout,omitempty"` //订单支付超时时间 OrderPayTimeout int `json:"order_pay_timeout,omitempty"` //订单支付超时时间
VipInfoText string `json:"vip_info_text,omitempty"` // 会员页面充值说明 VipInfoText string `json:"vip_info_text,omitempty"` // 会员页面充值说明
DefaultModels []int `json:"default_models,omitempty"` // 默认开通的 AI 模型
MjPower int `json:"mj_power,omitempty"` // MJ 绘画消耗算力 MjPower int `json:"mj_power,omitempty"` // MJ 绘画消耗算力
MjActionPower int `json:"mj_action_power,omitempty"` // MJ 操作(放大,变换)消耗算力 MjActionPower int `json:"mj_action_power,omitempty"` // MJ 操作(放大,变换)消耗算力
@ -164,6 +163,7 @@ type SystemConfig struct {
Copyright string `json:"copyright"` // 版权信息 Copyright string `json:"copyright"` // 版权信息
MarkMapText string `json:"mark_map_text"` // 思维导入的默认文本 MarkMapText string `json:"mark_map_text"` // 思维导入的默认文本
EnabledVerify bool `json:"enabled_verify"` // 是否启用验证码 EnabledVerify bool `json:"enabled_verify"` // 是否启用验证码
EmailWhiteList []string `json:"email_white_list"` // 邮箱白名单列表 EmailWhiteList []string `json:"email_white_list"` // 邮箱白名单列表
TranslateModelId int `json:"translate_model_id"` // 用来做提示词翻译的大模型 id
} }

View File

@ -24,30 +24,31 @@ const (
// MjTask MidJourney 任务 // MjTask MidJourney 任务
type MjTask struct { type MjTask struct {
Id uint `json:"id"` // 任务ID Id uint `json:"id"` // 任务ID
TaskId string `json:"task_id"` // 中转任务ID TaskId string `json:"task_id"` // 中转任务ID
ClientId string `json:"client_id"` ClientId string `json:"client_id"`
ImgArr []string `json:"img_arr"` ImgArr []string `json:"img_arr"`
Type TaskType `json:"type"` Type TaskType `json:"type"`
UserId int `json:"user_id"` UserId int `json:"user_id"`
Prompt string `json:"prompt,omitempty"` Prompt string `json:"prompt,omitempty"`
NegPrompt string `json:"neg_prompt,omitempty"` NegPrompt string `json:"neg_prompt,omitempty"`
Params string `json:"full_prompt"` Params string `json:"full_prompt"`
Index int `json:"index,omitempty"` Index int `json:"index,omitempty"`
MessageId string `json:"message_id,omitempty"` MessageId string `json:"message_id,omitempty"`
MessageHash string `json:"message_hash,omitempty"` MessageHash string `json:"message_hash,omitempty"`
RetryCount int `json:"retry_count"` ChannelId string `json:"channel_id"` // 渠道ID用来区分是哪个渠道创建的任务一个任务的 create 和 action 操作必须要再同一个渠道
ChannelId string `json:"channel_id"` // 渠道ID用来区分是哪个渠道创建的任务一个任务的 create 和 action 操作必须要再同一个渠道 Mode string `json:"mode"` // 绘画模式relax, fast, turbo
Mode string `json:"mode"` // 绘画模式relax, fast, turbo TranslateModelId int `json:"translate_model_id"` // 提示词翻译模型ID
} }
type SdTask struct { type SdTask struct {
Id int `json:"id"` // job 数据库ID Id int `json:"id"` // job 数据库ID
Type TaskType `json:"type"` Type TaskType `json:"type"`
ClientId string `json:"client_id"` ClientId string `json:"client_id"`
UserId int `json:"user_id"` UserId int `json:"user_id"`
Params SdTaskParams `json:"params"` Params SdTaskParams `json:"params"`
RetryCount int `json:"retry_count"` RetryCount int `json:"retry_count"`
TranslateModelId int `json:"translate_model_id"` // 提示词翻译模型ID
} }
type SdTaskParams struct { type SdTaskParams struct {
@ -81,7 +82,8 @@ type DallTask struct {
Size string `json:"size"` Size string `json:"size"`
Style string `json:"style"` Style string `json:"style"`
Power int `json:"power"` Power int `json:"power"`
TranslateModelId int `json:"translate_model_id"` // 提示词翻译模型ID
} }
type SunoTask struct { type SunoTask struct {
@ -109,14 +111,15 @@ const (
) )
type VideoTask struct { type VideoTask struct {
ClientId string `json:"client_id"` ClientId string `json:"client_id"`
Id uint `json:"id"` Id uint `json:"id"`
Channel string `json:"channel"` Channel string `json:"channel"`
UserId int `json:"user_id"` UserId int `json:"user_id"`
Type string `json:"type"` Type string `json:"type"`
TaskId string `json:"task_id"` TaskId string `json:"task_id"`
Prompt string `json:"prompt"` // 提示词 Prompt string `json:"prompt"` // 提示词
Params VideoParams `json:"params"` Params VideoParams `json:"params"`
TranslateModelId int `json:"translate_model_id"` // 提示词翻译模型ID
} }
type VideoParams struct { type VideoParams struct {

View File

@ -371,7 +371,7 @@ func (h *ChatHandler) doRequest(ctx context.Context, req types.ApiRequest, sessi
} else { } else {
client = http.DefaultClient client = http.DefaultClient
} }
logger.Debugf("Sending %s request, API KEY:%s, PROXY: %s, Model: %s", apiKey.ApiURL, apiURL, apiKey.ProxyURL, req.Model) logger.Infof("Sending %s request, API KEY:%s, PROXY: %s, Model: %s", apiKey.ApiURL, apiURL, apiKey.ProxyURL, req.Model)
request.Header.Set("Authorization", fmt.Sprintf("Bearer %s", apiKey.Value)) request.Header.Set("Authorization", fmt.Sprintf("Bearer %s", apiKey.Value))
// 更新API KEY 最后使用时间 // 更新API KEY 最后使用时间
h.DB.Model(&model.ApiKey{}).Where("id", apiKey.Id).UpdateColumn("last_used_at", time.Now().Unix()) h.DB.Model(&model.ApiKey{}).Where("id", apiKey.Id).UpdateColumn("last_used_at", time.Now().Unix())

View File

@ -30,29 +30,25 @@ func NewChatModelHandler(app *core.AppServer, db *gorm.DB) *ChatModelHandler {
func (h *ChatModelHandler) List(c *gin.Context) { func (h *ChatModelHandler) List(c *gin.Context) {
var items []model.ChatModel var items []model.ChatModel
var chatModels = make([]vo.ChatModel, 0) var chatModels = make([]vo.ChatModel, 0)
var res *gorm.DB
session := h.DB.Session(&gorm.Session{}).Where("enabled", true) session := h.DB.Session(&gorm.Session{}).Where("enabled", true)
t := c.Query("type") t := c.Query("type")
if t != "" { if t != "" {
session = session.Where("type", t) session = session.Where("type", t)
} }
// 如果用户没有登录,则加载所有开放模型
if !h.IsLogin(c) { session = session.Where("open", true)
res = session.Where("open", true).Order("sort_num ASC").Find(&items) if h.IsLogin(c) {
} else {
user, _ := h.GetLoginUser(c) user, _ := h.GetLoginUser(c)
var models []int var models []int
err := utils.JsonDecode(user.ChatModels, &models) err := utils.JsonDecode(user.ChatModels, &models)
if err != nil {
resp.ERROR(c, "当前用户没有订阅任何模型")
return
}
// 查询用户有权限访问的模型以及所有开放的模型 // 查询用户有权限访问的模型以及所有开放的模型
res = h.DB.Where("enabled = ?", true).Where( if err == nil {
h.DB.Where("id IN ?", models).Or("open", true), session = session.Or("id IN ?", models)
).Order("sort_num ASC").Find(&items) }
} }
res := session.Order("sort_num ASC").Find(&items)
if res.Error == nil { if res.Error == nil {
for _, item := range items { for _, item := range items {
var cm vo.ChatModel var cm vo.ChatModel

View File

@ -84,14 +84,15 @@ func (h *DallJobHandler) Image(c *gin.Context) {
} }
h.dallService.PushTask(types.DallTask{ h.dallService.PushTask(types.DallTask{
ClientId: data.ClientId, ClientId: data.ClientId,
JobId: job.Id, JobId: job.Id,
UserId: uint(userId), UserId: uint(userId),
Prompt: data.Prompt, Prompt: data.Prompt,
Quality: data.Quality, Quality: data.Quality,
Size: data.Size, Size: data.Size,
Style: data.Style, Style: data.Style,
Power: job.Power, Power: job.Power,
TranslateModelId: h.App.SysConfig.TranslateModelId,
}) })
resp.SUCCESS(c) resp.SUCCESS(c)
} }

View File

@ -113,10 +113,13 @@ func (h *FunctionHandler) WeiBo(c *gin.Context) {
SetHeader("AppId", h.config.AppId). SetHeader("AppId", h.config.AppId).
SetHeader("Authorization", fmt.Sprintf("Bearer %s", h.config.Token)). SetHeader("Authorization", fmt.Sprintf("Bearer %s", h.config.Token)).
SetSuccessResult(&res).Get(url) SetSuccessResult(&res).Get(url)
if err != nil || r.IsErrorState() { if err != nil {
resp.ERROR(c, fmt.Sprintf("%v%v", err, r.Err)) resp.ERROR(c, fmt.Sprintf("%v", err))
return return
} }
if r.IsErrorState() {
resp.ERROR(c, fmt.Sprintf("error http code status: %v", r.Status))
}
if res.Code != types.Success { if res.Code != types.Success {
resp.ERROR(c, res.Message) resp.ERROR(c, res.Message)

View File

@ -87,7 +87,7 @@ func (h *MarkMapHandler) Generate(c *gin.Context) {
请直接生成结果不要任何解释性语句 请直接生成结果不要任何解释性语句
`}) `})
messages = append(messages, types.Message{Role: "user", Content: fmt.Sprintf("请生成一份有关【%s】一份思维导图要求结构清晰有条理", data.Prompt)}) messages = append(messages, types.Message{Role: "user", Content: fmt.Sprintf("请生成一份有关【%s】一份思维导图要求结构清晰有条理", data.Prompt)})
content, err := utils.SendOpenAIMessage(h.DB, messages, chatModel.Value, chatModel.KeyId) content, err := utils.SendOpenAIMessage(h.DB, messages, data.ModelId)
if err != nil { if err != nil {
resp.ERROR(c, fmt.Sprintf("请求 OpenAI API 失败: %s", err)) resp.ERROR(c, fmt.Sprintf("请求 OpenAI API 失败: %s", err))
return return

View File

@ -176,16 +176,17 @@ func (h *MidJourneyHandler) Image(c *gin.Context) {
} }
h.mjService.PushTask(types.MjTask{ h.mjService.PushTask(types.MjTask{
Id: job.Id, Id: job.Id,
ClientId: data.ClientId, ClientId: data.ClientId,
TaskId: taskId, TaskId: taskId,
Type: types.TaskType(data.TaskType), Type: types.TaskType(data.TaskType),
Prompt: data.Prompt, Prompt: data.Prompt,
NegPrompt: data.NegPrompt, NegPrompt: data.NegPrompt,
Params: params, Params: params,
UserId: userId, UserId: userId,
ImgArr: data.ImgArr, ImgArr: data.ImgArr,
Mode: h.App.SysConfig.MjMode, Mode: h.App.SysConfig.MjMode,
TranslateModelId: h.App.SysConfig.TranslateModelId,
}) })
// update user's power // update user's power
@ -226,13 +227,12 @@ func (h *MidJourneyHandler) Upscale(c *gin.Context) {
userId := utils.IntValue(utils.InterfaceToString(idValue), 0) userId := utils.IntValue(utils.InterfaceToString(idValue), 0)
taskId, _ := h.snowflake.Next(true) taskId, _ := h.snowflake.Next(true)
job := model.MidJourneyJob{ job := model.MidJourneyJob{
Type: types.TaskUpscale.String(), Type: types.TaskUpscale.String(),
ReferenceId: data.MessageId, UserId: userId,
UserId: userId, TaskId: taskId,
TaskId: taskId, Progress: 0,
Progress: 0, Power: h.App.SysConfig.MjActionPower,
Power: h.App.SysConfig.MjActionPower, CreatedAt: time.Now(),
CreatedAt: time.Now(),
} }
if res := h.DB.Create(&job); res.Error != nil || res.RowsAffected == 0 { if res := h.DB.Create(&job); res.Error != nil || res.RowsAffected == 0 {
resp.ERROR(c, "添加任务失败:"+res.Error.Error()) resp.ERROR(c, "添加任务失败:"+res.Error.Error())
@ -281,14 +281,13 @@ func (h *MidJourneyHandler) Variation(c *gin.Context) {
userId := utils.IntValue(utils.InterfaceToString(idValue), 0) userId := utils.IntValue(utils.InterfaceToString(idValue), 0)
taskId, _ := h.snowflake.Next(true) taskId, _ := h.snowflake.Next(true)
job := model.MidJourneyJob{ job := model.MidJourneyJob{
Type: types.TaskVariation.String(), Type: types.TaskVariation.String(),
ChannelId: data.ChannelId, ChannelId: data.ChannelId,
ReferenceId: data.MessageId, UserId: userId,
UserId: userId, TaskId: taskId,
TaskId: taskId, Progress: 0,
Progress: 0, Power: h.App.SysConfig.MjActionPower,
Power: h.App.SysConfig.MjActionPower, CreatedAt: time.Now(),
CreatedAt: time.Now(),
} }
if res := h.DB.Create(&job); res.Error != nil || res.RowsAffected == 0 { if res := h.DB.Create(&job); res.Error != nil || res.RowsAffected == 0 {
resp.ERROR(c, "添加任务失败:"+res.Error.Error()) resp.ERROR(c, "添加任务失败:"+res.Error.Error())

View File

@ -0,0 +1,59 @@
package handler
// * +++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
// * Copyright 2023 The Geek-AI Authors. All rights reserved.
// * Use of this source code is governed by a Apache-2.0 license
// * that can be found in the LICENSE file.
// * @Author yangjian102621@163.com
// * +++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
import (
"fmt"
"geekai/core"
"geekai/core/types"
"geekai/service"
"geekai/service/oss"
"geekai/service/suno"
"geekai/utils"
"geekai/utils/resp"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
)
// 提示词生成 handler
// 使用 AI 生成绘画指令,歌词,视频生成指令等
type PromptHandler struct {
BaseHandler
sunoService *suno.Service
uploader *oss.UploaderManager
userService *service.UserService
}
func NewPromptHandler(app *core.AppServer, db *gorm.DB, userService *service.UserService) *PromptHandler {
return &PromptHandler{
BaseHandler: BaseHandler{
App: app,
DB: db,
},
userService: userService,
}
}
// Lyric 生成歌词
func (h *PromptHandler) Lyric(c *gin.Context) {
var data struct {
Prompt string `json:"prompt"`
}
if err := c.ShouldBindJSON(&data); err != nil {
resp.ERROR(c, types.InvalidArgs)
return
}
content, err := utils.OpenAIRequest(h.DB, fmt.Sprintf(service.LyricPromptTemplate, data.Prompt), h.App.SysConfig.TranslateModelId)
if err != nil {
resp.ERROR(c, err.Error())
return
}
resp.SUCCESS(c, content)
}

View File

@ -109,29 +109,37 @@ func (h *SdJobHandler) Image(c *gin.Context) {
resp.ERROR(c, "error with generate task id: "+err.Error()) resp.ERROR(c, "error with generate task id: "+err.Error())
return return
} }
params := types.SdTaskParams{
TaskId: taskId, task := types.SdTask{
Prompt: data.Prompt, ClientId: data.ClientId,
NegPrompt: data.NegPrompt, Type: types.TaskImage,
Steps: data.Steps, Params: types.SdTaskParams{
Sampler: data.Sampler, TaskId: taskId,
FaceFix: data.FaceFix, Prompt: data.Prompt,
CfgScale: data.CfgScale, NegPrompt: data.NegPrompt,
Seed: data.Seed, Steps: data.Steps,
Height: data.Height, Sampler: data.Sampler,
Width: data.Width, FaceFix: data.FaceFix,
HdFix: data.HdFix, CfgScale: data.CfgScale,
HdRedrawRate: data.HdRedrawRate, Seed: data.Seed,
HdScale: data.HdScale, Height: data.Height,
HdScaleAlg: data.HdScaleAlg, Width: data.Width,
HdSteps: data.HdSteps, HdFix: data.HdFix,
HdRedrawRate: data.HdRedrawRate,
HdScale: data.HdScale,
HdScaleAlg: data.HdScaleAlg,
HdSteps: data.HdSteps,
},
UserId: userId,
TranslateModelId: h.App.SysConfig.TranslateModelId,
} }
job := model.SdJob{ job := model.SdJob{
UserId: userId, UserId: userId,
Type: types.TaskImage.String(), Type: types.TaskImage.String(),
TaskId: params.TaskId, TaskId: taskId,
Params: utils.JsonEncode(params), Params: utils.JsonEncode(task.Params),
TaskInfo: utils.JsonEncode(task),
Prompt: data.Prompt, Prompt: data.Prompt,
Progress: 0, Progress: 0,
Power: h.App.SysConfig.SdPower, Power: h.App.SysConfig.SdPower,
@ -143,13 +151,8 @@ func (h *SdJobHandler) Image(c *gin.Context) {
return return
} }
h.sdService.PushTask(types.SdTask{ task.Id = int(job.Id)
Id: int(job.Id), h.sdService.PushTask(task)
ClientId: data.ClientId,
Type: types.TaskImage,
Params: params,
UserId: userId,
})
// update user's power // update user's power
err = h.userService.DecreasePower(job.UserId, job.Power, model.PowerLog{ err = h.userService.DecreasePower(job.UserId, job.Power, model.PowerLog{

View File

@ -334,40 +334,3 @@ func (h *SunoHandler) Play(c *gin.Context) {
} }
h.DB.Model(&model.SunoJob{}).Where("song_id", songId).UpdateColumn("play_times", gorm.Expr("play_times + ?", 1)) h.DB.Model(&model.SunoJob{}).Where("song_id", songId).UpdateColumn("play_times", gorm.Expr("play_times + ?", 1))
} }
const genLyricTemplate = `
你是一位才华横溢的作曲家拥有丰富的情感和细腻的笔触你对文字有着独特的感悟力能将各种情感和意境巧妙地融入歌词中
请以%s为主题创作一首歌曲歌曲时间不要太短3分钟左右不要输出任何解释性的内容
输出格式如下
歌曲名称
第一节
{{歌词内容}}
副歌
{{歌词内容}}
第二节
{{歌词内容}}
副歌
{{歌词内容}}
尾声
{{歌词内容}}
`
// Lyric 生成歌词
func (h *SunoHandler) Lyric(c *gin.Context) {
var data struct {
Prompt string `json:"prompt"`
}
if err := c.ShouldBindJSON(&data); err != nil {
resp.ERROR(c, types.InvalidArgs)
return
}
content, err := utils.OpenAIRequest(h.DB, fmt.Sprintf(genLyricTemplate, data.Prompt), "gpt-4o-mini", 0)
if err != nil {
resp.ERROR(c, err.Error())
return
}
resp.SUCCESS(c, content)
}

View File

@ -132,14 +132,13 @@ func (h *UserHandler) Register(c *gin.Context) {
salt := utils.RandString(8) salt := utils.RandString(8)
user := model.User{ user := model.User{
Username: data.Username, Username: data.Username,
Password: utils.GenPassword(data.Password, salt), Password: utils.GenPassword(data.Password, salt),
Avatar: "/images/avatar/user.png", Avatar: "/images/avatar/user.png",
Salt: salt, Salt: salt,
Status: true, Status: true,
ChatRoles: utils.JsonEncode([]string{"gpt"}), // 默认只订阅通用助手角色 ChatRoles: utils.JsonEncode([]string{"gpt"}), // 默认只订阅通用助手角色
ChatModels: utils.JsonEncode(h.App.SysConfig.DefaultModels), // 默认开通的模型 Power: h.App.SysConfig.InitPower,
Power: h.App.SysConfig.InitPower,
} }
// check if the username is existing // check if the username is existing
@ -417,16 +416,15 @@ func (h *UserHandler) CLoginCallback(c *gin.Context) {
salt := utils.RandString(8) salt := utils.RandString(8)
password := fmt.Sprintf("%d", utils.RandomNumber(8)) password := fmt.Sprintf("%d", utils.RandomNumber(8))
user = model.User{ user = model.User{
Username: fmt.Sprintf("%s@%d", loginType, utils.RandomNumber(10)), Username: fmt.Sprintf("%s@%d", loginType, utils.RandomNumber(10)),
Password: utils.GenPassword(password, salt), Password: utils.GenPassword(password, salt),
Avatar: fmt.Sprintf("%s", data["avatar"]), Avatar: fmt.Sprintf("%s", data["avatar"]),
Salt: salt, Salt: salt,
Status: true, Status: true,
ChatRoles: utils.JsonEncode([]string{"gpt"}), // 默认只订阅通用助手角色 ChatRoles: utils.JsonEncode([]string{"gpt"}), // 默认只订阅通用助手角色
ChatModels: utils.JsonEncode(h.App.SysConfig.DefaultModels), // 默认开通的模型 Power: h.App.SysConfig.InitPower,
Power: h.App.SysConfig.InitPower, OpenId: fmt.Sprintf("%s", data["openid"]),
OpenId: fmt.Sprintf("%s", data["openid"]), Nickname: fmt.Sprintf("%s", data["nickname"]),
Nickname: fmt.Sprintf("%s", data["nickname"]),
} }
tx = h.DB.Create(&user) tx = h.DB.Create(&user)

View File

@ -96,12 +96,13 @@ func (h *VideoHandler) LumaCreate(c *gin.Context) {
// 创建任务 // 创建任务
h.videoService.PushTask(types.VideoTask{ h.videoService.PushTask(types.VideoTask{
ClientId: data.ClientId, ClientId: data.ClientId,
Id: job.Id, Id: job.Id,
UserId: userId, UserId: userId,
Type: types.VideoLuma, Type: types.VideoLuma,
Prompt: data.Prompt, Prompt: data.Prompt,
Params: params, Params: params,
TranslateModelId: h.App.SysConfig.TranslateModelId,
}) })
// update user's power // update user's power

View File

@ -484,7 +484,6 @@ func main() {
group.POST("update", h.Update) group.POST("update", h.Update)
group.GET("detail", h.Detail) group.GET("detail", h.Detail)
group.GET("play", h.Play) group.GET("play", h.Play)
group.POST("lyric", h.Lyric)
}), }),
fx.Provide(handler.NewVideoHandler), fx.Provide(handler.NewVideoHandler),
fx.Invoke(func(s *core.AppServer, h *handler.VideoHandler) { fx.Invoke(func(s *core.AppServer, h *handler.VideoHandler) {
@ -518,6 +517,11 @@ func main() {
fx.Invoke(func(s *core.AppServer, h *handler.WebsocketHandler) { fx.Invoke(func(s *core.AppServer, h *handler.WebsocketHandler) {
s.Engine.Any("/api/ws", h.Client) s.Engine.Any("/api/ws", h.Client)
}), }),
fx.Provide(handler.NewPromptHandler),
fx.Invoke(func(s *core.AppServer, h *handler.PromptHandler) {
group := s.Engine.Group("/api/prompt")
group.POST("/lyric", h.Lyric)
}),
fx.Invoke(func(s *core.AppServer, db *gorm.DB) { fx.Invoke(func(s *core.AppServer, db *gorm.DB) {
go func() { go func() {
err := s.Run(db) err := s.Run(db)

View File

@ -114,7 +114,7 @@ func (s *Service) Image(task types.DallTask, sync bool) (string, error) {
prompt := task.Prompt prompt := task.Prompt
// translate prompt // translate prompt
if utils.HasChinese(prompt) { if utils.HasChinese(prompt) {
content, err := utils.OpenAIRequest(s.db, fmt.Sprintf(service.RewritePromptTemplate, prompt), "gpt-4o-mini", 0) content, err := utils.OpenAIRequest(s.db, fmt.Sprintf(service.TranslatePromptTemplate, prompt), task.TranslateModelId)
if err == nil { if err == nil {
prompt = content prompt = content
logger.Debugf("重写后提示词:%s", prompt) logger.Debugf("重写后提示词:%s", prompt)

View File

@ -58,7 +58,7 @@ func (s *Service) Run() {
// translate prompt // translate prompt
if utils.HasChinese(task.Prompt) { if utils.HasChinese(task.Prompt) {
content, err := utils.OpenAIRequest(s.db, fmt.Sprintf(service.TranslatePromptTemplate, task.Prompt), "gpt-4o-mini", 0) content, err := utils.OpenAIRequest(s.db, fmt.Sprintf(service.TranslatePromptTemplate, task.Prompt), task.TranslateModelId)
if err == nil { if err == nil {
task.Prompt = content task.Prompt = content
} else { } else {
@ -67,7 +67,7 @@ func (s *Service) Run() {
} }
// translate negative prompt // translate negative prompt
if task.NegPrompt != "" && utils.HasChinese(task.NegPrompt) { if task.NegPrompt != "" && utils.HasChinese(task.NegPrompt) {
content, err := utils.OpenAIRequest(s.db, fmt.Sprintf(service.TranslatePromptTemplate, task.NegPrompt), "gpt-4o-mini", 0) content, err := utils.OpenAIRequest(s.db, fmt.Sprintf(service.TranslatePromptTemplate, task.NegPrompt), task.TranslateModelId)
if err == nil { if err == nil {
task.NegPrompt = content task.NegPrompt = content
} else { } else {
@ -275,7 +275,6 @@ func (s *Service) SyncTaskProgress() {
} }
oldProgress := job.Progress oldProgress := job.Progress
job.Progress = utils.IntValue(strings.Replace(task.Progress, "%", "", 1), 0) job.Progress = utils.IntValue(strings.Replace(task.Progress, "%", "", 1), 0)
job.Prompt = task.PromptEn
if task.ImageUrl != "" { if task.ImageUrl != "" {
job.OrgURL = task.ImageUrl job.OrgURL = task.ImageUrl
} }

View File

@ -48,6 +48,19 @@ func NewService(db *gorm.DB, manager *oss.UploaderManager, levelDB *store.LevelD
} }
func (s *Service) Run() { func (s *Service) Run() {
// 将数据库中未提交的人物加载到队列
var jobs []model.SdJob
s.db.Where("progress", 0).Find(&jobs)
for _, v := range jobs {
var task types.SdTask
err := utils.JsonDecode(v.TaskInfo, &task)
if err != nil {
logger.Errorf("decode task info with error: %v", err)
continue
}
task.Id = int(v.Id)
s.PushTask(task)
}
logger.Infof("Starting Stable-Diffusion job consumer") logger.Infof("Starting Stable-Diffusion job consumer")
go func() { go func() {
for { for {
@ -60,7 +73,7 @@ func (s *Service) Run() {
// translate prompt // translate prompt
if utils.HasChinese(task.Params.Prompt) { if utils.HasChinese(task.Params.Prompt) {
content, err := utils.OpenAIRequest(s.db, fmt.Sprintf(service.RewritePromptTemplate, task.Params.Prompt), "gpt-4o-mini", 0) content, err := utils.OpenAIRequest(s.db, fmt.Sprintf(service.TranslatePromptTemplate, task.Params.Prompt), task.TranslateModelId)
if err == nil { if err == nil {
task.Params.Prompt = content task.Params.Prompt = content
} else { } else {
@ -70,7 +83,7 @@ func (s *Service) Run() {
// translate negative prompt // translate negative prompt
if task.Params.NegPrompt != "" && utils.HasChinese(task.Params.NegPrompt) { if task.Params.NegPrompt != "" && utils.HasChinese(task.Params.NegPrompt) {
content, err := utils.OpenAIRequest(s.db, fmt.Sprintf(service.TranslatePromptTemplate, task.Params.NegPrompt), "gpt-4o-mini", 0) content, err := utils.OpenAIRequest(s.db, fmt.Sprintf(service.TranslatePromptTemplate, task.Params.NegPrompt), task.TranslateModelId)
if err == nil { if err == nil {
task.Params.NegPrompt = content task.Params.NegPrompt = content
} else { } else {
@ -161,7 +174,7 @@ func (s *Service) Txt2Img(task types.SdTask) error {
} }
apiURL := fmt.Sprintf("%s/sdapi/v1/txt2img", apiKey.ApiURL) apiURL := fmt.Sprintf("%s/sdapi/v1/txt2img", apiKey.ApiURL)
logger.Debugf("send image request to %s", apiURL) logger.Infof("send image request to %s", apiURL)
// send a request to sd api endpoint // send a request to sd api endpoint
go func() { go func() {
response, err := s.httpClient.R(). response, err := s.httpClient.R().

View File

@ -14,5 +14,79 @@ type NotifyMessage struct {
Message string `json:"message"` Message string `json:"message"`
} }
const RewritePromptTemplate = "Please rewrite the following text into AI painting prompt words, and please try to add detailed description of the picture, painting style, scene, rendering effect, picture light and other creative elements. Just output the final prompt word directly. Do not output any explanation lines. The text to be rewritten is: [%s]"
const TranslatePromptTemplate = "Translate the following painting prompt words into English keyword phrases. Without any explanation, directly output the keyword phrases separated by commas. The content to be translated is: [%s]" const TranslatePromptTemplate = "Translate the following painting prompt words into English keyword phrases. Without any explanation, directly output the keyword phrases separated by commas. The content to be translated is: [%s]"
const ImagePromptOptimizeTemplate = `
Create a highly effective prompt to provide to an AI image generation tool in order to create an artwork based on a desired concept.
Please specify details about the artwork, such as the style, subject, mood, and other important characteristics you want the resulting image to have.
Remember, prompts should always be output in English.
# Steps
1. **Subject Description**: Describe the main subject of the image clearly. Include as much detail as possible about what should be in the scene. For example, "a majestic lion roaring at sunrise" or "a futuristic city with flying cars."
2. **Art Style**: Specify the art style you envision. Possible options include 'realistic', 'impressionist', a specific artist name, or imaginative styles like "cyberpunk." This helps the AI achieve your visual expectations.
3. **Mood or Atmosphere**: Convey the feeling you want the image to evoke. For instance, peaceful, chaotic, epic, etc.
4. **Color Palette and Lighting**: Mention color preferences or lighting. For example, "vibrant with shades of blue and purple" or "dim and dramatic lighting."
5. **Optional Features**: You can add any additional attributes, such as background details, attention to textures, or any specific kind of framing.
# Output Format
- **Prompt Format**: A descriptive phrase that includes key aspects of the artwork (subject, style, mood, colors, lighting, any optional features).
Here is an example of how the final prompt should look:
"An ethereal landscape featuring towering ice mountains, in an impressionist style reminiscent of Claude Monet, with a serene mood. The sky is glistening with soft purples and whites, with a gentle morning sun illuminating the scene."
**Please input the prompt words directly in English, and do not input any other explanatory statements**
# Examples
1. **Input**:
- Subject: A white tiger in a dense jungle
- Art Style: Realistic
- Mood: Intense, mysterious
- Lighting: Dramatic contrast with light filtering through leaves
**Output Prompt**: "A realistic rendering of a white tiger stealthily moving through a dense jungle, with an intense, mysterious mood. The lighting creates strong contrasts as beams of sunlight filter through a thick canopy of leaves."
2. **Input**:
- Subject: An enchanted castle on a floating island
- Art Style: Fantasy
- Mood: Majestic, magical
- Colors: Bright blues, greens, and gold
**Output Prompt**: "A majestic fantasy castle on a floating island above the clouds, with bright blues, greens, and golds to create a magical, dreamy atmosphere. Textured cobblestone details and glistening waters surround the scene."
# Notes
- Ensure that you mix different aspects to get a comprehensive and visually compelling prompt.
- Be as descriptive as possible as it often helps generate richer, more detailed images.
- If you want the image to resemble a particular artist's work, be sure to mention the artist explicitly. e.g., "in the style of Van Gogh."
The theme of the creation is:%s
`
const LyricPromptTemplate = `
你是一位才华横溢的作曲家拥有丰富的情感和细腻的笔触你对文字有着独特的感悟力能将各种情感和意境巧妙地融入歌词中
请以%s为主题创作一首歌曲歌曲时间不要太短3分钟左右不要输出任何解释性的内容
输出格式如下
歌曲名称
第一节
{{歌词内容}}
副歌
{{歌词内容}}
第二节
{{歌词内容}}
副歌
{{歌词内容}}
尾声
{{歌词内容}}
`

View File

@ -87,7 +87,7 @@ func (s *Service) Run() {
// translate prompt // translate prompt
if utils.HasChinese(task.Prompt) { if utils.HasChinese(task.Prompt) {
content, err := utils.OpenAIRequest(s.db, fmt.Sprintf(service.TranslatePromptTemplate, task.Prompt), "gpt-4o-mini", 0) content, err := utils.OpenAIRequest(s.db, fmt.Sprintf(service.TranslatePromptTemplate, task.Prompt), task.TranslateModelId)
if err == nil { if err == nil {
task.Prompt = content task.Prompt = content
} else { } else {

View File

@ -7,6 +7,7 @@ type SdJob struct {
Type string Type string
UserId int UserId int
TaskId string TaskId string
TaskInfo string // 原始任务信息
ImgURL string ImgURL string
Progress int Progress int
Prompt string Prompt string

View File

@ -1,21 +1,20 @@
package vo package vo
type MidJourneyJob struct { type MidJourneyJob struct {
Id uint `json:"id"` Id uint `json:"id"`
Type string `json:"type"` Type string `json:"type"`
UserId int `json:"user_id"` UserId int `json:"user_id"`
ChannelId string `json:"channel_id"` ChannelId string `json:"channel_id"`
TaskId string `json:"task_id"` TaskId string `json:"task_id"`
MessageId string `json:"message_id"` MessageId string `json:"message_id"`
ReferenceId string `json:"reference_id"` ImgURL string `json:"img_url"`
ImgURL string `json:"img_url"` OrgURL string `json:"org_url"`
OrgURL string `json:"org_url"` Hash string `json:"hash"`
Hash string `json:"hash"` Progress int `json:"progress"`
Progress int `json:"progress"` Prompt string `json:"prompt"`
Prompt string `json:"prompt"` UseProxy bool `json:"use_proxy"`
UseProxy bool `json:"use_proxy"` Publish bool `json:"publish"`
Publish bool `json:"publish"` ErrMsg string `json:"err_msg"`
ErrMsg string `json:"err_msg"` Power int `json:"power"`
Power int `json:"power"` CreatedAt int64 `json:"created_at"`
CreatedAt int64 `json:"created_at"`
} }

View File

@ -45,20 +45,25 @@ type apiRes struct {
} `json:"choices"` } `json:"choices"`
} }
func OpenAIRequest(db *gorm.DB, prompt string, modelName string, keyId int) (string, error) { func OpenAIRequest(db *gorm.DB, prompt string, modelId int) (string, error) {
messages := make([]interface{}, 1) messages := make([]interface{}, 1)
messages[0] = types.Message{ messages[0] = types.Message{
Role: "user", Role: "user",
Content: prompt, Content: prompt,
} }
return SendOpenAIMessage(db, messages, modelName, keyId) return SendOpenAIMessage(db, messages, modelId)
} }
func SendOpenAIMessage(db *gorm.DB, messages []interface{}, modelName string, keyId int) (string, error) { func SendOpenAIMessage(db *gorm.DB, messages []interface{}, modelId int) (string, error) {
var chatModel model.ChatModel
db.Where("id", modelId).First(&chatModel)
if chatModel.Name == "" {
chatModel.Name = "gpt-4o-mini" // 默认使用 gpt-4o-mini
}
var apiKey model.ApiKey var apiKey model.ApiKey
session := db.Session(&gorm.Session{}).Where("type", "chat").Where("enabled", true) session := db.Session(&gorm.Session{}).Where("type", "chat").Where("enabled", true)
if keyId > 0 { if chatModel.KeyId > 0 {
session = session.Where("id", keyId) session = session.Where("id", chatModel.KeyId)
} }
err := session.First(&apiKey).Error err := session.First(&apiKey).Error
if err != nil { if err != nil {
@ -71,11 +76,11 @@ func SendOpenAIMessage(db *gorm.DB, messages []interface{}, modelName string, ke
client.SetProxyURL(apiKey.ApiURL) client.SetProxyURL(apiKey.ApiURL)
} }
apiURL := fmt.Sprintf("%s/v1/chat/completions", apiKey.ApiURL) apiURL := fmt.Sprintf("%s/v1/chat/completions", apiKey.ApiURL)
logger.Debugf("Sending %s request, API KEY:%s, PROXY: %s, Model: %s", apiKey.ApiURL, apiURL, apiKey.ProxyURL, modelName) logger.Infof("Sending %s request, API KEY:%s, PROXY: %s, Model: %s", apiKey.ApiURL, apiURL, apiKey.ProxyURL, chatModel.Name)
r, err := client.R().SetHeader("Body-Type", "application/json"). r, err := client.R().SetHeader("Body-Type", "application/json").
SetHeader("Authorization", "Bearer "+apiKey.Value). SetHeader("Authorization", "Bearer "+apiKey.Value).
SetBody(types.ApiRequest{ SetBody(types.ApiRequest{
Model: modelName, Model: chatModel.Name,
Temperature: 0.9, Temperature: 0.9,
MaxTokens: 1024, MaxTokens: 1024,
Stream: false, Stream: false,

View File

@ -0,0 +1 @@
ALTER TABLE `chatgpt_sd_jobs` ADD `task_info` TEXT NOT NULL COMMENT '任务详情' AFTER `task_id`;

View File

@ -615,7 +615,7 @@ const createLyric = () => {
return showMessageError("请输入歌词描述") return showMessageError("请输入歌词描述")
} }
isGenerating.value = true isGenerating.value = true
httpPost("/api/suno/lyric", {prompt: data.value.lyrics}).then(res => { httpPost("/api/prompt/lyric", {prompt: data.value.lyrics}).then(res => {
const lines = res.data.split('\n'); const lines = res.data.split('\n');
data.value.title = lines.shift().replace(/\*/g,"") data.value.title = lines.shift().replace(/\*/g,"")
lines.shift() lines.shift()

View File

@ -153,14 +153,13 @@
</template> </template>
</el-input> </el-input>
</el-form-item> </el-form-item>
<el-form-item label="默认AI模型" prop="default_models"> <el-form-item label="默认翻译模型">
<template #default> <template #default>
<div class="tip-input"> <div class="tip-input">
<el-select <el-select
v-model="system['default_models']" v-model.number="system['translate_model_id']"
multiple
:filterable="true" :filterable="true"
placeholder="选择AI模型多选" placeholder="选择一个默认模型来翻译提示词"
style="width: 100%" style="width: 100%"
> >
<el-option <el-option