diff --git a/short_drama.db b/short_drama.db index 7513994..ff1f5a6 100644 Binary files a/short_drama.db and b/short_drama.db differ diff --git a/shortdrama/agent/chat_model.go b/shortdrama/agent/chat_model.go index 72b2d65..110ac84 100644 --- a/shortdrama/agent/chat_model.go +++ b/shortdrama/agent/chat_model.go @@ -49,7 +49,7 @@ func CallChatModel(ctx context.Context, cfg *ModelConfig, req *ChatRequest) (*Ch return nil, err } - url := trimSlashes(cfg.BaseURL) + "/v1/chat/completions" + url := trimSlashes(cfg.BaseURL) var lastErr error maxRetries := cfg.MaxRetries diff --git a/shortdrama/controller/config_controller.go b/shortdrama/controller/config_controller.go index 0d34d63..2bff190 100644 --- a/shortdrama/controller/config_controller.go +++ b/shortdrama/controller/config_controller.go @@ -89,13 +89,7 @@ func (c *config) SaveUserConfig(ctx context.Context, req *dto.SaveUserModelConfi if userId <= 0 { return nil, fmt.Errorf("未登录") } - cfg := &entity.UserModelConfig{ - ModelConfigId: req.ModelConfigId, - ApiKey: req.ApiKey, - Temperature: req.Temperature, - MaxTokens: req.MaxTokens, - } - return nil, service.ConfigService.SaveUserConfig(ctx, userId, cfg) + return nil, service.ConfigService.SaveUserConfigs(ctx, userId, req.Configs) } func (c *config) GetUserModelList(ctx context.Context, req *dto.GetUserModelListReq) (res *dto.GetUserModelListRes, err error) { diff --git a/shortdrama/dao/model_config_dao.go b/shortdrama/dao/model_config_dao.go index fea8a69..31a307d 100644 --- a/shortdrama/dao/model_config_dao.go +++ b/shortdrama/dao/model_config_dao.go @@ -21,11 +21,9 @@ func init() { "id INTEGER PRIMARY KEY AUTOINCREMENT,"+ "model_type TEXT NOT NULL DEFAULT '',"+ "model_name TEXT NOT NULL DEFAULT '',"+ - "base_url TEXT NOT NULL DEFAULT '',"+ "schema TEXT NOT NULL DEFAULT '',"+ "price INTEGER NOT NULL DEFAULT 0,"+ "price_unit TEXT NOT NULL DEFAULT 'second',"+ - "task_callback_url TEXT NOT NULL DEFAULT '',"+ "created_at DATETIME DEFAULT (datetime('now','localtime')),"+ "updated_at DATETIME DEFAULT (datetime('now','localtime'))"+ ")") @@ -33,12 +31,12 @@ func init() { g.Log().Warningf(ctx, "创建模型配置表失败: %v", err) } - // 检测旧表结构:含 api_key / chat_api_key / is_active 等旧列 → 迁移 + // 检测旧表结构:含 base_url / api_key / chat_api_key / is_active / task_callback_url 等旧列 → 迁移 hasOldColumns := false if r, _ := g.DB().GetAll(ctx, "PRAGMA table_info("+public.TableNameModelConfig+")"); r != nil { for _, col := range r { name := col["name"].String() - if name == "api_key" || name == "chat_api_key" || name == "is_active" || name == "price_per_second" || name == "max_tokens" { + if name == "base_url" || name == "api_key" || name == "chat_api_key" || name == "is_active" || name == "price_per_second" || name == "max_tokens" || name == "task_callback_url" { hasOldColumns = true break } @@ -53,11 +51,9 @@ func init() { "id INTEGER PRIMARY KEY AUTOINCREMENT,"+ "model_type TEXT NOT NULL DEFAULT '',"+ "model_name TEXT NOT NULL DEFAULT '',"+ - "base_url TEXT NOT NULL DEFAULT '',"+ "schema TEXT NOT NULL DEFAULT '',"+ "price INTEGER NOT NULL DEFAULT 0,"+ "price_unit TEXT NOT NULL DEFAULT 'second',"+ - "task_callback_url TEXT NOT NULL DEFAULT '',"+ "created_at DATETIME DEFAULT (datetime('now','localtime')),"+ "updated_at DATETIME DEFAULT (datetime('now','localtime'))"+ ")") @@ -73,16 +69,14 @@ func init() { if oldRow != nil && !oldRow.IsEmpty() { _, _ = g.DB().Model(newTable).Ctx(ctx).Data(g.Map{ "model_type": "chat", "model_name": oldRow["chat_model_name"], - "base_url": oldRow["chat_base_url"], - "schema": oldRow["chat_schema"], - "price": 0, "price_unit": "video", + "schema": oldRow["chat_schema"], + "price": 0, "price_unit": "video", "created_at": now, "updated_at": now, }).Insert() _, _ = g.DB().Model(newTable).Ctx(ctx).Data(g.Map{ "model_type": "video", "model_name": oldRow["video_model_name"], - "base_url": oldRow["video_base_url"], - "schema": oldRow["video_schema"], "task_callback_url": oldRow["video_task_callback_url"], - "price": gconv.Int(oldRow["price_per_second"]), "price_unit": "second", + "schema": oldRow["video_schema"], + "price": gconv.Int(oldRow["price_per_second"]), "price_unit": "second", "created_at": now, "updated_at": now, }).Insert() } @@ -100,11 +94,9 @@ func init() { } _, _ = g.DB().Model(newTable).Ctx(ctx).Data(g.Map{ "model_type": mt, "model_name": or["model_name"], - "base_url": or["base_url"], - "schema": or["schema"], - "price": priceVal, "price_unit": unit, - "task_callback_url": or["task_callback_url"], - "created_at": now, "updated_at": now, + "schema": or["schema"], + "price": priceVal, "price_unit": unit, + "created_at": now, "updated_at": now, }).Insert() } } @@ -118,8 +110,8 @@ func init() { if count == 0 { now := gtime.Now().Format("Y-m-d H:i:s") defaults := []g.Map{ - {"model_type": "chat", "model_name": "", "base_url": "", "schema": "{\"properties\":{\"parameters\":{\"properties\":{\"max_tokens\":{\"type\":\"integer\",\"description\":\"最大Token数\",\"default\":4096,\"maximum\":65535}}}}}}", "price": 0, "price_unit": "video", "created_at": now, "updated_at": now}, - {"model_type": "video", "model_name": "", "base_url": "", "schema": "", "price": 10, "price_unit": "second", "created_at": now, "updated_at": now}, + {"model_type": "chat", "model_name": "", "schema": "{\"properties\":{\"parameters\":{\"properties\":{\"max_tokens\":{\"type\":\"integer\",\"description\":\"最大Token数\",\"default\":4096,\"maximum\":65535}}}}}}", "price": 0, "price_unit": "video", "created_at": now, "updated_at": now}, + {"model_type": "video", "model_name": "", "schema": "", "price": 10, "price_unit": "second", "created_at": now, "updated_at": now}, } for _, m := range defaults { if _, e := g.DB().Model(public.TableNameModelConfig).Ctx(ctx).Data(m).Insert(); e != nil { diff --git a/shortdrama/dao/user_model_config_dao.go b/shortdrama/dao/user_model_config_dao.go index 6dcf09a..ae84269 100644 --- a/shortdrama/dao/user_model_config_dao.go +++ b/shortdrama/dao/user_model_config_dao.go @@ -28,7 +28,10 @@ func init() { id INTEGER PRIMARY KEY AUTOINCREMENT, user_id INTEGER NOT NULL DEFAULT 0, model_config_id INTEGER NOT NULL DEFAULT 0, + model_type TEXT NOT NULL DEFAULT '', api_key TEXT NOT NULL DEFAULT '', + base_url TEXT NOT NULL DEFAULT '', + task_callback_url TEXT NOT NULL DEFAULT '', temperature REAL NOT NULL DEFAULT 0, max_tokens INTEGER NOT NULL DEFAULT 0, created_at DATETIME DEFAULT (datetime('now','localtime')), @@ -40,46 +43,76 @@ func init() { return } - // 表已存在,检测是否需要迁移(移除 is_active 列或修复列类型) + // 表已存在,检测是否需要迁移(新增 base_url/task_callback_url 列或修复列类型) r, _ := g.DB().GetAll(ctx, "PRAGMA table_info("+public.TableNameUserModelConfig+")") var hasIsActive bool - var needsMigrate bool + var missingBaseUrl bool + var apiKeyNullable bool + var missingModelType bool for _, col := range r { name := col["name"].String() if name == "is_active" { hasIsActive = true } if name == "api_key" && col["notnull"].Int() == 0 { - needsMigrate = true + apiKeyNullable = true } } - if hasIsActive { - needsMigrate = true + // 检测是否缺少 model_type 列 + if _, err := g.DB().Exec(ctx, "SELECT model_type FROM "+public.TableNameUserModelConfig+" LIMIT 1"); err != nil { + missingModelType = true } - if needsMigrate { - g.Log().Info(ctx, "检测到 user_model_config 列类型异常,开始迁移...") + // 检测是否缺少 base_url 列 + if _, err := g.DB().Exec(ctx, "SELECT base_url FROM "+public.TableNameUserModelConfig+" LIMIT 1"); err != nil { + missingBaseUrl = true + } + + if hasIsActive || missingBaseUrl || apiKeyNullable || missingModelType { + g.Log().Info(ctx, "检测到 user_model_config 表结构需要迁移,开始迁移...") newTable := public.TableNameUserModelConfig + "_new" - _, _ = g.DB().Exec(ctx, `CREATE TABLE `+newTable+` ( + + if _, err := g.DB().Exec(ctx, `CREATE TABLE `+newTable+` ( id INTEGER PRIMARY KEY AUTOINCREMENT, user_id INTEGER NOT NULL DEFAULT 0, model_config_id INTEGER NOT NULL DEFAULT 0, + model_type TEXT NOT NULL DEFAULT '', api_key TEXT NOT NULL DEFAULT '', + base_url TEXT NOT NULL DEFAULT '', + task_callback_url TEXT NOT NULL DEFAULT '', temperature REAL NOT NULL DEFAULT 0, max_tokens INTEGER NOT NULL DEFAULT 0, created_at DATETIME DEFAULT (datetime('now','localtime')), updated_at DATETIME DEFAULT (datetime('now','localtime')) - )`) - _, _ = g.DB().Exec(ctx, - `INSERT INTO `+newTable+`(user_id,model_config_id,api_key,temperature,max_tokens,created_at,updated_at) - SELECT user_id,COALESCE(model_config_id,0), - COALESCE(api_key,''), - COALESCE(temperature,0), - COALESCE(max_tokens,0), - COALESCE(created_at,datetime('now','localtime')), - COALESCE(updated_at,datetime('now','localtime')) - FROM `+public.TableNameUserModelConfig) - _, _ = g.DB().Exec(ctx, "DROP TABLE "+public.TableNameUserModelConfig) - _, _ = g.DB().Exec(ctx, "ALTER TABLE "+newTable+" RENAME TO "+public.TableNameUserModelConfig) + )`); err != nil { + g.Log().Warningf(ctx, "创建新表失败: %v", err) + return + } + + // 注意: 旧表中不存在 base_url/task_callback_url 列,只能从 model_config 表 JOIN 获取 + if _, err := g.DB().Exec(ctx, + `INSERT INTO `+newTable+`(user_id,model_config_id,model_type,api_key,base_url,task_callback_url,temperature,max_tokens,created_at,updated_at) + SELECT uc.user_id,uc.model_config_id, + COALESCE(mc.model_type,''), + COALESCE(uc.api_key,''), + COALESCE(mc.base_url,''), + COALESCE(mc.task_callback_url,''), + COALESCE(uc.temperature,0), + COALESCE(uc.max_tokens,0), + COALESCE(uc.created_at,datetime('now','localtime')), + COALESCE(uc.updated_at,datetime('now','localtime')) + FROM `+public.TableNameUserModelConfig+` uc + LEFT JOIN `+public.TableNameModelConfig+` mc ON mc.id = uc.model_config_id`); err != nil { + g.Log().Error(ctx, "迁移数据失败:", err) + // 不回滚,后续 DROP/RENAME 仍执行以防残留 + } + + if _, err := g.DB().Exec(ctx, "DROP TABLE "+public.TableNameUserModelConfig); err != nil { + g.Log().Warningf(ctx, "删除旧表失败: %v", err) + } + if _, err := g.DB().Exec(ctx, "ALTER TABLE "+newTable+" RENAME TO "+public.TableNameUserModelConfig); err != nil { + g.Log().Error(ctx, "重命名表失败:", err) + return + } g.Log().Info(ctx, "user_model_config 表结构迁移完成") } } @@ -116,12 +149,9 @@ func (d *userModelConfigDao) GetByModelConfigId(ctx context.Context, userId, mod // GetByModelType 获取指定用户对指定模型类型的配置 func (d *userModelConfigDao) GetByModelType(ctx context.Context, userId int64, modelType string) (res *entity.UserModelConfig, err error) { - r, err := g.DB().Model(public.TableNameUserModelConfig+" uc"). - Ctx(ctx). - InnerJoin(public.TableNameModelConfig+" mc", "mc.id = uc.model_config_id"). - Where("uc.user_id", userId). - Where("mc.model_type", modelType). - Fields("uc.*"). + r, err := g.DB().Model(public.TableNameUserModelConfig).Ctx(ctx). + Where("user_id", userId). + Where("model_type", modelType). Limit(1). One() if err != nil { @@ -136,15 +166,11 @@ func (d *userModelConfigDao) GetByModelType(ctx context.Context, userId int64, m } // DeleteSameType 删除同用户同模型类型的其他配置(保留当前配置) -func (d *userModelConfigDao) DeleteSameType(ctx context.Context, userId, modelConfigId, excludeId int64) error { +func (d *userModelConfigDao) DeleteSameType(ctx context.Context, userId int64, modelType string, excludeId int64) error { _, err := g.DB().Exec(ctx, `DELETE FROM `+public.TableNameUserModelConfig+` - WHERE user_id=? AND id!=? AND model_config_id IN ( - SELECT id FROM `+public.TableNameModelConfig+` WHERE model_type=( - SELECT model_type FROM `+public.TableNameModelConfig+` WHERE id=? - ) - )`, - userId, excludeId, modelConfigId) + WHERE user_id=? AND id!=? AND model_type=?`, + userId, excludeId, modelType) return err } @@ -159,6 +185,9 @@ func (d *userModelConfigDao) Save(ctx context.Context, data *entity.UserModelCon delete(m, "id") m["created_at"] = gtime.Now().Format("Y-m-d H:i:s") m["updated_at"] = gtime.Now().Format("Y-m-d H:i:s") - _, err := g.DB().Model(public.TableNameUserModelConfig).Ctx(ctx).Data(m).Insert() + lid, err := g.DB().Model(public.TableNameUserModelConfig).Ctx(ctx).Data(m).InsertAndGetId() + if err == nil && lid > 0 { + data.Id = lid + } return err } diff --git a/shortdrama/model/dto/config_dto.go b/shortdrama/model/dto/config_dto.go index d93d9fc..ad670a3 100644 --- a/shortdrama/model/dto/config_dto.go +++ b/shortdrama/model/dto/config_dto.go @@ -23,15 +23,13 @@ type GetModelConfigListRes struct { // SaveModelConfigReq 保存单个模型配置 type SaveModelConfigReq struct { - g.Meta `path:"/model" method:"post" tags:"模型配置" summary:"保存模型配置"` - Id int64 `json:"id"` - ModelType string `json:"modelType" v:"required" dc:"chat=对话模型 video=视频模型"` - ModelName string `json:"modelName" v:"required" dc:"模型名称"` - BaseUrl string `json:"baseUrl" v:"required" dc:"API接口地址"` - Schema *gjson.Json `json:"schema" dc:"请求体JSON Schema"` - Price int `json:"price" dc:"价格(分)"` - PriceUnit string `json:"priceUnit" dc:"价格单位(second=每秒/video=每次视频)"` - TaskCallbackUrl string `json:"taskCallbackUrl" dc:"回调地址(仅视频模型)"` + g.Meta `path:"/model" method:"post" tags:"模型配置" summary:"保存模型配置"` + Id int64 `json:"id"` + ModelType string `json:"modelType" v:"required" dc:"chat=对话模型 video=视频模型"` + ModelName string `json:"modelName" v:"required" dc:"模型名称"` + Schema *gjson.Json `json:"schema" dc:"请求体JSON Schema"` + Price int `json:"price" dc:"价格(分)"` + PriceUnit string `json:"priceUnit" dc:"价格单位(second=每秒/video=每次视频)"` } type GetUserModelConfigReq struct { @@ -43,12 +41,19 @@ type GetUserModelConfigRes struct { *entity.UserModelConfig } +// SaveUserModelConfigItem 单条用户模型配置 +type SaveUserModelConfigItem struct { + ModelConfigId int64 `json:"modelConfigId" dc:"关联模型配置ID"` + ApiKey string `json:"apiKey" dc:"用户API密钥"` + BaseUrl string `json:"baseUrl" dc:"API接口地址(用户覆盖)"` + TaskCallbackUrl string `json:"taskCallbackUrl" dc:"回调地址(仅视频模型,用户覆盖)"` + Temperature float64 `json:"temperature" dc:"温度参数"` + MaxTokens int `json:"maxTokens" dc:"最大Token数(不超过系统配置)"` +} + type SaveUserModelConfigReq struct { - g.Meta `path:"/user-model" method:"post" tags:"模型配置" summary:"保存用户模型配置"` - ModelConfigId int64 `json:"modelConfigId" dc:"关联模型配置ID"` - ApiKey string `json:"apiKey" dc:"用户API密钥"` - Temperature float64 `json:"temperature" dc:"温度参数"` - MaxTokens int `json:"maxTokens" dc:"最大Token数(不超过系统配置)"` + g.Meta `path:"/user-model" method:"post" tags:"模型配置" summary:"批量保存用户模型配置"` + Configs []*SaveUserModelConfigItem `json:"configs" dc:"模型配置列表"` } // GetUserModelListReq 获取用户模型配置列表(分页) @@ -62,14 +67,16 @@ type GetUserModelListReq struct { // UserModelItem 系统模型项及用户配置状态 type UserModelItem struct { - ModelConfigId int64 `json:"modelConfigId" dc:"系统模型配置ID"` - ModelType string `json:"modelType" dc:"chat=对话模型 video=视频模型"` - ModelName string `json:"modelName" dc:"模型名称"` - Configured bool `json:"configured" dc:"用户是否已配置"` - UserApiKey string `json:"userApiKey,omitempty" dc:"用户配置的API Key"` - Temperature float64 `json:"temperature,omitempty" dc:"用户配置的温度"` - UserMaxTokens int `json:"userMaxTokens,omitempty" dc:"用户配置的max_tokens"` - SystemMaxTokens int `json:"systemMaxTokens,omitempty" dc:"系统max_tokens上限"` + ModelConfigId int64 `json:"modelConfigId" dc:"系统模型配置ID"` + ModelType string `json:"modelType" dc:"chat=对话模型 video=视频模型"` + ModelName string `json:"modelName" dc:"模型名称"` + Configured bool `json:"configured" dc:"用户是否已配置"` + UserApiKey string `json:"userApiKey,omitempty" dc:"用户配置的API Key"` + UserBaseUrl string `json:"userBaseUrl,omitempty" dc:"用户配置的请求地址"` + UserTaskCallbackUrl string `json:"userTaskCallbackUrl,omitempty" dc:"用户配置的回调地址"` + Temperature float64 `json:"temperature,omitempty" dc:"用户配置的温度"` + UserMaxTokens int `json:"userMaxTokens,omitempty" dc:"用户配置的max_tokens"` + SystemMaxTokens int `json:"systemMaxTokens,omitempty" dc:"系统max_tokens上限"` } // GetUserModelListRes 用户模型配置列表响应 diff --git a/shortdrama/model/entity/model_config.go b/shortdrama/model/entity/model_config.go index f612f92..1fe0037 100644 --- a/shortdrama/model/entity/model_config.go +++ b/shortdrama/model/entity/model_config.go @@ -10,12 +10,12 @@ type ModelConfig struct { ModelType string `orm:"model_type" json:"modelType" dc:"chat=对话模型 video=视频模型"` ModelName string `orm:"model_name" json:"modelName" dc:"模型名称"` ApiKey string `json:"apiKey" dc:"API密钥"` - BaseUrl string `orm:"base_url" json:"baseUrl" dc:"API接口地址"` Schema string `orm:"schema" json:"schema" dc:"请求体JSON Schema"` MaxTokens int `orm:"max_tokens" json:"maxTokens" dc:"最大Token数(仅对话模型)"` Price int `orm:"price" json:"price" dc:"价格(分)"` PriceUnit string `orm:"price_unit" json:"priceUnit" dc:"价格单位(second=每秒/video=每次视频)"` - TaskCallbackUrl string `orm:"task_callback_url" json:"taskCallbackUrl" dc:"回调地址(仅视频模型)"` + BaseUrl string `json:"baseUrl" dc:"API接口地址(仅运行时合并用户配置用,不持久化)"` + TaskCallbackUrl string `json:"taskCallbackUrl" dc:"回调地址(仅运行时合并用户配置用,不持久化)"` Temperature float64 `json:"temperature" dc:"温度参数(仅运行时合并用户配置用,不持久化)"` CreatedAt *gtime.Time `orm:"created_at" json:"createdAt" dc:"创建时间"` UpdatedAt *gtime.Time `orm:"updated_at" json:"updatedAt" dc:"更新时间"` diff --git a/shortdrama/model/entity/user_model_config.go b/shortdrama/model/entity/user_model_config.go index aa03322..0240fca 100644 --- a/shortdrama/model/entity/user_model_config.go +++ b/shortdrama/model/entity/user_model_config.go @@ -6,12 +6,15 @@ import "github.com/gogf/gf/v2/os/gtime" // 每行对应一个用户对一个系统模型的 API Key 等私有配置。 // 每个模型类型(chat/video)最多保留一条用户配置。 type UserModelConfig struct { - Id int64 `orm:"id" json:"id" dc:"配置ID"` - UserId int64 `orm:"user_id" json:"userId" dc:"用户ID"` - ModelConfigId int64 `orm:"model_config_id" json:"modelConfigId" dc:"关联模型配置ID"` - ApiKey string `orm:"api_key" json:"apiKey" dc:"用户API密钥"` - Temperature float64 `orm:"temperature" json:"temperature" dc:"温度参数"` - MaxTokens int `orm:"max_tokens" json:"maxTokens" dc:"最大Token数(不超过系统配置)"` - CreatedAt *gtime.Time `orm:"created_at" json:"createdAt" dc:"创建时间"` - UpdatedAt *gtime.Time `orm:"updated_at" json:"updatedAt" dc:"更新时间"` + Id int64 `orm:"id" json:"id" dc:"配置ID"` + UserId int64 `orm:"user_id" json:"userId" dc:"用户ID"` + ModelConfigId int64 `orm:"model_config_id" json:"modelConfigId" dc:"关联模型配置ID"` + ModelType string `orm:"model_type" json:"modelType" dc:"模型类型(chat/video)"` + ApiKey string `orm:"api_key" json:"apiKey" dc:"用户API密钥"` + BaseUrl string `orm:"base_url" json:"baseUrl" dc:"API接口地址(用户覆盖)"` + TaskCallbackUrl string `orm:"task_callback_url" json:"taskCallbackUrl" dc:"回调地址(仅视频模型,用户覆盖)"` + Temperature float64 `orm:"temperature" json:"temperature" dc:"温度参数"` + MaxTokens int `orm:"max_tokens" json:"maxTokens" dc:"最大Token数(不超过系统配置)"` + CreatedAt *gtime.Time `orm:"created_at" json:"createdAt" dc:"创建时间"` + UpdatedAt *gtime.Time `orm:"updated_at" json:"updatedAt" dc:"更新时间"` } diff --git a/shortdrama/service/config_service.go b/shortdrama/service/config_service.go index b68549f..405f51c 100644 --- a/shortdrama/service/config_service.go +++ b/shortdrama/service/config_service.go @@ -151,6 +151,63 @@ func (s *configService) GetUserConfig(ctx context.Context, userId int64, modelTy return cfg } +// SaveUserConfigs 批量保存用户模型配置 +// 先查询用户已有全部配置,再按 model_type 判断更新或新增 +func (s *configService) SaveUserConfigs(ctx context.Context, userId int64, items []*dto.SaveUserModelConfigItem) error { + // 加载该用户所有现有配置(按 model_type 索引) + existingList, _ := dao.UserModelConfig.GetByUserId(ctx, userId) + existingByType := make(map[string]*entity.UserModelConfig) + for _, ec := range existingList { + if ec.ModelType != "" { + existingByType[ec.ModelType] = ec + } + } + + // 加载系统模型列表(用于查询 model_type) + modelList := s.getModelList(ctx) + modelTypeMap := make(map[int64]string) + for _, m := range modelList { + modelTypeMap[m.Id] = m.ModelType + } + + for _, item := range items { + modelType := modelTypeMap[item.ModelConfigId] + cfg := &entity.UserModelConfig{ + ModelConfigId: item.ModelConfigId, + ModelType: modelType, + ApiKey: item.ApiKey, + BaseUrl: item.BaseUrl, + TaskCallbackUrl: item.TaskCallbackUrl, + Temperature: item.Temperature, + MaxTokens: item.MaxTokens, + } + + if existing, ok := existingByType[modelType]; ok { + // 同类型已有配置 → 更新(保留 id,覆盖字段) + existing.ModelConfigId = item.ModelConfigId + existing.ApiKey = item.ApiKey + existing.BaseUrl = item.BaseUrl + existing.TaskCallbackUrl = item.TaskCallbackUrl + existing.Temperature = item.Temperature + existing.MaxTokens = item.MaxTokens + if err := dao.UserModelConfig.Save(ctx, existing); err != nil { + return err + } + } else { + // 无同类配置 → 新增 + cfg.UserId = userId + if err := dao.UserModelConfig.Save(ctx, cfg); err != nil { + return err + } + } + } + + // 清除该用户所有模型类型的配置缓存 + _, _ = gcache.Remove(ctx, cacheKeyUserConfigPrefix+gconv.String(userId)+"_chat") + _, _ = gcache.Remove(ctx, cacheKeyUserConfigPrefix+gconv.String(userId)+"_video") + return nil +} + // SaveUserConfig 保存用户私有模型配置 // 如果已存在 user_id + model_config_id 的记录则更新,否则新增 // 保存时自动清理同类型其他配置 @@ -166,6 +223,12 @@ func (s *configService) SaveUserConfig(ctx context.Context, userId int64, cfg *e if cfg.ApiKey == "" { cfg.ApiKey = existing.ApiKey } + if cfg.BaseUrl == "" { + cfg.BaseUrl = existing.BaseUrl + } + if cfg.TaskCallbackUrl == "" { + cfg.TaskCallbackUrl = existing.TaskCallbackUrl + } if cfg.Temperature == 0 { cfg.Temperature = existing.Temperature } @@ -179,7 +242,7 @@ func (s *configService) SaveUserConfig(ctx context.Context, userId int64, cfg *e } // 清理同模型类型的其他用户配置 - if err := dao.UserModelConfig.DeleteSameType(ctx, userId, cfg.ModelConfigId, cfg.Id); err != nil { + if err := dao.UserModelConfig.DeleteSameType(ctx, userId, cfg.ModelType, cfg.Id); err != nil { g.Log().Error(ctx, "清理同类型模型配置失败:", err) } // 清除该用户所有模型类型的配置缓存 @@ -202,10 +265,16 @@ func (s *configService) GetMergedConfig(ctx context.Context, userId int64, model return &merged } - // 用户 API Key 覆盖系统 API Key + // 用户 API Key / BaseUrl / TaskCallbackUrl 覆盖系统配置 if userCfg.ApiKey != "" { merged.ApiKey = userCfg.ApiKey } + if userCfg.BaseUrl != "" { + merged.BaseUrl = userCfg.BaseUrl + } + if userCfg.TaskCallbackUrl != "" { + merged.TaskCallbackUrl = userCfg.TaskCallbackUrl + } // Temperature 从用户配置读取(仅 chat 模型) if userCfg.Temperature > 0 { merged.Temperature = userCfg.Temperature @@ -244,8 +313,10 @@ func (s *configService) GetUserModelList(ctx context.Context, userId int64) *dto SystemMaxTokens: sysMax, } if uc, ok := cfgMap[m.Id]; ok { - item.Configured = uc.ApiKey != "" + item.Configured = true item.UserApiKey = uc.ApiKey + item.UserBaseUrl = uc.BaseUrl + item.UserTaskCallbackUrl = uc.TaskCallbackUrl item.Temperature = uc.Temperature item.UserMaxTokens = uc.MaxTokens } @@ -278,8 +349,10 @@ func (s *configService) GetUserModelListPage(ctx context.Context, userId int64, SystemMaxTokens: sysMax, } if uc, ok := cfgMap[m.Id]; ok { - item.Configured = uc.ApiKey != "" + item.Configured = true item.UserApiKey = uc.ApiKey + item.UserBaseUrl = uc.BaseUrl + item.UserTaskCallbackUrl = uc.TaskCallbackUrl item.Temperature = uc.Temperature item.UserMaxTokens = uc.MaxTokens } diff --git a/shortdrama/service/drama_service.go b/shortdrama/service/drama_service.go index 9d16491..07a3c31 100644 --- a/shortdrama/service/drama_service.go +++ b/shortdrama/service/drama_service.go @@ -1284,21 +1284,9 @@ func (s *dramaService) pollPendingVideos(ctx context.Context) { } -// buildPollURL 从视频创建 URL 构建任务查询 URL(DashScope API 规范) -// 创建端点: /api/v1/services/aigc/video-generation/video-synthesis -// 查询端点: /api/v1/tasks/{taskId} -func buildPollURL(baseURL, taskId string) string { - // 查找 "/api/v1/" 路径并替换后面的内容为 tasks/{taskId} - if idx := strings.Index(baseURL, "/api/v1/"); idx > 0 { - return baseURL[:idx] + "/api/v1/tasks/" + taskId - } - // fallback:直接拼接 - return strings.TrimRight(baseURL, "/") + "/" + taskId -} - // pollVideoTaskOnce 单次查询视频生成任务状态 func (s *dramaService) pollVideoTaskOnce(ctx context.Context, modelCfg *entity.ModelConfig, taskId string) (string, error) { - queryURL := buildPollURL(modelCfg.BaseUrl, taskId) + queryURL := strings.TrimRight(modelCfg.TaskCallbackUrl, "/") + "/" + taskId req, err := http.NewRequestWithContext(ctx, "GET", queryURL, nil) if err != nil { @@ -1316,7 +1304,10 @@ func (s *dramaService) pollVideoTaskOnce(ctx context.Context, modelCfg *entity.M data, _ := io.ReadAll(resp.Body) var result struct { - Output struct { + Status string `json:"status"` + TaskStatus string `json:"task_status"` + VideoURL string `json:"video_url"` + Output *struct { TaskStatus string `json:"task_status"` VideoURL string `json:"video_url"` Code string `json:"code"` @@ -1331,55 +1322,46 @@ func (s *dramaService) pollVideoTaskOnce(ctx context.Context, modelCfg *entity.M return "", fmt.Errorf("解析响应失败: %s", string(data)) } - videoURL := result.Output.VideoURL - if videoURL == "" && len(result.Output.Results) > 0 { - videoURL = result.Output.Results[0].VideoURL - if videoURL == "" { - videoURL = result.Output.Results[0].URL - } + // 兼容多种响应格式:status / task_status / output.task_status + status := result.Status + if status == "" { + status = result.TaskStatus + } + if status == "" && result.Output != nil { + status = result.Output.TaskStatus } - status := mapVideoTaskStatus(result.Output.TaskStatus) - buildErrMsg := func() string { - c, m := result.Output.Code, result.Output.Message - if c != "" || m != "" { - return fmt.Sprintf("%s(code=%s, msg=%s)", result.Output.TaskStatus, c, m) + // 获取视频URL + videoURL := result.VideoURL + if videoURL == "" && result.Output != nil { + videoURL = result.Output.VideoURL + if videoURL == "" && len(result.Output.Results) > 0 { + videoURL = result.Output.Results[0].VideoURL + if videoURL == "" { + videoURL = result.Output.Results[0].URL + } } - return result.Output.TaskStatus } switch status { - case videoTaskSucceeded: + case "SUCCEEDED", "succeeded": if videoURL != "" { return videoURL, nil } - // SUCCEEDED 但没有视频URL,输出完整响应帮助调试 g.Log().Warningf(ctx, "任务 %s 状态为 SUCCEEDED 但未返回视频URL,完整响应: %s", taskId, string(data)) return "", fmt.Errorf("SUCCEEDED但视频URL为空,完整响应: %s", string(data)) - case videoTaskFailed: - return "", fmt.Errorf("FAILED: %s", buildErrMsg()) - case videoTaskRunning, videoTaskPending: + case "FAILED", "failed": + errMsg := status + if result.Output != nil && result.Output.Message != "" { + errMsg = result.Output.Message + } + return "", fmt.Errorf("FAILED: %s", errMsg) + case "RUNNING", "running", "PENDING", "pending": return "", fmt.Errorf("RUNNING") } - g.Log().Warningf(ctx, "任务 %s 返回未知状态,完整响应: %s", taskId, string(data)) - return "", fmt.Errorf("未知任务状态: %s,完整响应: %s", result.Output.TaskStatus, string(data)) -} - -// mapVideoTaskStatus 将供应商任务状态映射为内部状态 -func mapVideoTaskStatus(s string) videoTaskStatus { - switch s { - case "PENDING": - return videoTaskPending - case "RUNNING": - return videoTaskRunning - case "SUCCEEDED": - return videoTaskSucceeded - case "FAILED": - return videoTaskFailed - default: - return videoTaskUnknown - } + g.Log().Warningf(ctx, "任务 %s 返回未知状态 %s,完整响应: %s", taskId, status, string(data)) + return "", fmt.Errorf("未知任务状态: %s,完整响应: %s", status, string(data)) } // ==================== Video Generation ==================== @@ -1619,7 +1601,6 @@ func createVideoTask(ctx context.Context, apiKey, baseURL, modelName, prompt, ne } req.Header.Set("Authorization", "Bearer "+apiKey) req.Header.Set("Content-Type", "application/json") - req.Header.Set("X-DashScope-Async", "enable") client := &http.Client{Timeout: 30 * time.Second} resp, err := client.Do(req) @@ -1630,22 +1611,26 @@ func createVideoTask(ctx context.Context, apiKey, baseURL, modelName, prompt, ne respData, _ := io.ReadAll(resp.Body) var result struct { - Output struct { + ID string `json:"id"` + TaskID string `json:"task_id"` + Output *struct { TaskID string `json:"task_id"` } `json:"output"` - Code string `json:"code"` - Message string `json:"message"` } if err := json.Unmarshal(respData, &result); err != nil { return "", "", fmt.Errorf("解析响应失败: %s", string(respData)) } - if result.Code != "" { - return "", "", fmt.Errorf("请求失败(code=%s): %s", result.Code, string(respData)) + taskID := result.ID + if taskID == "" { + taskID = result.TaskID } - if result.Output.TaskID == "" { - return "", "", fmt.Errorf("任务ID为空") + if taskID == "" && result.Output != nil { + taskID = result.Output.TaskID } - return result.Output.TaskID, string(payload), nil + if taskID == "" { + return "", "", fmt.Errorf("无法从响应中获取任务ID: %s", string(respData)) + } + return taskID, string(payload), nil } // resubmitVideoTask 使用已有的 bodyJSON 重新提交视频任务(跳过 Agent,直接调 API) @@ -1681,7 +1666,6 @@ func resubmitVideoTask(ctx context.Context, apiKey, baseURL string, bodyJSON []b } req.Header.Set("Authorization", "Bearer "+apiKey) req.Header.Set("Content-Type", "application/json") - req.Header.Set("X-DashScope-Async", "enable") client := &http.Client{Timeout: 30 * time.Second} resp, err := client.Do(req) @@ -1692,22 +1676,26 @@ func resubmitVideoTask(ctx context.Context, apiKey, baseURL string, bodyJSON []b respData, _ := io.ReadAll(resp.Body) var result struct { - Output struct { + ID string `json:"id"` + TaskID string `json:"task_id"` + Output *struct { TaskID string `json:"task_id"` } `json:"output"` - Code string `json:"code"` - Message string `json:"message"` } if err := json.Unmarshal(respData, &result); err != nil { return "", nil, fmt.Errorf("解析响应失败: %s", string(respData)) } - if result.Code != "" { - return "", nil, fmt.Errorf("请求失败(code=%s): %s", result.Code, string(respData)) + taskID := result.ID + if taskID == "" { + taskID = result.TaskID } - if result.Output.TaskID == "" { - return "", nil, fmt.Errorf("任务ID为空") + if taskID == "" && result.Output != nil { + taskID = result.Output.TaskID } - return result.Output.TaskID, payload, nil + if taskID == "" { + return "", nil, fmt.Errorf("无法从响应中获取任务ID: %s", string(respData)) + } + return taskID, payload, nil } // resolveSizeFromSchema 从 video_schema.sizes 中根据分辨率和宽高比查找尺寸字符串